diff --git a/.gitignore b/.gitignore index 3df4381..b29e834 100644 --- a/.gitignore +++ b/.gitignore @@ -13,6 +13,7 @@ predictions/ predictions.csv *.shards/ clusters.csv +*.duckdb # Local working notes, not part of the repo CLAUDE.md diff --git a/README.md b/README.md index e612b9e..981b282 100644 --- a/README.md +++ b/README.md @@ -108,6 +108,7 @@ result = linker.cluster("predictions.parquet", truth="sample.truth.csv") print(result.quality) ``` +`cpplink-viewer` (`pip install "cpplink[viewer]"`) then serves the clusters as a page: each cluster's members side by side with every disagreeing cell highlighted, and every prediction behind it as the ledger the scorer produced. Between `init` and `estimate` sit the diagnostics that cost seconds and decide the quality of the result: `profile` says what each column can be worth before any model exists, `levels` places the fuzzy thresholds from the data, `explain-blocking` prices every source without enumerating a pair, and `recall` measures what blocking reaches. See [Getting started](https://4ment.github.io/cpplink/getting-started/) for the whole pipeline with its output explained, [Commands](https://4ment.github.io/cpplink/commands/) for every command at a glance, the [schema reference](https://4ment.github.io/cpplink/reference/schema/) for every field, and [From Python](https://4ment.github.io/cpplink/python/) for the package. diff --git a/docs/commands/cluster.md b/docs/commands/cluster.md index be6dd39..3649ed3 100644 --- a/docs/commands/cluster.md +++ b/docs/commands/cluster.md @@ -223,6 +223,10 @@ The `cluster_id` is the representative record's own `unique_id`, so the output s record the others collapse onto. Only records in a cluster of at least `--min-size` are written; singletons are excluded by default. +## Reading the clusters + +[The cluster viewer](../viewer.md) serves this file beside the predictions and their waterfalls: the members of each cluster side by side with every disagreeing cell highlighted, and every prediction touching it, including the ones clustering at a higher threshold overruled. + ## Cost The union–find is `uint32` parent plus `uint8` rank — **5 bytes a record**, 100 MB at 20M rows — diff --git a/docs/commands/explain.md b/docs/commands/explain.md index 0b117e8..cbe9d13 100644 --- a/docs/commands/explain.md +++ b/docs/commands/explain.md @@ -216,8 +216,8 @@ A level's label, `m` and `u` are the model's rather than the pair's, so the file Ids are resolved through a sorted index over the id column, 4 bytes a record, so a file of millions of predictions costs one load and one pass. A prediction naming an id no record holds is skipped and counted, and so is one whose stored `gamma` the comparisons no longer produce, which means the schema changed after `predict` ran and the ledgers explain today's schema rather than the file's weights. -This is what the cluster viewers in `tools/` draw from: `tools/cluster_view.py --waterfalls` embeds a ledger per prediction, and `tools/cluster_server.py --waterfalls` loads the file beside the predictions so a click is one lookup. -Nothing in either recomputes a bit of the weight, and neither runs this binary. +This is what [the cluster viewer](../viewer.md) draws from: `cpplink-viewer --waterfalls` loads the file beside the predictions so a click is one lookup. +Nothing in it recomputes a bit of the weight, and it does not run this binary. ## See also diff --git a/docs/python.md b/docs/python.md index 848d1da..7cf624f 100644 --- a/docs/python.md +++ b/docs/python.md @@ -168,6 +168,12 @@ None of them calls back into Python, so another thread can run while they do. Errors the core reports as `false` and a message become `cpplink.Error`, a `RuntimeError` carrying the core's own text. +## The cluster viewer + +`cpplink-viewer`, the `cpplink_viewer` package in the same wheel, serves the clusters a run produced as a page; `pip install "cpplink[viewer]"` adds DuckDB, its one dependency beyond the binding's. +It reads the files the pipeline wrote and never the compiled module, so it also runs from a checkout with nothing built. +See [The cluster viewer](viewer.md). + ## Not in this version - **In-memory input.** The `Linker` reads parquet files, as the command does. A pyarrow, polars or pandas table through `__arrow_c_stream__` is the next step, and needs the loader split so a table can be appended without a file. diff --git a/docs/viewer.md b/docs/viewer.md new file mode 100644 index 0000000..fe7256a --- /dev/null +++ b/docs/viewer.md @@ -0,0 +1,67 @@ +# The cluster viewer + +`cpplink-viewer` serves a page over the clusters a run produced: a list of clusters on the left, and for the one picked, its members side by side with every disagreeing cell highlighted, and every prediction touching it with the one picked drawn as the ledger the scorer produced. +It reads what the pipeline wrote and runs none of it: the clusters and the predictions come from the files, the waterfalls from `explain --predictions`, and nothing in the viewer recomputes a bit of a weight. + +```sh +cpplink predict --schema schema.json --model model.json --out predictions.parquet sample.parquet +cpplink cluster --schema schema.json --predictions predictions.parquet --out clusters.csv sample.parquet +cpplink explain --schema schema.json --model model.json --predictions predictions.parquet --out waterfalls.parquet sample.parquet +cpplink-viewer --schema schema.json --clusters clusters.csv --predictions predictions.parquet \ + --waterfalls waterfalls.parquet --model model.json --truth sample.truth.csv --open sample.parquet +``` + +The page is served at `http://127.0.0.1:8770/` (`--port` moves it, `--open` opens a browser on it). + +## Install + +The viewer is the `cpplink_viewer` package, shipped in the same wheel as the binding, with DuckDB as its one extra dependency: + +```sh +pip install "cpplink[viewer]" +``` + +It never imports the compiled module, so it also runs from a checkout with nothing built, against a run the command line produced: + +```sh +pip install duckdb pyarrow numpy +PYTHONPATH=python python -m cpplink_viewer --schema schema.json --clusters clusters.csv sample.parquet +``` + +## What it reads + +| Option | File | Written by | +|---|---|---| +| `data` | the parquet file(s) the run read, in order | you | +| `--schema` | the schema, for the id column and the columns to show | `init` | +| `--clusters` | `unique_id, cluster_id, cluster_size`, csv or parquet | `cluster --out` | +| `--predictions` | one csv or parquet file, or the shard directory | `predict --out` | +| `--waterfalls` | one wide row per prediction, csv or parquet | `explain --predictions --out` | +| `--model` | the level labels and rates the waterfall shows | `estimate --out` | +| `--truth` | known pairs, to colour members by true entity | you, or `gen-sample --truth` | + +Only `--schema` and the data are required, plus one of `--clusters` or `--threshold`. +Without `--predictions` the page shows the members and nothing about why they are together; without `--waterfalls` and `--model` it shows each prediction's weight but not the ledger behind it. +A shard directory names records by row, so it costs one pass over the id column that a merged file does not. + +## The cache + +The first run builds a DuckDB file (`--cache`, `cpplink_viewer.duckdb` by default) holding only what the page can ever show: the clustered records, never the singletons, their predictions, the waterfalls of those predictions, and one row of statistics per cluster. +At 20M records with 2.9M of them clustered the parquet is read exactly once and every later request touches a table seven times smaller than the file. +The cache is keyed on every input's path, size and modification time and on the options that shape it, so a changed input rebuilds it and an unchanged one is reused; `--rebuild` forces it. +`--memory` and `--threads` are DuckDB's limits while building. + +## Thresholds and rejected predictions + +`--threshold` keeps only the predictions at or above it. +With `--clusters`, the file is taken as what `cluster --threshold` wrote at that threshold and read as it stands. +Without one, the predictions are clustered here by the same union-find in the same order, so the cache holds exactly the partition that command would write, named by the same representatives; the test suite holds the two to the same file. + +A prediction whose two records clustering put in different clusters, above the write threshold and below the clustering one, is the prediction most worth reading, so it is kept and listed under both of its clusters as `rejected`, each naming the other cluster its second record went to. +`--min-size` and `--max-size` bound the clusters listed; `--max-rows` bounds the members shown per cluster, the rest being counted only. + +## The list + +The list is sorted by disagreement (columns whose values differ within the cluster), size, weakest prediction or id, and filtered to every cluster, the split ones, the ones held together only transitively (fewer predictions than pairs), or the ones mixing entities against the truth file. +The search box reads any value, one column, the cluster id, or a comma-separated list of record ids matched whole, in which case it says which of them named a record. +The `network` checkbox draws the selected cluster's predictions as a graph, which is where a chain shows itself as a chain. diff --git a/environment.yml b/environment.yml index 150addd..f1a660a 100644 --- a/environment.yml +++ b/environment.yml @@ -21,3 +21,5 @@ dependencies: - numpy - pytest - ruff + # The viewer's cache; the `viewer` extra of the wheel. + - python-duckdb diff --git a/mkdocs.yml b/mkdocs.yml index 7f453d8..07e9299 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -36,6 +36,7 @@ nav: - Home: index.md - Getting started: getting-started.md - From Python: python.md + - The cluster viewer: viewer.md - Concepts: - The model: model.md - Comparisons: comparisons.md diff --git a/pyproject.toml b/pyproject.toml index 67e2b05..aae8c5a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -14,6 +14,11 @@ dependencies = ["numpy", "pyarrow"] [project.optional-dependencies] test = ["pytest"] +# `cpplink-viewer` serves the clusters of a run over a DuckDB cache. +viewer = ["duckdb"] + +[project.scripts] +cpplink-viewer = "cpplink_viewer.server:main" [project.urls] Homepage = "https://github.com/4ment/cpplink" @@ -26,7 +31,7 @@ metadata.version.input = "CMakeLists.txt" metadata.version.regex = 'project\(cpplink VERSION (?P[0-9.]+)' cmake.args = ["-DCPPLINK_BUILD_PYTHON=ON"] cmake.build-type = "Release" -wheel.packages = ["python/cpplink"] +wheel.packages = ["python/cpplink", "python/cpplink_viewer"] # Keep the extension's build out of build/, which is the CLI's. build-dir = "build/python" @@ -62,4 +67,4 @@ indent-style = "space" line-ending = "auto" [tool.ruff.lint.isort] -known-first-party = ["cpplink"] +known-first-party = ["cpplink", "cpplink_viewer"] diff --git a/python/cpplink_viewer/__init__.py b/python/cpplink_viewer/__init__.py new file mode 100644 index 0000000..1cfbb94 --- /dev/null +++ b/python/cpplink_viewer/__init__.py @@ -0,0 +1,13 @@ +# Copyright 2026 Mathieu Fourment +# SPDX-License-Identifier: MIT +"""``cpplink_viewer``: the cluster viewer, served over a DuckDB cache of a run. + +It reads what a run wrote and never the compiled module, so it runs from a +checkout with no build (``PYTHONPATH=python python -m cpplink_viewer``) as well +as from the wheel (``cpplink-viewer``). DuckDB is the ``viewer`` extra. +""" + +from .cache import Options, open_cache +from .server import Viewer, main, make_server + +__all__ = ["Options", "Viewer", "main", "make_server", "open_cache"] diff --git a/python/cpplink_viewer/__main__.py b/python/cpplink_viewer/__main__.py new file mode 100644 index 0000000..de4eba5 --- /dev/null +++ b/python/cpplink_viewer/__main__.py @@ -0,0 +1,10 @@ +# Copyright 2026 Mathieu Fourment +# SPDX-License-Identifier: MIT +"""``python -m cpplink_viewer [options] data.parquet``: the viewer, served.""" + +import sys + +from .server import main + +if __name__ == "__main__": + sys.exit(main()) diff --git a/python/cpplink_viewer/cache.py b/python/cpplink_viewer/cache.py new file mode 100644 index 0000000..c64862b --- /dev/null +++ b/python/cpplink_viewer/cache.py @@ -0,0 +1,531 @@ +# Copyright 2026 Mathieu Fourment +# SPDX-License-Identifier: MIT +"""The DuckDB cache the viewer serves from, built once per set of inputs. + +It holds only what the page can ever show: the clustered records, never the +singletons, so it is a fraction of the dataset. At 20M records with 2.9M of them +clustered the parquet is read exactly once and every later request touches a +table seven times smaller than the file. +""" + +from __future__ import annotations + +import json +import os +import time +from dataclasses import dataclass + +import numpy as np +import pyarrow as pa +import pyarrow.compute as pc +import pyarrow.csv as pacsv +import pyarrow.parquet as pq + +from .inputs import EDGE_DTYPE, EDGE_MAGIC, read_truth, schema_columns + +EDGE_CHUNK = 1 << 20 +CACHE_VERSION = 8 + + +@dataclass +class Options: + """What the cache is built from and how; the command line fills one.""" + + data: list[str] + schema: str + clusters: str | None = None + predictions: str | None = None + truth: str | None = None + waterfalls: str | None = None + model: str | None = None + cache: str = "cpplink_viewer.duckdb" + rebuild: bool = False + min_size: int = 2 + max_size: int = 0 + threshold: float | None = None + max_rows: int = 200 + memory: str = "4GB" + threads: int = 0 + + def check(self): + """The reason these options cannot build a cache, or None.""" + if bool(self.waterfalls) != bool(self.model): + return "--waterfalls and --model go together" + if self.waterfalls and not self.predictions: + return "--waterfalls needs --predictions" + if self.threshold is not None: + if not self.predictions: + return "--threshold prunes the predictions, so it needs --predictions" + if not self.clusters and self.min_size < 2: + return ( + "--threshold lists only what a prediction reaches, so a " + "singleton is never shown; --min-size must be at least 2" + ) + elif not self.clusters: + return "--clusters is needed without --threshold" + return None + + +def connect(path, read_only=False): + """A DuckDB connection, naming the extra to install when there is no DuckDB.""" + try: + import duckdb + except ImportError as missing: + raise SystemExit( + "the viewer needs duckdb: pip install duckdb, or pip install cpplink[viewer]" + ) from missing + return duckdb.connect(path, read_only=read_only) + + +def quoted(name): + return '"' + name.replace('"', '""') + '"' + + +def scan(path, types): + """The duckdb table function that reads one csv or parquet file. + + Both of a run's outputs are one file whose extension picks the format, so + that is what picks the reader here. The ids are pinned to VARCHAR rather than + sniffed, because a file of numeric ids would otherwise read as integers and + join against nothing. + """ + path = os.path.abspath(path) + if path.endswith(".parquet"): + return f"read_parquet('{path}')" + pinned = ", ".join(f"'{name}': 'VARCHAR'" for name in types) + return f"read_csv('{path}', header = true, types = {{{pinned}}})" + + +def fingerprint(opts, columns): + """What the cache was built from, so a changed input rebuilds it.""" + parts = [ + CACHE_VERSION, + columns, + opts.max_rows, + opts.min_size, + opts.max_size, + opts.threshold, + ] + for path in [opts.clusters, opts.truth, opts.waterfalls] + list(opts.data): + if path: + parts.append( + [os.path.abspath(path), os.path.getmtime(path), os.path.getsize(path)] + ) + if opts.predictions: + if os.path.isdir(opts.predictions): + for name in sorted(os.listdir(opts.predictions)): + if name.endswith(".bin"): + full = os.path.join(opts.predictions, name) + parts.append([full, os.path.getmtime(full), os.path.getsize(full)]) + else: + parts.append( + [ + os.path.abspath(opts.predictions), + os.path.getmtime(opts.predictions), + os.path.getsize(opts.predictions), + ] + ) + return json.dumps(parts, sort_keys=True) + + +def shard_blocks(directory): + """Each shard of a directory as one array of its records, sorted by name, + which is the order `cpplink cluster` reads them in.""" + for name in sorted(n for n in os.listdir(directory) if n.endswith(".bin")): + with open(os.path.join(directory, name), "rb") as handle: + if handle.read(len(EDGE_MAGIC)) != EDGE_MAGIC: + raise SystemExit(f"{name}: not a cpplink prediction shard") + while True: + data = handle.read(EDGE_CHUNK * EDGE_DTYPE.itemsize) + if not data: + break + yield np.frombuffer(data, dtype=EDGE_DTYPE) + + +def load_shard_edges(conn, directory): + """Every shard edge into `raw_edges`, a chunk at a time. + + The shards hold row indices rather than ids, which is why this is worth + doing once: the join that names them costs a pass over the record ids, and + a served request must not pay it. + """ + conn.execute("CREATE TABLE raw_edges (a UINTEGER, b UINTEGER, w DOUBLE)") + total = 0 + for block in shard_blocks(directory): + chunk = pa.table({"a": block["a"], "b": block["b"], "w": block["w"]}) + conn.register("chunk", chunk) + conn.execute("INSERT INTO raw_edges SELECT a, b, w FROM chunk") + conn.unregister("chunk") + total += block.size + return total + + +def read_shard_edges(directory, floor): + """Every shard edge at or above `floor` as one (a, b, w) table of rows, in + the order `cpplink cluster` reads them: shard by shard, sorted by name.""" + blocks = [block[block["w"] >= floor] for block in shard_blocks(directory)] + block = np.concatenate(blocks) if blocks else np.empty(0, dtype=EDGE_DTYPE) + return pa.table({"a": block["a"], "b": block["b"], "w": block["w"]}) + + +def read_file_edges(path, floor): + """Every prediction of a merged file at or above `floor` as one (a, b, w) + table of ids, in file order, which DuckDB would not promise once the build + connection stops preserving it.""" + if path.endswith(".parquet"): + table = pq.read_table(path, columns=["id_a", "id_b", "match_weight"]) + else: + table = pacsv.read_csv( + path, + convert_options=pacsv.ConvertOptions( + include_columns=["id_a", "id_b", "match_weight"], + column_types={ + "id_a": pa.string(), + "id_b": pa.string(), + "match_weight": pa.float64(), + }, + ), + ) + table = table.filter(pc.greater_equal(table["match_weight"], floor)) + return pa.table( + { + "a": table["id_a"].cast(pa.string()), + "b": table["id_b"].cast(pa.string()), + "w": table["match_weight"], + } + ) + + +def id_rows(conn, files, id_column, keys, by): + """`(uid, rid)` of the records `keys` names, by uid or by rid. + + One pass over the id column, which is the index `cpplink cluster` builds to + read a merged file back, and what a shard needs to name its rows. + """ + conn.register("keys", pa.table({"key": keys})) + parts, offset = [], 0 + for path in files: + parts.append(f""" + SELECT CAST({quoted(id_column)} AS VARCHAR) AS uid, + {offset} + file_row_number AS rid + FROM read_parquet('{path}', file_row_number = true)""") + offset += pq.ParquetFile(path).metadata.num_rows + found = conn.execute(f""" + SELECT uid, rid FROM ({" UNION ALL ".join(parts)}) + WHERE {by} IN (SELECT key FROM keys) + """).to_arrow_table() + conn.unregister("keys") + return found + + +def union_find(edges, count): + """The core's union-find, edge for edge: union by rank, a tie keeping the + first end's root, so each cluster comes out named by the representative + `cpplink cluster` names it by. Returns the root of every index. + + This is the one piece of the core the viewer reimplements, and a test holds + it to the partition `cpplink cluster --threshold` writes. + """ + parent = list(range(count)) + rank = [0] * count + + def find(x): + while parent[x] != x: + parent[x] = parent[parent[x]] + x = parent[x] + return x + + for a, b in edges: + ra, rb = find(a), find(b) + if ra == rb: + continue + if rank[ra] < rank[rb]: + ra, rb = rb, ra + parent[rb] = ra + if rank[ra] == rank[rb]: + rank[ra] += 1 + return [find(x) for x in range(count)] + + +def recluster(conn, opts, files, id_column, say): + """`members` and `named_edges` from the predictions at or above the + threshold, as `cpplink cluster --threshold` would write them. + + A prediction naming a record no input holds is skipped as the core skips + it, which is why the id pass comes before the union-find rather than after. + """ + say("reading the predictions at the threshold") + if os.path.isdir(opts.predictions): + edges = read_shard_edges(opts.predictions, opts.threshold) + keys = pc.unique( + pa.concat_arrays([edges["a"].combine_chunks(), edges["b"].combine_chunks()]) + ) + names = id_rows(conn, files, id_column, keys, "rid") + by_rid = dict( + zip(names["rid"].to_pylist(), names["uid"].to_pylist(), strict=True) + ) + a = [by_rid.get(r) for r in edges["a"].to_pylist()] + b = [by_rid.get(r) for r in edges["b"].to_pylist()] + else: + edges = read_file_edges(opts.predictions, opts.threshold) + keys = pc.unique( + pa.concat_arrays([edges["a"].combine_chunks(), edges["b"].combine_chunks()]) + ) + names = id_rows(conn, files, id_column, keys, "uid") + known = set(names["uid"].to_pylist()) + a = [u if u in known else None for u in edges["a"].to_pylist()] + b = [u if u in known else None for u in edges["b"].to_pylist()] + weights = edges["w"].to_pylist() + kept = [ + (x, y, w) + for x, y, w in zip(a, b, weights, strict=True) + if x is not None and y is not None + ] + + say("clustering them") + index = {} + for x, y, _ in kept: + index.setdefault(x, len(index)) + index.setdefault(y, len(index)) + root = union_find(((index[x], index[y]) for x, y, _ in kept), len(index)) + uids = list(index) + size = {} + for r in root: + size[r] = size.get(r, 0) + 1 + rows = [ + (uids[i], uids[root[i]], size[root[i]]) + for i in range(len(uids)) + if size[root[i]] >= opts.min_size + and (not opts.max_size or size[root[i]] <= opts.max_size) + ] + conn.register( + "assignment", + pa.table( + { + "uid": pa.array([r[0] for r in rows], pa.string()), + "cluster_id": pa.array([r[1] for r in rows], pa.string()), + "size": pa.array([r[2] for r in rows], pa.int64()), + } + ), + ) + conn.execute("CREATE TABLE members AS SELECT uid, cluster_id, size FROM assignment") + conn.unregister("assignment") + conn.register( + "kept", + pa.table( + { + "a": pa.array([k[0] for k in kept], pa.string()), + "b": pa.array([k[1] for k in kept], pa.string()), + "weight": pa.array([k[2] for k in kept], pa.float64()), + } + ), + ) + conn.execute("CREATE TABLE named_edges AS SELECT a, b, weight FROM kept") + conn.unregister("kept") + + +def build(conn, opts, id_column, columns, say): + """Fill the cache: members, their records, their edges, per-cluster stats.""" + files = [os.path.abspath(p) for p in opts.data] + picked = ", ".join(quoted(c) for c in columns) + + if opts.threshold is not None and not opts.clusters: + recluster(conn, opts, files, id_column, say) + else: + say("reading the cluster assignment") + ceiling = f" AND cluster_size <= {opts.max_size}" if opts.max_size else "" + conn.execute(f""" + CREATE TABLE members AS + SELECT CAST(unique_id AS VARCHAR) AS uid, + CAST(cluster_id AS VARCHAR) AS cluster_id, + CAST(cluster_size AS BIGINT) AS size + FROM {scan(opts.clusters, ("unique_id", "cluster_id"))} + WHERE cluster_size >= {opts.min_size}{ceiling} + """) + + say("reading the clustered records") + # A record keeps its store row (`rid`): the position in the inputs read in + # order, which is what a shard names. file_row_number is that position + # inside one file, and the offsets between the inputs are the only thing to + # carry. + parts, offset = [], 0 + for path in files: + parts.append(f""" + SELECT m.cluster_id, d.{quoted(id_column)} AS uid, {picked}, + lower(concat_ws(' ', d.{quoted(id_column)}, {picked})) AS text, + {offset} + file_row_number AS rid + FROM read_parquet('{path}', file_row_number = true) d + JOIN members m ON m.uid = d.{quoted(id_column)}""") + offset += pq.ParquetFile(path).metadata.num_rows + conn.execute("CREATE TABLE records AS " + " UNION ALL ".join(parts)) + + if opts.predictions: + if opts.threshold is not None and not opts.clusters: + pass # `named_edges` is what `recluster` clustered + elif os.path.isdir(opts.predictions): + say("naming the predictions") + # A shard names rows, and the clustered records carry theirs. + load_shard_edges(conn, opts.predictions) + conn.execute(""" + CREATE TABLE named_edges AS + SELECT ra.uid AS a, rb.uid AS b, e.w AS weight + FROM raw_edges e + JOIN records ra ON ra.rid = e.a + JOIN records rb ON rb.rid = e.b + """) + conn.execute("DROP TABLE raw_edges") + else: + # A merged file already names records by unique_id. + say("reading the predictions") + conn.execute(f""" + CREATE TABLE named_edges AS + SELECT CAST(id_a AS VARCHAR) AS a, CAST(id_b AS VARCHAR) AS b, + match_weight AS weight + FROM {scan(opts.predictions, ("id_a", "id_b"))} + """) + # Both ends clustered, because clustering at a threshold above the one + # the run wrote at leaves predictions that cross two clusters or land + # outside every one. A pair across two clusters is kept with both + # cluster ids, since it is the prediction clustering overruled and the + # one most worth explaining; a pair with an end no cluster holds has no + # record to show and is dropped. `--threshold` drops the predictions + # below it here, so the cache never holds them; when the clusters were + # computed here they are the components of these very predictions, so + # every pair lands inside one cluster and none is rejected. + floor = ( + f" WHERE e.weight >= {opts.threshold}" if opts.threshold is not None else "" + ) + conn.execute(f""" + CREATE TABLE pairs AS + SELECT e.a, e.b, e.weight, ra.cluster_id AS cluster_a, + rb.cluster_id AS cluster_b + FROM named_edges e + JOIN records ra ON ra.uid = e.a + JOIN records rb ON rb.uid = e.b{floor} + """) + conn.execute("DROP TABLE named_edges") + conn.execute(""" + CREATE TABLE edges AS + SELECT cluster_a AS cluster_id, a, b, weight FROM pairs + WHERE cluster_a = cluster_b + """) + conn.execute("CREATE INDEX edges_by_cluster ON edges (cluster_id)") + # A rejected pair is asked for from either of its clusters. + conn.execute("CREATE INDEX pairs_by_cluster_a ON pairs (cluster_a)") + conn.execute("CREATE INDEX pairs_by_cluster_b ON pairs (cluster_b)") + conn.execute("CREATE INDEX pairs_by_a ON pairs (a)") + conn.execute("CREATE INDEX pairs_by_b ON pairs (b)") + if opts.waterfalls: + # Only the rows the page can ever ask for: the pairs kept above. + say("reading the waterfalls") + conn.execute(f""" + CREATE TABLE waterfalls AS + SELECT w.* REPLACE (CAST(w.id_a AS VARCHAR) AS id_a, + CAST(w.id_b AS VARCHAR) AS id_b) + FROM {scan(opts.waterfalls, ("id_a", "id_b"))} w + JOIN pairs p ON p.a = CAST(w.id_a AS VARCHAR) + AND p.b = CAST(w.id_b AS VARCHAR) + """) + conn.execute("CREATE INDEX waterfalls_by_pair ON waterfalls (id_a, id_b)") + else: + conn.execute( + "CREATE TABLE edges (cluster_id VARCHAR, a VARCHAR, b VARCHAR, weight DOUBLE)" + ) + conn.execute( + "CREATE TABLE pairs (a VARCHAR, b VARCHAR, weight DOUBLE, " + "cluster_a VARCHAR, cluster_b VARCHAR)" + ) + + if opts.truth: + say("closing the known pairs") + uids = [row[0] for row in conn.execute("SELECT uid FROM members").fetchall()] + groups = read_truth(opts.truth, set(uids)) + table = pa.table( + { + "uid": pa.array(list(groups.keys()), pa.string()), + "grp": pa.array(list(groups.values()), pa.int64()), + } + ) + conn.register("groups", table) + conn.execute("CREATE TABLE truth AS SELECT uid, grp FROM groups") + conn.unregister("groups") + else: + conn.execute("CREATE TABLE truth (uid VARCHAR, grp BIGINT)") + + say("counting what each cluster agrees on") + # count(DISTINCT x) skips nulls, which is the page's rule: a column with one + # value and some gaps still agrees. + splits = " + ".join( + f"CASE WHEN count(DISTINCT {quoted(c)}) > 1 THEN 1 ELSE 0 END" for c in columns + ) + conn.execute(f""" + CREATE TABLE stats AS + WITH per_cluster AS ( + SELECT cluster_id, count(*) AS size, {splits} AS discord + FROM records GROUP BY cluster_id), + per_edge AS ( + SELECT cluster_id, count(*) AS edge_count, min(weight) AS weakest + FROM edges GROUP BY cluster_id), + per_truth AS ( + SELECT r.cluster_id, + count(DISTINCT coalesce(CAST(t.grp AS VARCHAR), 'x' || r.uid)) + AS entities + FROM records r LEFT JOIN truth t ON t.uid = r.uid + GROUP BY r.cluster_id) + SELECT c.cluster_id AS id, c.size, c.discord, + coalesce(e.edge_count, 0) AS edge_count, e.weakest, + coalesce(t.entities, 0) AS entities + FROM per_cluster c + LEFT JOIN per_edge e ON e.cluster_id = c.cluster_id + LEFT JOIN per_truth t ON t.cluster_id = c.cluster_id + """) + conn.execute("CREATE INDEX stats_by_id ON stats (id)") + conn.execute("CREATE INDEX records_by_cluster ON records (cluster_id)") + conn.execute("CREATE INDEX records_by_uid ON records (uid)") + + +def announce(line): + """The default progress logger: one line, flushed, so a server started in + the background still shows what it is doing.""" + print(line, flush=True) + + +def open_cache(opts, log=announce): + """A read-only connection to the cache for `opts`, built first if it is + missing, stale, or `opts.rebuild` says so. Returns `(conn, id_column, columns)`.""" + id_column, columns = schema_columns(opts.schema) + stamp = fingerprint(opts, columns) + + fresh = opts.rebuild or not os.path.exists(opts.cache) + if not fresh: + try: + conn = connect(opts.cache, read_only=True) + fresh = conn.execute("SELECT stamp FROM meta").fetchone()[0] != stamp + conn.close() + except Exception: + fresh = True + + if fresh: + if os.path.exists(opts.cache): + os.remove(opts.cache) + start = time.time() + conn = connect(opts.cache) + conn.execute(f"SET memory_limit = '{opts.memory}'") + conn.execute("SET preserve_insertion_order = false") + if opts.threads: + conn.execute(f"SET threads = {opts.threads}") + + def say(what): + log(f" {time.time() - start:6.1f}s {what}") + + log(f"building {opts.cache}") + build(conn, opts, id_column, columns, say) + conn.execute("CREATE TABLE meta (stamp VARCHAR)") + conn.execute("INSERT INTO meta VALUES (?)", [stamp]) + conn.close() + size = os.path.getsize(opts.cache) / 1e6 + log(f" {time.time() - start:6.1f}s done, {size:.0f} MB") + + conn = connect(opts.cache, read_only=True) + conn.execute(f"SET memory_limit = '{opts.memory}'") + return conn, id_column, columns diff --git a/python/cpplink_viewer/inputs.py b/python/cpplink_viewer/inputs.py new file mode 100644 index 0000000..77755f2 --- /dev/null +++ b/python/cpplink_viewer/inputs.py @@ -0,0 +1,178 @@ +# Copyright 2026 Mathieu Fourment +# SPDX-License-Identifier: MIT +"""The run's files as the viewer reads them: schema, shards, truth, model, waterfalls. + +Nothing here imports the compiled module. The viewer reads what a run wrote and +so runs against a checkout with no build, or against a wheel with no Arrow. +""" + +from __future__ import annotations + +import json + +import numpy as np +import pyarrow as pa +import pyarrow.compute as pc +import pyarrow.csv as pacsv + +# The binary shard `predict` writes: `kEdgeMagic` and one `kEdgeBytes` record per +# prediction, (row a, row b, packed gamma, match weight). A test reads a shard the +# core wrote through this, so the two cannot drift silently. +EDGE_MAGIC = b"CPPLNKE1" +EDGE_DTYPE = np.dtype([("a", " --out ` writes, one +wide row per prediction; it is loaded into the cache beside the predictions and +a click is one lookup, so the ledger is the scorer's own arithmetic and nothing +is computed or spawned here. The level labels and rates come from `--model`. +Predictions clustering kept apart, above the write threshold and below the +clustering one, are kept and listed under both of their clusters as `rejected`. + +`--threshold` keeps only the predictions at or above it. With `--clusters` the +file is taken as what `cpplink cluster --threshold` wrote at that threshold and +read as it stands, which is the fast path; without one the predictions are +clustered here, by the same union-find over the same predictions in the same +order, so the cache holds exactly the partition that command writes, named by +the same representatives. +""" + +from __future__ import annotations + +import argparse +import json +import os +import sys +import threading +import webbrowser +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from importlib import resources +from urllib.parse import parse_qs, urlparse + +from .cache import Options, announce, open_cache, quoted +from .inputs import cell, ledger, read_model_levels + +# The list's order: the chosen key in either direction, a cluster with no +# value for it (no predictions, so no weakest edge) last either way, and the +# ties broken the same way whichever direction the key runs. +ORDERS = { + "discord": ("discord", "size DESC, id"), + "size": ("size", "discord DESC, id"), + "weakest": ("weakest", "size DESC, id"), + "id": ("id", ""), +} + +FILTERS = { + "all": "TRUE", + "split": "discord > 0", + "transitive": "size * (size - 1) / 2 > edge_count", + "mixed": "entities > 1", +} + +REJECTED_LIMIT = 1000 # rejected predictions listed per cluster, weakest first +LIST_LIMIT = 500 # clusters a list request may ask for at once + + +def order_clause(sort, desc): + key, ties = ORDERS.get(sort, ORDERS["discord"]) + clause = f"{key} {'DESC' if desc else 'ASC'} NULLS LAST" + return clause + (", " + ties if ties else "") + + +def like_pattern(text): + escaped = text.lower().replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + return "%" + escaped + "%" + + +class Viewer: + """The queries the page makes, over one shared read-only connection.""" + + def __init__(self, conn, columns, source, max_rows, has_truth, levels=None): + self.conn = conn + self.columns = columns + self.source = source + self.max_rows = max_rows + self.has_truth = has_truth + self.levels = levels # the model's level table, when waterfalls are held + self.lock = threading.Lock() + with self.lock: + self.totals = conn.execute( + "SELECT count(*), coalesce(sum(size), 0) FROM stats" + ).fetchone() + self.pair_total = conn.execute("SELECT count(*) FROM pairs").fetchone()[0] + self.waterfall_columns = [ + r[0] + for r in conn.execute( + "SELECT column_name FROM information_schema.columns " + "WHERE table_name = 'waterfalls' ORDER BY ordinal_position" + ).fetchall() + ] + + @classmethod + def open(cls, opts, log=announce): + """A viewer over the cache `opts` names, built first if it has to be.""" + reason = opts.check() + if reason: + raise SystemExit(reason) + conn, _, columns = open_cache(opts, log) + levels = read_model_levels(opts.model) if opts.model else None + source = ", ".join(os.path.basename(p) for p in opts.data) + return cls(conn, columns, source, opts.max_rows, bool(opts.truth), levels) + + def query(self, sql, params=()): + with self.lock: + return self.conn.execute(sql, params).fetchall() + + def summary(self): + return { + "columns": self.columns, + "source": self.source, + "totals": { + "clusters": self.totals[0], + "records": self.totals[1], + "pairs": self.pair_total, + }, + "has_truth": self.has_truth, + "has_model": bool(self.waterfall_columns), + } + + def waterfall(self, a, b): + """The ledger of one prediction, from the wide row the run's explain wrote.""" + if not self.waterfall_columns or self.levels is None: + return None + rows = self.query("SELECT * FROM waterfalls WHERE id_a = ? AND id_b = ?", [a, b]) + if not rows: + return None + row = dict(zip(self.waterfall_columns, rows[0], strict=True)) + return ledger(row, self.levels) + + def rejected(self, cluster_id): + """The predictions clustering overruled that touch one cluster. + + Each names the other cluster its second record went to, and the truth's + verdict where there is one, so the page can list them beside the + cluster's own edges without a lookup per row. + """ + rows = self.query( + "SELECT p.a, p.b, p.weight, p.cluster_a, p.cluster_b, ta.grp, tb.grp " + "FROM pairs p LEFT JOIN truth ta ON ta.uid = p.a " + "LEFT JOIN truth tb ON tb.uid = p.b " + "WHERE (p.cluster_a = ? OR p.cluster_b = ?) AND p.cluster_a <> p.cluster_b " + f"ORDER BY p.weight, p.a, p.b LIMIT {REJECTED_LIMIT + 1}", + [cluster_id, cluster_id], + ) + more = max(0, len(rows) - REJECTED_LIMIT) + found = [] + for a, b, w, ca, cb, ga, gb in rows[:REJECTED_LIMIT]: + same = (ga is not None and ga == gb) if self.has_truth else None + found.append( + { + "a": a, + "b": b, + "w": round(w, 3), + "other": cb if ca == cluster_id else ca, + "same": same, + } + ) + return found, more + + def pair(self, a, b): + head = self.query( + "SELECT a, b, weight, cluster_a, cluster_b FROM pairs WHERE a = ? AND b = ?", + [a, b], + ) + if not head: + return None + picked = ", ".join(quoted(c) for c in self.columns) + rows = { + r[0]: r + for r in self.query( + f"SELECT uid, {picked} FROM records WHERE uid IN (?, ?)", [a, b] + ) + } + if a not in rows or b not in rows: + return None + answer = { + "a": a, + "b": b, + "w": round(head[0][2], 3), + "ca": head[0][3], + "cb": head[0][4], + "rows": [ + {"id": uid, "v": [cell(v) for v in rows[uid][1:]]} for uid in (a, b) + ], + } + if self.has_truth: + groups = dict( + self.query("SELECT uid, grp FROM truth WHERE uid IN (?, ?)", [a, b]) + ) + same = a in groups and b in groups and groups[a] == groups[b] + answer["rows"][0]["t"] = 0 + answer["rows"][1]["t"] = 0 if same else 1 + answer["waterfall"] = self.waterfall(a, b) + return answer + + def listing(self, q, field, column, sort, desc, only, offset, limit): + """One page of the cluster list. + + `field` says what `q` is read against: `any` is a substring of the + concatenated record, `cluster` of the cluster id, `col` of the named + column, and `id` a comma-separated list of record ids matched whole, in + which case the answer also says which of them named a record. + """ + where = FILTERS.get(only, "TRUE") + params = [] + found = None + if q and field == "id": + ids = list(dict.fromkeys(s.strip() for s in q.split(",") if s.strip())) + marks = ", ".join("?" * len(ids)) + where += ( + f" AND id IN (SELECT cluster_id FROM records " + f"WHERE CAST(uid AS VARCHAR) IN ({marks}))" + ) + params.extend(ids) + found = [ + r[0] + for r in self.query( + f"SELECT DISTINCT CAST(uid AS VARCHAR) FROM records " + f"WHERE CAST(uid AS VARCHAR) IN ({marks})", + ids, + ) + ] + elif q and field == "cluster": + where += " AND lower(id) LIKE ? ESCAPE '\\'" + params.append(like_pattern(q)) + elif q and field == "col": + if column not in self.columns: + raise ValueError(f"no column {column!r}") + where += ( + f" AND id IN (SELECT cluster_id FROM records " + f"WHERE lower(CAST({quoted(column)} AS VARCHAR)) LIKE ? ESCAPE '\\')" + ) + params.append(like_pattern(q)) + elif q: + where += ( + " AND id IN (SELECT cluster_id FROM records " + "WHERE text LIKE ? ESCAPE '\\')" + ) + params.append(like_pattern(q)) + matched = self.query(f"SELECT count(*) FROM stats WHERE {where}", params)[0][0] + rows = self.query( + f"SELECT id, size, discord, edge_count, weakest, entities FROM stats " + f"WHERE {where} ORDER BY {order_clause(sort, desc)} " + f"LIMIT {int(limit)} OFFSET {int(offset)}", + params, + ) + answer = { + "matched": matched, + "clusters": [ + { + "id": r[0], + "size": r[1], + "discord": r[2], + "edge_count": r[3], + "weakest": r[4], + "entities": r[5], + } + for r in rows + ], + } + if found is not None: + answer["found"] = found + return answer + + def cluster(self, cluster_id): + head = self.query( + "SELECT id, size, discord, edge_count, weakest, entities " + "FROM stats WHERE id = ?", + [cluster_id], + ) + if not head: + return None + picked = ", ".join(quoted(c) for c in self.columns) + rows = self.query( + f"SELECT uid, {picked} FROM records WHERE cluster_id = ? " + f"ORDER BY uid LIMIT {int(self.max_rows)}", + [cluster_id], + ) + ids = [r[0] for r in rows] + groups, labels = {}, {} + if self.has_truth: + for uid, grp in self.query( + "SELECT uid, grp FROM truth WHERE uid IN " + "(SELECT uid FROM records WHERE cluster_id = ?)", + [cluster_id], + ): + groups[uid] = grp + for uid in ids: + key = groups.get(uid, "x" + uid) + labels.setdefault(key, len(labels)) + # The whole cluster's distinct counts, not the page of rows shown. + counts = self.query( + "SELECT " + + ", ".join(f"count(DISTINCT {quoted(c)})" for c in self.columns) + + " FROM records WHERE cluster_id = ?", + [cluster_id], + )[0] + edges = self.query( + "SELECT a, b, weight FROM edges WHERE cluster_id = ? AND a IN " + "(SELECT uid FROM records WHERE cluster_id = ?)", + [cluster_id, cluster_id], + ) + keep = set(ids) + answer = { + "id": head[0][0], + "size": head[0][1], + "discord": head[0][2], + "edge_count": head[0][3], + "weakest": head[0][4], + "entities": head[0][5], + "dist": list(counts), + "rows": [{"id": r[0], "v": [cell(v) for v in r[1:]]} for r in rows], + "edges": [ + [a, b, round(w, 3)] for a, b, w in edges if a in keep and b in keep + ], + } + if self.has_truth: + for row in answer["rows"]: + row["t"] = labels[groups.get(row["id"], "x" + row["id"])] + answer["rejected"], more = self.rejected(cluster_id) + if more: + answer["rejected_more"] = more + return answer + + +def page_html(): + """The viewer page, shipped beside this module.""" + return ( + resources.files(__package__).joinpath("viewer.html").read_text(encoding="utf-8") + ) + + +def handler_for(viewer, page): + class Handler(BaseHTTPRequestHandler): + def log_message(self, *_): + pass + + def send_json(self, payload, status=200): + body = json.dumps(payload).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def do_GET(self): + url = urlparse(self.path) + query = {k: v[0] for k, v in parse_qs(url.query).items()} + try: + if url.path in ("/", "/index.html"): + body = page.encode() + self.send_response(200) + self.send_header("Content-Type", "text/html; charset=utf-8") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + elif url.path == "/api/summary": + self.send_json(viewer.summary()) + elif url.path == "/api/list": + self.send_json( + viewer.listing( + query.get("q", ""), + query.get("field", "any"), + query.get("column", ""), + query.get("sort", "discord"), + query.get("dir", "desc") == "desc", + query.get("only", "all"), + int(query.get("offset", 0)), + min(int(query.get("limit", 100)), LIST_LIMIT), + ) + ) + elif url.path == "/api/cluster": + found = viewer.cluster(query.get("id", "")) + self.send_json( + found or {"error": "no such cluster"}, 200 if found else 404 + ) + elif url.path == "/api/pair": + found = viewer.pair(query.get("a", ""), query.get("b", "")) + self.send_json( + found or {"error": "no such prediction"}, 200 if found else 404 + ) + else: + self.send_json({"error": "not found"}, 404) + except BrokenPipeError: + pass + except Exception as trouble: # a bad query should not kill the server + self.send_json({"error": str(trouble)}, 500) + + return Handler + + +def make_server(viewer, port=0, host="127.0.0.1"): + """An HTTP server over `viewer`, bound and not yet serving; port 0 picks a + free one, which is what a test wants.""" + return ThreadingHTTPServer((host, port), handler_for(viewer, page_html())) + + +def parse_args(argv=None): + ap = argparse.ArgumentParser( + prog="cpplink-viewer", + description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter, + ) + ap.add_argument("data", nargs="+", help="the parquet file(s) the run read, in order") + ap.add_argument("--schema", required=True) + ap.add_argument("--clusters", help="the csv or parquet cpplink cluster wrote") + ap.add_argument( + "--predictions", + help="the run's predictions -- one csv or parquet file, or the shard " + "directory -- for within-cluster weights", + ) + ap.add_argument("--truth", help="known pairs, to colour members by true entity") + ap.add_argument( + "--waterfalls", + help="the csv or parquet cpplink explain --predictions wrote; draws each " + "prediction's waterfall", + ) + ap.add_argument( + "--model", + help="the model the run scored with, for the level labels and rates the " + "waterfall shows", + ) + ap.add_argument("--cache", default=Options.cache) + ap.add_argument("--rebuild", action="store_true", help="rebuild the cache first") + ap.add_argument("--min-size", type=int, default=Options.min_size) + ap.add_argument( + "--max-size", type=int, default=Options.max_size, help="0 = no ceiling" + ) + ap.add_argument( + "--threshold", + type=float, + help="keep only the predictions at or above this match_weight; with " + "--clusters that file is read as the clustering at this threshold, " + "without one the predictions are clustered here exactly as " + "cpplink cluster --threshold would", + ) + ap.add_argument( + "--max-rows", + type=int, + default=Options.max_rows, + help="members shown per cluster; the rest are counted only", + ) + ap.add_argument( + "--memory", default=Options.memory, help="what DuckDB may use while building" + ) + ap.add_argument( + "--threads", type=int, default=Options.threads, help="0 = DuckDB's own default" + ) + ap.add_argument("--port", type=int, default=8770) + ap.add_argument("--open", action="store_true", help="open a browser on it") + args = ap.parse_args(argv) + opts = Options( + data=args.data, + schema=args.schema, + clusters=args.clusters, + predictions=args.predictions, + truth=args.truth, + waterfalls=args.waterfalls, + model=args.model, + cache=args.cache, + rebuild=args.rebuild, + min_size=args.min_size, + max_size=args.max_size, + threshold=args.threshold, + max_rows=args.max_rows, + memory=args.memory, + threads=args.threads, + ) + reason = opts.check() + if reason: + ap.error(reason) + return opts, args.port, args.open + + +def main(argv=None): + opts, port, open_browser = parse_args(argv) + viewer = Viewer.open(opts) + server = make_server(viewer, port) + where = f"http://127.0.0.1:{server.server_address[1]}/" + clusters, records = viewer.totals + announce(f"{clusters:,} clusters over {records:,} records at {where}") + if open_browser: + webbrowser.open(where) + try: + server.serve_forever() + except KeyboardInterrupt: + print("\nstopped") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/python/cpplink_viewer/viewer.html b/python/cpplink_viewer/viewer.html new file mode 100644 index 0000000..4384a90 --- /dev/null +++ b/python/cpplink_viewer/viewer.html @@ -0,0 +1,1274 @@ + + + + + +cpplink clusters + + + +
+

cpplink clusters

+ + + j k cluster · n p pair · / search + + +
+
+
+
+
+ + +
+ +
+ + + +
+
+ + +
+
+
+
+
Pick a cluster.
+
+ + + diff --git a/python/tests/test_viewer.py b/python/tests/test_viewer.py new file mode 100644 index 0000000..f80c597 --- /dev/null +++ b/python/tests/test_viewer.py @@ -0,0 +1,318 @@ +# Copyright 2026 Mathieu Fourment +# SPDX-License-Identifier: MIT +"""`cpplink_viewer` serves what a run wrote, and where it reimplements the core +(the union-find, the shard format) it agrees with it.""" + +from __future__ import annotations + +import csv +import json +import threading +from collections import defaultdict +from pathlib import Path +from urllib.request import urlopen + +import numpy as np +import pyarrow.parquet as pq +import pytest +from conftest import ROWS + +import cpplink + +duckdb = pytest.importorskip("duckdb") + +from cpplink_viewer import Options, Viewer, make_server # noqa: E402 +from cpplink_viewer.cache import read_shard_edges # noqa: E402 +from cpplink_viewer.inputs import EDGE_DTYPE, EDGE_MAGIC # noqa: E402 + +THRESHOLD = 10 # what conftest's predict wrote at +HIGHER = 60 # a threshold that prunes a good share of the fixture's predictions + + +class Run: + """A clustered run with its waterfalls, and the viewer's options over it.""" + + def __init__(self, sample, root: Path) -> None: + self.sample = sample + self.clusters = root / "clusters.csv" + self.cluster_result = sample.linker.cluster(sample.predictions, out=self.clusters) + self.waterfalls = root / "waterfalls.parquet" + result = cpplink.run( + [ + "explain", + "--schema", + str(sample.schema_path), + "--model", + str(sample.model_path), + "--predictions", + str(sample.predictions), + "--out", + str(self.waterfalls), + str(sample.parquet), + ] + ) + assert result.code == 0, result.stderr + self.shards = root / "shards" + sample.linker.predict(sample.model, self.shards, threshold=THRESHOLD) + self.root = root + + def options(self, **overrides) -> Options: + opts = Options( + data=[str(self.sample.parquet)], + schema=str(self.sample.schema_path), + clusters=str(self.clusters), + predictions=str(self.sample.predictions), + truth=str(self.sample.truth), + waterfalls=str(self.waterfalls), + model=str(self.sample.model_path), + cache=str(self.root / "cache.duckdb"), + ) + for key, value in overrides.items(): + setattr(opts, key, value) + return opts + + def written_clusters(self) -> dict[str, str]: + with open(self.clusters, newline="") as handle: + return {r["unique_id"]: r["cluster_id"] for r in csv.DictReader(handle)} + + +@pytest.fixture(scope="module") +def run(sample, tmp_path_factory: pytest.TempPathFactory) -> Run: + return Run(sample, tmp_path_factory.mktemp("viewer")) + + +@pytest.fixture(scope="module") +def viewer(run: Run) -> Viewer: + return Viewer.open(run.options(), log=lambda _: None) + + +def test_summary_counts_the_cluster_file(run: Run, viewer: Viewer) -> None: + summary = viewer.summary() + written = run.written_clusters() + assert summary["totals"]["clusters"] == len(set(written.values())) + assert summary["totals"]["records"] == len(written) + assert summary["totals"]["pairs"] == pq.read_metadata(run.sample.predictions).num_rows + with open(run.sample.schema_path) as handle: + declared = json.load(handle)["columns"] + assert summary["columns"] == [c["name"] for c in declared if "derive" not in c] + assert summary["has_truth"] and summary["has_model"] + + +def test_listing_pages_through_every_cluster(viewer: Viewer) -> None: + total = viewer.summary()["totals"]["clusters"] + seen = [] + offset = 0 + while True: + page = viewer.listing("", "any", "", "size", True, "all", offset, 7) + assert page["matched"] == total + if not page["clusters"]: + break + seen.extend(c["id"] for c in page["clusters"]) + offset += len(page["clusters"]) + assert len(seen) == total and len(set(seen)) == total + # The filters are subsets, and an id search names the record it found. + split = viewer.listing("", "any", "", "discord", True, "split", 0, 500) + assert 0 < split["matched"] <= total + assert all(c["discord"] > 0 for c in split["clusters"]) + first = viewer.cluster(seen[0]) + uid = first["rows"][0]["id"] + by_id = viewer.listing(uid + ", no-such-id", "id", "", "discord", True, "all", 0, 10) + assert by_id["matched"] == 1 and by_id["found"] == [uid] + assert by_id["clusters"][0]["id"] == seen[0] + with pytest.raises(ValueError, match="no column"): + viewer.listing("x", "col", "no_such_column", "discord", True, "all", 0, 10) + + +def test_cluster_holds_its_members_and_edges(run: Run, viewer: Viewer) -> None: + written = run.written_clusters() + members = defaultdict(set) + for uid, cid in written.items(): + members[cid].add(uid) + largest = max(members, key=lambda c: len(members[c])) + cluster = viewer.cluster(largest) + assert cluster["size"] == len(members[largest]) + assert {r["id"] for r in cluster["rows"]} == members[largest] + assert len(cluster["dist"]) == len(viewer.columns) + assert all(len(r["v"]) == len(viewer.columns) for r in cluster["rows"]) + assert cluster["edge_count"] == len(cluster["edges"]) >= len(members[largest]) - 1 + assert all( + a in members[largest] and b in members[largest] for a, b, _ in cluster["edges"] + ) + assert cluster["weakest"] == pytest.approx( + min(w for _, _, w in cluster["edges"]), abs=1e-3 + ) + # With truth, every member carries an entity label. + assert all("t" in r for r in cluster["rows"]) + assert cluster["entities"] >= 1 + assert viewer.cluster("no-such-cluster") is None + + +def test_pair_is_the_scorer_ledger_read_back(run: Run, viewer: Viewer) -> None: + table = pq.read_table( + run.sample.predictions, columns=["id_a", "id_b", "match_weight"] + ) + row = table.slice(0, 1).to_pylist()[0] + pair = viewer.pair(row["id_a"], row["id_b"]) + assert pair is not None + assert pair["w"] == pytest.approx(row["match_weight"], abs=1e-3) + ledger = pair["waterfall"] + assert ledger["weight"] == pytest.approx(row["match_weight"], abs=1e-9) + # The ledger is the explain report's: same steps, and the running total + # ends at the weight. + explanation = run.sample.linker.explain( + row["id_a"], row["id_b"], model=run.sample.model + ) + assert [s["name"] for s in ledger["steps"]] == run.sample.schema.comparison_names + assert ledger["steps"][-1]["running"] + sum( + i["bits"] for i in ledger["interactions"] + ) == pytest.approx(ledger["weight"], abs=1e-9) + assert ledger["prior"] == pytest.approx(explanation.waterfall.prior, abs=1e-9) + assert {r["id"] for r in pair["rows"]} == {row["id_a"], row["id_b"]} + assert viewer.pair("no", "such") is None + + +def test_threshold_reclusters_as_the_core_does(run: Run) -> None: + """The Python union-find, over the file and over the shards, produces the + partition `cluster --threshold` writes, named by the same representatives.""" + for predictions in (run.sample.predictions, run.shards): + # Each source is clustered by the core itself: the shards are a second + # run's, so their edge order, and with it the representatives, is its own. + expected_path = run.root / f"clusters_{HIGHER}_{predictions.name}.csv" + run.sample.linker.cluster(predictions, threshold=HIGHER, out=expected_path) + with open(expected_path, newline="") as handle: + expected = {r["unique_id"]: r["cluster_id"] for r in csv.DictReader(handle)} + assert expected # the fixture must leave something above the higher threshold + opts = run.options( + clusters=None, + predictions=str(predictions), + threshold=HIGHER, + cache=str(run.root / f"cache_{predictions.name}.duckdb"), + ) + viewer = Viewer.open(opts, log=lambda _: None) + got = dict(viewer.query("SELECT uid, cluster_id FROM records")) + assert got == expected + # Clustered at the threshold, no prediction crosses two clusters. + assert viewer.query( + "SELECT count(*) FROM pairs WHERE cluster_a <> cluster_b" + ) == [(0,)] + viewer.conn.close() + + +def test_a_prediction_across_two_clusters_is_rejected_under_both(run: Run) -> None: + """A prediction clustering overruled is listed under both of its clusters, + naming the other. The fixture has none, so one is added: the run's + predictions as csv plus a pair between two clusters.""" + written = run.written_clusters() + first, second = sorted(set(written.values()))[:2] + a = min(u for u, c in written.items() if c == first) + b = min(u for u, c in written.items() if c == second) + table = pq.read_table( + run.sample.predictions, columns=["id_a", "id_b", "match_weight"] + ) + predictions = run.root / "with_crossing.csv" + with open(predictions, "w", newline="") as handle: + out = csv.writer(handle) + out.writerow(["id_a", "id_b", "match_weight"]) + out.writerows(list(r.values()) for r in table.to_pylist()) + out.writerow([a, b, 20.0]) + opts = run.options( + predictions=str(predictions), + waterfalls=None, + model=None, + cache=str(run.root / "cache_rejected.duckdb"), + ) + viewer = Viewer.open(opts, log=lambda _: None) + assert viewer.summary()["totals"]["pairs"] == table.num_rows + 1 + assert viewer.query("SELECT a, b FROM pairs WHERE cluster_a <> cluster_b") == [(a, b)] + for mine, other in ((first, second), (second, first)): + rejected = viewer.cluster(mine)["rejected"] + assert [(r["a"], r["b"], r["w"], r["other"]) for r in rejected] == [ + (a, b, 20.0, other) + ] + assert rejected[0]["same"] is False # two entities, by the truth file + pair = viewer.pair(a, b) + assert (pair["ca"], pair["cb"]) == (first, second) and pair["waterfall"] is None + # Filtering at a threshold above it drops it from the cache altogether. + opts.threshold = 30.0 + opts.cache = str(run.root / "cache_rejected_30.duckdb") + viewer = Viewer.open(opts, log=lambda _: None) + assert viewer.query("SELECT count(*) FROM pairs WHERE cluster_a <> cluster_b") == [ + (0,) + ] + assert viewer.summary()["totals"]["pairs"] == table.num_rows + + +def test_shards_are_read_as_the_core_wrote_them(run: Run) -> None: + shards = sorted(run.shards.glob("shard-*.bin")) + assert shards + with open(shards[0], "rb") as handle: + assert handle.read(len(EDGE_MAGIC)) == EDGE_MAGIC + table = read_shard_edges(str(run.shards), THRESHOLD) + merged = pq.read_table(run.sample.predictions, columns=["match_weight"]) + assert table.num_rows == merged.num_rows + assert sorted(table["w"].to_pylist()) == pytest.approx( + sorted(merged["match_weight"].to_pylist()), abs=1e-9 + ) + assert EDGE_DTYPE.itemsize == 20 + assert np.asarray(table["a"]).max() < ROWS + + +def test_cache_is_reused_until_an_input_changes(run: Run) -> None: + opts = run.options(cache=str(run.root / "cache_reuse.duckdb")) + said = [] + Viewer.open(opts, log=said.append).conn.close() + assert any("building" in line for line in said) + said.clear() + Viewer.open(opts, log=said.append).conn.close() + assert said == [] + opts.max_rows += 1 # part of the fingerprint + Viewer.open(opts, log=said.append).conn.close() + assert any("building" in line for line in said) + + +def test_refused_options_say_why(run: Run) -> None: + assert Options(data=["x"], schema="s").check() == ( + "--clusters is needed without --threshold" + ) + assert ( + "go together" in Options(data=["x"], schema="s", clusters="c", model="m").check() + ) + assert ( + "needs --predictions" + in Options(data=["x"], schema="s", clusters="c", threshold=1.0).check() + ) + assert run.options().check() is None + with pytest.raises(SystemExit, match="--clusters is needed"): + Viewer.open(Options(data=["x"], schema="s")) + + +def test_http_serves_the_page_and_the_api(viewer: Viewer) -> None: + server = make_server(viewer, port=0) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + where = f"http://127.0.0.1:{server.server_address[1]}" + with urlopen(where + "/") as answer: + page = answer.read().decode() + assert "cpplink clusters" in page + assert "__CLUSTER_DATA__" not in page + with urlopen(where + "/api/summary") as answer: + assert json.load(answer) == viewer.summary() + with urlopen(where + "/api/list?sort=size&limit=3") as answer: + listing = json.load(answer) + assert len(listing["clusters"]) == 3 + first = listing["clusters"][0]["id"] + with urlopen(where + f"/api/cluster?id={first}") as answer: + cluster = json.load(answer) + assert cluster["id"] == first + edge = cluster["edges"][0] + with urlopen(where + f"/api/pair?a={edge[0]}&b={edge[1]}") as answer: + pair = json.load(answer) + assert pair["w"] == edge[2] + with pytest.raises(Exception, match="404"): + urlopen(where + "/api/cluster?id=no-such-cluster") + with pytest.raises(Exception, match="404"): + urlopen(where + "/nothing") + finally: + server.shutdown() + server.server_close()