From 0f3a2d92d864d887035da4cef4ed6355939c70be Mon Sep 17 00:00:00 2001 From: Emil Date: Fri, 24 Jul 2026 14:38:11 +0300 Subject: [PATCH] Add distributed similarity search --- STATUS.md | 22 +- docs/ctx-07-distributed-workload-protocol.md | 21 +- docs/worker-daemon-task.md | 12 +- scimesh/distributed/__init__.py | 10 +- scimesh/distributed/registry.py | 14 +- scimesh/distributed/similarity_search.py | 380 +++++++++++++++++++ scimesh/worker/daemon.py | 14 +- scimesh/worker/runners.py | 8 + tests/test_distributed_similarity_search.py | 184 +++++++++ tests/test_worker_daemon.py | 131 +++++-- 10 files changed, 749 insertions(+), 47 deletions(-) create mode 100644 scimesh/distributed/similarity_search.py create mode 100644 tests/test_distributed_similarity_search.py diff --git a/STATUS.md b/STATUS.md index 1b4d7cb..cebfd27 100644 --- a/STATUS.md +++ b/STATUS.md @@ -32,8 +32,8 @@ Docker PostgreSQL stack on 2026-07-23. | CTX-04 Worker registry and HTTP API | Implemented | Registration, claim, heartbeat, result, failure, and status endpoints. | | CTX-05 Artifact storage | Implemented | Coordinator-owned inputs/results, checksum verification, and upload flow. | | CTX-06 Python Worker live-contract alignment | Implemented | Worker completed a real uploaded shard via HTTP on 2026-07-23. | -| CTX-07 Distributed workload protocol | Implemented | Versioned Python contract models, registry, strict plan validation, and deterministic reduction ordering are in `scimesh/distributed/`. The concrete molecular planner/reducer remains CTX-08/09. | -| CTX-08 Distributed similarity-search | Not started | Local reference exists. | +| CTX-07 Distributed workload protocol | Implemented | Versioned Python contract models, registry, strict plan validation, and deterministic reduction ordering are in `scimesh/distributed/`. | +| CTX-08 Distributed similarity-search | Implemented (scientific layer) | Python planner resolves `query_id` once, creates deterministic shard plans, worker adapter emits exact partial top-k CSVs/metrics, and reducer matches the local reference. Coordinator persistence/orchestration remains CTX-09. | | CTX-09 Reducer and final-result API | Not started | Depends on CTX-07 and CTX-08. | | CTX-10 Distributed similarity-graph | Not started | Local reference exists. | | CTX-11 Dashboard/operator view | Implemented (diagnostic scope) | Protected local view: job/task/worker status, validated similarity-search upload, diagnostic partial-artifact download, and bounded polling. Final-result reduction remains CTX-09. | @@ -41,20 +41,22 @@ Docker PostgreSQL stack on 2026-07-23. ## Next recommended assignment -Assign **CTX-08** to the workload role: implement the molecular -`similarity-search` planner and worker adapter on top of the accepted CTX-07 -contract. +Assign **CTX-09** to the coordinator role: materialize planned shards, +persist them transactionally, invoke the registered reducer once, and expose a +durable final artifact. ## Known constraints -- The CTX-07 protocol is implemented, but no concrete molecular planner or - reducer is registered yet; the operator UI labels `partial_result` files as - diagnostic and cannot present them as final output. +- The Python `similarity-search` planner/reducer is implemented, but the Go + coordinator does not yet invoke it or persist its final artifact. The + operator UI labels `partial_result` files as diagnostic and cannot present + them as final output. Use the local `scimesh` CLI for complete workload results. - The worker/coordinator flow currently accepts both underscore API workload names and hyphenated CLI names while the contract is consolidated. -- A real-stack worker test uses a small `query_smiles` shard. Resolving a - `query_id` once and sharing it across shards belongs to CTX-07. +- A real-stack worker test uses a small `query_smiles` shard. The Python + planner resolves `query_id` once and shares `query_smiles`; connecting that + planner to uploaded coordinator jobs belongs to CTX-09. - The coordinator accepts uploaded distributed jobs only for `similarity-search` with `query_smiles`. It rejects `similarity-graph` until CTX-10 supplies cross-shard pair planning. diff --git a/docs/ctx-07-distributed-workload-protocol.md b/docs/ctx-07-distributed-workload-protocol.md index 1c80892..5720b38 100644 --- a/docs/ctx-07-distributed-workload-protocol.md +++ b/docs/ctx-07-distributed-workload-protocol.md @@ -4,9 +4,11 @@ This document is the implementation contract for CTX-07. Its generic protocol, registry, strict JSON models, and deterministic reduction ordering are -implemented in `scimesh/distributed/`. It does not implement a molecular -planner, reducer, API endpoint, database migration, or final artifact. Until -CTX-08 and CTX-09 are complete, shard CSVs remain diagnostic partial results. +implemented in `scimesh/distributed/`. CTX-08 implements the molecular +similarity-search planner, worker adapter, and pure reducer on top of it. This +document does not implement a coordinator API endpoint, database migration, or +durable final artifact. Until CTX-09 is complete, shard CSVs remain diagnostic +partial results. The protocol gives local scientific workloads a coordinator-independent way to validate a job, plan artifact-backed tasks, and later reduce completed outputs. @@ -154,7 +156,10 @@ rank,chembl_id,canonical_smiles,similarity ``` - `rank` is one-based local rank. -- `similarity` uses the local CLI's six-decimal formatting. +- `similarity` uses a round-trip decimal representation of the computed float + (for Python, `repr(similarity)`). This preserves exact cross-shard ranking; + the reducer writes the user-facing final CSV with the local CLI's six-decimal + display formatting. - Rows are sorted by `(-similarity, chembl_id, canonical_smiles)` for `threshold_direction=greater`, or `(similarity, chembl_id, canonical_smiles)` for `less`. @@ -203,7 +208,7 @@ multiplicity. Reduction is independent of worker completion order and uses ## Deferred work -CTX-08 implements the similarity-search planner, runner adapter, reducer, and -comparison against the local CLI. CTX-09 persists the final artifact and job -state. CTX-10 defines graph-specific triangular block plans; it must not reuse -the search shard scheme without its pair-coverage invariants. +CTX-09 materializes the planned shard files as coordinator artifacts, invokes +the registered reducer once, and persists its final artifact/job state. CTX-10 +defines graph-specific triangular block plans; it must not reuse the search +shard scheme without its pair-coverage invariants. diff --git a/docs/worker-daemon-task.md b/docs/worker-daemon-task.md index e4d2f07..2a0e122 100644 --- a/docs/worker-daemon-task.md +++ b/docs/worker-daemon-task.md @@ -122,7 +122,10 @@ Content-Type: application/json }, "metrics": { "elapsed_seconds": 12.4, - "processed_rows": 10000 + "scanned_rows": 10000, + "valid_molecules": 9876, + "invalid_smiles": 124, + "matches_emitted": 20 } } ``` @@ -189,8 +192,11 @@ class Runner(Protocol): """Run one task and return output artifacts plus safe metrics.""" ``` -`SciMeshRunner` should map `workload` and validated parameters to the existing -SciMesh CLI. For example, a `similarity-search` task invokes: +`SciMeshRunner` maps an allowlisted workload and validated parameters to the +local SciMesh reference functions. A planned `similarity-search` task contains +a resolved `query_smiles` (never `query_id`) and writes one exact local top-k +partial CSV plus the metrics above. Legacy single-shard tasks may still use the +CLI compatibility path: ```text scimesh similarity-search --query-id ... --output /result.csv diff --git a/scimesh/distributed/__init__.py b/scimesh/distributed/__init__.py index e194dc9..40b5026 100644 --- a/scimesh/distributed/__init__.py +++ b/scimesh/distributed/__init__.py @@ -12,17 +12,25 @@ from .models import ( FinalResult, PlannedTask, ) -from .registry import DistributedWorkloadRegistry, PlanningService, WorkloadDescription +from .registry import ( + DistributedWorkloadRegistry, + PlanningService, + WorkloadDescription, + default_distributed_registry, +) +from .similarity_search import SimilaritySearchDistributedWorkload from .workload import DistributedWorkload __all__ = [ "ArtifactReference", "CompletedPartial", + "default_distributed_registry", "DistributedPlan", "DistributedWorkload", "DistributedWorkloadRegistry", "FinalResult", "PlannedTask", "PlanningService", + "SimilaritySearchDistributedWorkload", "WorkloadDescription", ] diff --git a/scimesh/distributed/registry.py b/scimesh/distributed/registry.py index 4de6a83..868b27f 100644 --- a/scimesh/distributed/registry.py +++ b/scimesh/distributed/registry.py @@ -50,7 +50,8 @@ class PlanningService: It writes neither jobs nor artifacts. A Go coordinator bridge can therefore validate and produce a plan before opening its own all-or-nothing persistence - transaction; CTX-08/09 will implement that concrete bridge and reducers. + transaction; CTX-09 will implement that concrete bridge and durable result + orchestration. """ def __init__(self, registry: DistributedWorkloadRegistry) -> None: @@ -94,3 +95,14 @@ class PlanningService: if not isinstance(result, FinalResult): raise ValueError("distributed reducer must return a FinalResult") return result + + +def default_distributed_registry() -> DistributedWorkloadRegistry: + """Return the currently supported distributed scientific workloads.""" + # Delayed import keeps the generic registry independent of concrete RDKit + # workloads and avoids making the contract layer import application setup. + from .similarity_search import SimilaritySearchDistributedWorkload + + registry = DistributedWorkloadRegistry() + registry.register(SimilaritySearchDistributedWorkload()) + return registry diff --git a/scimesh/distributed/similarity_search.py b/scimesh/distributed/similarity_search.py new file mode 100644 index 0000000..e73c9c6 --- /dev/null +++ b/scimesh/distributed/similarity_search.py @@ -0,0 +1,380 @@ +"""Distributed planning and reduction for exact molecular similarity search.""" + +from __future__ import annotations + +import csv +import hashlib +import heapq +import math +from pathlib import Path +from typing import Any, Iterator, Mapping, Sequence +from uuid import UUID, uuid5 + +from rdkit import Chem + +from scimesh.chemistry.dataset import MoleculeRecord, find_molecule_by_id, parse_smiles +from scimesh.chemistry.fingerprints import FP_RADIUS, FP_SIZE +from scimesh.workloads.similarity_search import ( + SimilarityMatch, + _HeapEntry, + search_similar, + write_search_results, +) + +from .models import ArtifactReference, CompletedPartial, DistributedPlan, FinalResult, PlannedTask + + +_TSV_CONTENT_TYPE = "text/tab-separated-values" +_CSV_CONTENT_TYPE = "text/csv" +_SEARCH_COLUMNS = ("rank", "chembl_id", "canonical_smiles", "similarity") +_REQUIRED_COLUMNS = {"chembl_id", "canonical_smiles"} + + +def write_similarity_search_partial(output_path: Path, matches: Sequence[SimilarityMatch]) -> None: + """Write a worker partial with a round-trip score, not display rounding. + + The public final CSV continues to use the local CLI's six-decimal display. + A reducer needs the full binary float representation to rank candidates + from separate shards exactly as the single-process reference does. + """ + output_path.parent.mkdir(parents=True, exist_ok=True) + with output_path.open("w", encoding="utf-8", newline="") as destination: + writer = csv.DictWriter(destination, fieldnames=_SEARCH_COLUMNS) + writer.writeheader() + for rank, match in enumerate(matches, start=1): + writer.writerow({ + "rank": rank, + "chembl_id": match.molecule_id, + "canonical_smiles": match.smiles, + "similarity": repr(match.similarity), + }) + + +class SimilaritySearchDistributedWorkload: + """Planner/reducer for exact global top-k Tanimoto similarity search.""" + + name = "similarity-search" + description = "Exact top-k molecular similarity search over deterministic TSV shards." + + def validate_job(self, parameters: Mapping[str, object]) -> None: + allowed = { + "query_id", "query_smiles", "top_k", "threshold", + "threshold_direction", "max_rows", "progress_every", + } + unknown = set(parameters) - allowed + if unknown: + raise ValueError(f"unsupported similarity-search parameters: {', '.join(sorted(unknown))}") + query_id = parameters.get("query_id") + query_smiles = parameters.get("query_smiles") + if (query_id is None) == (query_smiles is None): + raise ValueError("exactly one of query_id or query_smiles is required") + if query_id is not None: + self._string(query_id, "query_id") + if query_smiles is not None: + self._string(query_smiles, "query_smiles") + self._positive_int(parameters.get("top_k", 20), "top_k") + if "max_rows" in parameters: + self._positive_int(parameters["max_rows"], "max_rows") + if "progress_every" in parameters: + self._nonnegative_int(parameters["progress_every"], "progress_every") + if "threshold" in parameters: + self._unit_interval(parameters["threshold"], "threshold") + if "threshold_direction" in parameters and parameters["threshold_direction"] not in {"greater", "less"}: + raise ValueError("threshold_direction must be 'greater' or 'less'") + + def plan( + self, + input_path: Path, + input_artifact_id: str, + parameters: Mapping[str, object], + shard_rows: int, + workspace: Path, + ) -> DistributedPlan: + self.validate_job(parameters) + if not input_path.is_file(): + raise ValueError("input_path must be a readable dataset file") + if isinstance(shard_rows, bool) or not isinstance(shard_rows, int) or shard_rows < 1: + raise ValueError("shard_rows must be a positive integer") + try: + input_id = UUID(input_artifact_id) + except ValueError as error: + raise ValueError("input_artifact_id must be a UUID") from error + + query_smiles, query_source = self._resolve_query(input_path, parameters) + resolved = self._resolved_parameters(parameters, query_smiles, query_source) + workspace.mkdir(parents=True, exist_ok=True) + shard_paths: list[Path] = [] + try: + shard_paths = self._write_shards(input_path, workspace, shard_rows, resolved.get("max_rows")) + tasks = tuple( + PlannedTask( + chunk_index=index, + input_artifact=ArtifactReference( + artifact_id=str(uuid5(input_id, f"scimesh:similarity-search:shard:{index}")), + sha256=_sha256_file(path), + content_type=_TSV_CONTENT_TYPE, + ), + parameters=self._task_parameters(resolved), + ) + for index, path in enumerate(shard_paths) + ) + except Exception: + for path in shard_paths: + path.unlink(missing_ok=True) + raise + return DistributedPlan(self.name, resolved, tasks) + + def reduce( + self, + partial_results: Sequence[CompletedPartial], + parameters: Mapping[str, object], + workspace: Path, + ) -> FinalResult: + """Merge materialized partial CSVs into one deterministic final CSV. + + The coordinator bridge materializes each downloaded artifact at + ``workspace / artifact_id`` before it calls this method. Those local + paths are an ephemeral bridge detail, never present in the plan or task + payload. CTX-09 owns the durable final-artifact upload and job state. + """ + if not partial_results: + raise ValueError("at least one partial result is required") + resolved = self._validate_resolved_parameters(parameters) + top_k = resolved["top_k"] + direction = resolved["threshold_direction"] + heap: list[_HeapEntry] = [] + + ordered_partials = tuple(sorted(partial_results, key=lambda partial: partial.chunk_index)) + indexes = [partial.chunk_index for partial in ordered_partials] + if len(indexes) != len(set(indexes)): + raise ValueError("partial results must have unique chunk_index values") + for partial in ordered_partials: + path = workspace / partial.artifact.artifact_id + if not path.is_file(): + raise ValueError("materialized partial result is missing") + if _sha256_file(path) != partial.artifact.sha256: + raise ValueError("materialized partial result checksum does not match its artifact reference") + for match in self._read_partial(path, direction): + rank_key = match.sort_key(direction) + entry = _HeapEntry(match, rank_key) + if len(heap) < top_k: + heapq.heappush(heap, entry) + elif rank_key < heap[0].rank_key: + heapq.heapreplace(heap, entry) + + matches = sorted((entry.match for entry in heap), key=lambda match: match.sort_key(direction)) + output = workspace / "result.csv" + write_search_results(output, matches) + final_id = uuid5( + UUID(ordered_partials[0].artifact.artifact_id), + "scimesh:similarity-search:final:" + ",".join( + partial.artifact.artifact_id for partial in ordered_partials + ), + ) + return FinalResult( + ArtifactReference(str(final_id), _sha256_file(output), _CSV_CONTENT_TYPE), + {"matches_emitted": len(matches), "partial_count": len(ordered_partials)}, + ) + + def _resolve_query( + self, input_path: Path, parameters: Mapping[str, object] + ) -> tuple[str, dict[str, str]]: + query_id = parameters.get("query_id") + if isinstance(query_id, str): + record = find_molecule_by_id(input_path, query_id) + return Chem.MolToSmiles(record.molecule, canonical=True), {"kind": "chembl_id", "value": query_id} + supplied = parameters["query_smiles"] + assert isinstance(supplied, str) # checked by validate_job + molecule = parse_smiles(supplied) + if molecule is None: + raise ValueError("query_smiles is invalid") + return Chem.MolToSmiles(molecule, canonical=True), {"kind": "smiles", "value": supplied} + + def _resolved_parameters( + self, parameters: Mapping[str, object], query_smiles: str, query_source: Mapping[str, str] + ) -> dict[str, object]: + resolved: dict[str, object] = { + "query_smiles": query_smiles, + "query_source": dict(query_source), + "top_k": self._positive_int(parameters.get("top_k", 20), "top_k"), + "threshold_direction": parameters.get("threshold_direction", "greater"), + "fingerprint": {"algorithm": "morgan", "radius": FP_RADIUS, "fp_size": FP_SIZE}, + } + if "threshold" in parameters: + resolved["threshold"] = self._unit_interval(parameters["threshold"], "threshold") + if "max_rows" in parameters: + resolved["max_rows"] = self._positive_int(parameters["max_rows"], "max_rows") + if "progress_every" in parameters: + resolved["progress_every"] = self._nonnegative_int(parameters["progress_every"], "progress_every") + return resolved + + def _validate_resolved_parameters(self, parameters: Mapping[str, object]) -> dict[str, object]: + query_smiles = self._string(parameters.get("query_smiles"), "query_smiles") + if parse_smiles(query_smiles) is None: + raise ValueError("query_smiles is invalid") + resolved = self._resolved_parameters( + parameters, + Chem.MolToSmiles(parse_smiles(query_smiles), canonical=True), + {"kind": "resolved", "value": query_smiles}, + ) + # A reducer receives immutable plan metadata, whose query source and + # fixed fingerprint are observational context rather than worker input. + if "fingerprint" in parameters: + fingerprint = parameters["fingerprint"] + if fingerprint != {"algorithm": "morgan", "radius": FP_RADIUS, "fp_size": FP_SIZE}: + raise ValueError("resolved fingerprint does not match SciMesh defaults") + return resolved + + @staticmethod + def _task_parameters(resolved: Mapping[str, object]) -> dict[str, object]: + # max_rows is applied before sharding. Passing it to each task would + # silently scan N rows per shard instead of the requested global prefix. + return { + key: value for key, value in resolved.items() + if key in {"query_smiles", "top_k", "threshold", "threshold_direction", "progress_every"} + } + + def _write_shards( + self, input_path: Path, workspace: Path, shard_rows: int, max_rows: object + ) -> list[Path]: + limit = int(max_rows) if isinstance(max_rows, int) else None + paths: list[Path] = [] + current: Path | None = None + destination = None + rows_in_shard = 0 + seen_rows = 0 + try: + with input_path.open("r", encoding="utf-8", newline="") as source: + reader = csv.DictReader(source, delimiter="\t") + fieldnames = reader.fieldnames or [] + if not _REQUIRED_COLUMNS.issubset(set(fieldnames)): + missing = sorted(_REQUIRED_COLUMNS - set(fieldnames)) + raise ValueError(f"dataset is missing required columns: {', '.join(missing)}") + for row in reader: + if limit is not None and seen_rows >= limit: + break + if destination is None or rows_in_shard == shard_rows: + if destination is not None: + destination.close() + current = workspace / f"shard-{len(paths)}.tsv" + destination = current.open("w", encoding="utf-8", newline="") + writer = csv.DictWriter(destination, fieldnames=fieldnames, delimiter="\t", lineterminator="\n") + writer.writeheader() + paths.append(current) + rows_in_shard = 0 + writer.writerow(row) + rows_in_shard += 1 + seen_rows += 1 + finally: + if destination is not None: + destination.close() + if not paths: + raise ValueError("dataset has no data rows") + return paths + + @staticmethod + def _read_partial(path: Path, direction: object) -> Iterator[SimilarityMatch]: + if not path.is_file(): + raise ValueError("materialized partial result is missing") + if direction not in {"greater", "less"}: + raise ValueError("threshold_direction must be 'greater' or 'less'") + previous_key: tuple[float, str, str] | None = None + with path.open("r", encoding="utf-8", newline="") as source: + reader = csv.DictReader(source) + if tuple(reader.fieldnames or ()) != _SEARCH_COLUMNS: + raise ValueError("partial result has an invalid CSV header") + for expected_rank, row in enumerate(reader, start=1): + if set(row) != set(_SEARCH_COLUMNS) or row["rank"] != str(expected_rank): + raise ValueError("partial result has an invalid rank") + try: + similarity = float(row["similarity"]) + except (TypeError, ValueError) as error: + raise ValueError("partial result has an invalid similarity") from error + if not math.isfinite(similarity) or not 0 <= similarity <= 1: + raise ValueError("partial result has an invalid similarity") + match = SimilarityMatch(similarity, row["chembl_id"], row["canonical_smiles"]) + key = match.sort_key(direction) + if previous_key is not None and key < previous_key: + raise ValueError("partial result is not sorted deterministically") + previous_key = key + yield match + + @staticmethod + def _string(value: object, name: str) -> str: + if not isinstance(value, str) or not value.strip() or len(value) > 200: + raise ValueError(f"{name} must be a non-empty string") + return value + + @staticmethod + def _positive_int(value: object, name: str) -> int: + if isinstance(value, bool) or not isinstance(value, int) or value < 1: + raise ValueError(f"{name} must be a positive integer") + return value + + @staticmethod + def _nonnegative_int(value: object, name: str) -> int: + if isinstance(value, bool) or not isinstance(value, int) or value < 0: + raise ValueError(f"{name} must be a non-negative integer") + return value + + @staticmethod + def _unit_interval(value: object, name: str) -> float: + if isinstance(value, bool) or not isinstance(value, (int, float)) or not math.isfinite(value) or not 0 <= value <= 1: + raise ValueError(f"{name} must be a number between 0 and 1") + return float(value) + + +def run_similarity_search_shard( + input_path: Path, parameters: Mapping[str, object], output_path: Path +) -> dict[str, int]: + """Run one planned shard using the local reference implementation. + + This is the worker adapter used by CTX-08. It deliberately accepts only + resolved ``query_smiles``: resolving an identifier independently in each + shard would make the distributed search scientifically invalid. + """ + allowed = {"query_smiles", "top_k", "threshold", "threshold_direction", "progress_every"} + unknown = set(parameters) - allowed + if unknown: + raise ValueError(f"unsupported similarity-search parameters: {', '.join(sorted(unknown))}") + query_smiles = parameters.get("query_smiles") + if not isinstance(query_smiles, str) or not query_smiles.strip(): + raise ValueError("query_smiles is required for a distributed shard") + molecule = parse_smiles(query_smiles) + if molecule is None: + raise ValueError("query_smiles is invalid") + top_k = SimilaritySearchDistributedWorkload._positive_int(parameters.get("top_k", 20), "top_k") + threshold = None + if "threshold" in parameters: + threshold = SimilaritySearchDistributedWorkload._unit_interval(parameters["threshold"], "threshold") + direction = parameters.get("threshold_direction", "greater") + if direction not in {"greater", "less"}: + raise ValueError("threshold_direction must be 'greater' or 'less'") + progress_every = 0 + if "progress_every" in parameters: + progress_every = SimilaritySearchDistributedWorkload._nonnegative_int( + parameters["progress_every"], "progress_every" + ) + result = search_similar( + input_path, + MoleculeRecord("query", query_smiles, molecule), + top_k=top_k, + progress_every=progress_every, + threshold=threshold, + threshold_direction=direction, + ) + write_similarity_search_partial(output_path, result.matches) + return { + "scanned_rows": result.stats.scanned, + "valid_molecules": result.stats.valid, + "invalid_smiles": result.stats.invalid, + "matches_emitted": len(result.matches), + } + + +def _sha256_file(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as source: + for block in iter(lambda: source.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() diff --git a/scimesh/worker/daemon.py b/scimesh/worker/daemon.py index f8e51a7..6ef0008 100644 --- a/scimesh/worker/daemon.py +++ b/scimesh/worker/daemon.py @@ -9,6 +9,7 @@ from pathlib import Path import random import re import shutil +import subprocess import threading import time from datetime import datetime, timezone @@ -202,12 +203,23 @@ class WorkerDaemon: def _report_failure(self, task: ClaimedTask, error: Exception) -> None: message = self._sanitize_error_message(error) try: - self.coordinator.fail(task, {"worker_id": self._worker_id(), "attempt": task.attempt, "error_code": type(error).__name__, "error_message": message}) + self.coordinator.fail(task, { + "worker_id": self._worker_id(), + "attempt": task.attempt, + "error_code": type(error).__name__, + "error_message": message, + "retryable": self._is_retryable(error), + }) except CoordinatorTransientError: raise except Exception: self._log("failed", task, error_type="FailureReportError") + @staticmethod + def _is_retryable(error: Exception) -> bool: + """Retry transient worker/transport failures, never invalid scientific input.""" + return not isinstance(error, (ValueError, FileNotFoundError, subprocess.CalledProcessError)) + def _sanitize_error_message(self, error: Exception) -> str: """Keep coordinator-visible failures useful without exposing local paths.""" message = str(error).replace(str(self.config.work_dir), "") diff --git a/scimesh/worker/runners.py b/scimesh/worker/runners.py index 46f219a..f924636 100644 --- a/scimesh/worker/runners.py +++ b/scimesh/worker/runners.py @@ -7,6 +7,8 @@ import subprocess import sys from typing import Protocol +from scimesh.distributed.similarity_search import run_similarity_search_shard + from .models import ClaimedTask, ProducedArtifact, RunResult @@ -34,6 +36,12 @@ class SciMeshRunner: query_id, query_smiles = params.get("query_id"), params.get("query_smiles") if (query_id is None) == (query_smiles is None): raise ValueError("exactly one of query_id or query_smiles is required") + if query_smiles is not None and "max_rows" not in params: + metrics = run_similarity_search_shard(input_path, params, output_path) + return RunResult((ProducedArtifact(output_path, "text/csv"),), metrics) + # Legacy URI jobs may still use query_id or an explicitly task-local + # max_rows value. CTX-08 plans never create those payloads; retain + # CLI execution only for backwards compatibility at this boundary. top_k = self._positive_int(params, "top_k", default=20) command += ["--query-id", self._string(params, "query_id")] if query_id is not None else ["--query-smiles", self._string(params, "query_smiles")] command += ["--top-k", str(top_k)] diff --git a/tests/test_distributed_similarity_search.py b/tests/test_distributed_similarity_search.py new file mode 100644 index 0000000..7759532 --- /dev/null +++ b/tests/test_distributed_similarity_search.py @@ -0,0 +1,184 @@ +"""Scientific reference tests for the CTX-08 distributed search workload.""" + +from __future__ import annotations + +import csv +import hashlib +from pathlib import Path +from uuid import NAMESPACE_URL, uuid5 + +import pytest + +from scimesh.chemistry.dataset import find_molecule_by_id +from scimesh.distributed import ( + ArtifactReference, + CompletedPartial, + PlanningService, + default_distributed_registry, +) +from scimesh.distributed.registry import DistributedWorkloadRegistry +from scimesh.distributed.similarity_search import ( + SimilaritySearchDistributedWorkload, + run_similarity_search_shard, + write_similarity_search_partial, +) +from scimesh.workloads.similarity_search import search_similar, write_search_results + + +def make_dataset(path: Path) -> None: + path.write_text( + "chembl_id\tcanonical_smiles\textra\n" + "CHEMBL_QUERY\tCCO\tquery\n" + "CHEMBL_A\tCCCO\ta\n" + "CHEMBL_B\tCCCC\tb\n" + "CHEMBL_INVALID\tnot-a-smiles\tbad\n" + "CHEMBL_DUPLICATE\tCCO\tduplicate\n" + "CHEMBL_C\tCCN\tc\n", + encoding="utf-8", + ) + + +def planner() -> tuple[PlanningService, SimilaritySearchDistributedWorkload]: + workload = SimilaritySearchDistributedWorkload() + registry = DistributedWorkloadRegistry() + registry.register(workload) + return PlanningService(registry), workload + + +def test_default_registry_exposes_only_supported_distributed_search() -> None: + assert [item.name for item in default_distributed_registry().descriptions()] == ["similarity-search"] + + +def checksum(path: Path) -> str: + return hashlib.sha256(path.read_bytes()).hexdigest() + + +def test_query_id_is_resolved_once_before_deterministic_shards( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + dataset = tmp_path / "chembl.tsv" + workspace = tmp_path / "workspace" + make_dataset(dataset) + service, _ = planner() + calls = 0 + real_find = find_molecule_by_id + + def count_find(path: Path, query_id: str) -> MoleculeRecord: + nonlocal calls + calls += 1 + return real_find(path, query_id) + + monkeypatch.setattr("scimesh.distributed.similarity_search.find_molecule_by_id", count_find) + input_id = str(uuid5(NAMESPACE_URL, "dataset")) + plan = service.plan( + "similarity-search", dataset, input_id, + {"query_id": "CHEMBL_QUERY", "top_k": 3, "max_rows": 5, "progress_every": 0}, + 2, workspace, + ) + + assert calls == 1 + assert plan.resolved_parameters["query_smiles"] == "CCO" + assert plan.resolved_parameters["query_source"] == {"kind": "chembl_id", "value": "CHEMBL_QUERY"} + assert [task.chunk_index for task in plan.tasks] == [0, 1, 2] + assert all("query_id" not in task.parameters for task in plan.tasks) + assert all("max_rows" not in task.parameters for task in plan.tasks) + assert all(task.parameters["query_smiles"] == "CCO" for task in plan.tasks) + assert [ + sum(1 for _ in path.open(encoding="utf-8")) - 1 + for path in sorted(workspace.glob("shard-*.tsv")) + ] == [2, 2, 1] + + +def test_distributed_reduction_matches_single_process_reference(tmp_path: Path) -> None: + dataset = tmp_path / "chembl.tsv" + workspace = tmp_path / "workspace" + make_dataset(dataset) + service, workload = planner() + plan = service.plan( + "similarity-search", dataset, str(uuid5(NAMESPACE_URL, "dataset")), + {"query_smiles": "CCO", "top_k": 3, "threshold": 0.0}, 2, workspace, + ) + + partials: list[CompletedPartial] = [] + # Worker two finishes the latter shards first. Worker one loses its first + # attempt for shard zero, then retries it last. The reducer must remain + # independent of both completion and retry order. + for task in reversed(plan.tasks): + shard = workspace / f"shard-{task.chunk_index}.tsv" + temporary_partial = workspace / f"worker-output-{task.chunk_index}.csv" + metrics = run_similarity_search_shard(shard, task.parameters, temporary_partial) + partial_id = str(uuid5(NAMESPACE_URL, f"partial:{task.chunk_index}")) + partials.append( + CompletedPartial( + task.chunk_index, + ArtifactReference( + partial_id, checksum(temporary_partial), "text/csv", + ), + metrics, + ) + ) + # The reducer materializes result files under their own coordinator IDs, + # not shard input IDs. Keep this fixture faithful to that boundary. + temporary_partial.rename(workspace / partial_id) + + final = workload.reduce(tuple(partials), plan.resolved_parameters, workspace) + reference = tmp_path / "reference.csv" + query_record = find_molecule_by_id(dataset, "CHEMBL_QUERY") + write_search_results(reference, search_similar(dataset, query_record, top_k=3, threshold=0.0).matches) + + assert (workspace / "result.csv").read_bytes() == reference.read_bytes() + assert final.metrics == {"matches_emitted": 3, "partial_count": 3} + rows = list(csv.DictReader((workspace / "result.csv").open(encoding="utf-8"))) + assert {row["chembl_id"] for row in rows}.isdisjoint({"CHEMBL_QUERY", "CHEMBL_DUPLICATE"}) + + +def test_reducer_rejects_unsorted_or_invalid_partial_csv(tmp_path: Path) -> None: + workspace = tmp_path / "workspace" + workspace.mkdir() + workload = SimilaritySearchDistributedWorkload() + artifact_id = str(uuid5(NAMESPACE_URL, "bad")) + partial_path = workspace / artifact_id + partial_path.write_text( + "rank,chembl_id,canonical_smiles,similarity\n1,A,CC,0.1\n2,B,CCC,0.9\n", + encoding="utf-8", + ) + artifact = ArtifactReference(artifact_id, checksum(partial_path), "text/csv") + + with pytest.raises(ValueError, match="not sorted"): + workload.reduce( + (CompletedPartial(0, artifact, {"scanned_rows": 2}),), + {"query_smiles": "CCO", "top_k": 2, "threshold_direction": "greater", "fingerprint": {"algorithm": "morgan", "radius": 2, "fp_size": 2048}}, + workspace, + ) + + +def test_partial_csv_preserves_exact_scores_for_global_ranking(tmp_path: Path) -> None: + partial = tmp_path / "partial.csv" + # Both values look identical in a six-decimal final CSV. The exact value + # must survive shard transport so the global reducer can still rank them. + from scimesh.workloads.similarity_search import SimilarityMatch + + write_similarity_search_partial( + partial, + [SimilarityMatch(0.50000049, "A", "CC"), SimilarityMatch(0.50000048, "B", "CCC")], + ) + values = list(csv.DictReader(partial.open(encoding="utf-8"))) + assert values[0]["similarity"] == repr(0.50000049) + assert values[1]["similarity"] == repr(0.50000048) + + +def test_planner_rejects_fingerprint_override_and_invalid_query(tmp_path: Path) -> None: + dataset = tmp_path / "chembl.tsv" + make_dataset(dataset) + service, _ = planner() + + with pytest.raises(ValueError, match="unsupported similarity-search parameters"): + service.plan( + "similarity-search", dataset, str(uuid5(NAMESPACE_URL, "dataset")), + {"query_smiles": "CCO", "fingerprint": {"radius": 1}}, 2, tmp_path / "workspace", + ) + with pytest.raises(ValueError, match="query_smiles is invalid"): + service.plan( + "similarity-search", dataset, str(uuid5(NAMESPACE_URL, "dataset")), + {"query_smiles": "invalid"}, 2, tmp_path / "workspace", + ) diff --git a/tests/test_worker_daemon.py b/tests/test_worker_daemon.py index 7bf6f65..5f77098 100644 --- a/tests/test_worker_daemon.py +++ b/tests/test_worker_daemon.py @@ -106,6 +106,90 @@ def test_claims_runs_uploads_and_submits_csv(tmp_path: Path) -> None: } +def test_worker_executes_a_resolved_similarity_search_shard(tmp_path: Path) -> None: + content = b"chembl_id\tcanonical_smiles\nQUERY\tCCO\nMATCH\tCCCO\nINVALID\tnot-a-smiles\n" + task = make_task(content) + task = ClaimedTask( + task.task_id, task.attempt, task.lease_expires_at, task.workload, task.input, + {"query_smiles": "CCO", "top_k": 5, "progress_every": 0}, + ) + worker, coordinator, artifacts, _, _ = daemon(tmp_path, task, content) + worker.runner = SciMeshRunner() + + assert worker.run_once() == RunOnceOutcome(claimed=True, completed=True) + output = artifacts.uploaded[0][2].read_text(encoding="utf-8") + assert output.startswith("rank,chembl_id,canonical_smiles,similarity\n") + metrics = coordinator.submissions[0]["metrics"] + assert metrics["scanned_rows"] == 3 + assert metrics["valid_molecules"] == 2 + assert metrics["invalid_smiles"] == 1 + assert metrics["matches_emitted"] == 1 + assert isinstance(metrics["elapsed_seconds"], float) + + +def test_two_workers_complete_resolved_shards_after_one_retry(tmp_path: Path) -> None: + content = b"chembl_id\tcanonical_smiles\nQUERY\tCCO\nMATCH\tCCCO\n" + first = make_task(content) + first = ClaimedTask( + "retry-task", 1, first.lease_expires_at, "similarity-search", first.input, + {"query_smiles": "CCO", "top_k": 5}, + ) + second = ClaimedTask( + "other-task", 1, first.lease_expires_at, "similarity-search", first.input, + {"query_smiles": "CCO", "top_k": 5}, + ) + + class RetryCoordinator(FakeCoordinator): + def __init__(self) -> None: + super().__init__(None) + self.queue = [first, second] + self.claimants: list[str] = [] + + def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None: + self.claimants.append(worker_id) + return self.queue.pop(0) if self.queue else None + + def fail(self, task: ClaimedTask, payload: dict) -> None: + self.failures.append(payload) + if task.task_id == "retry-task" and task.attempt == 1 and payload["retryable"]: + self.queue.append( + ClaimedTask( + task.task_id, 2, task.lease_expires_at, task.workload, task.input, + task.parameters, + ) + ) + + class FailFirstAttempt: + def __init__(self) -> None: + self.calls = 0 + self.delegate = SciMeshRunner() + + def run(self, task: ClaimedTask, task_dir: Path) -> RunResult: + self.calls += 1 + if self.calls == 1: + raise RuntimeError("simulated retryable shard failure") + return self.delegate.run(task, task_dir) + + coordinator = RetryCoordinator() + artifacts = FakeArtifacts(content) + worker_a = WorkerDaemon( + WorkerConfig("https://example.test", "worker-a", tmp_path / "worker-a"), + coordinator, artifacts, FailFirstAttempt(), + ) + worker_b = WorkerDaemon( + WorkerConfig("https://example.test", "worker-b", tmp_path / "worker-b"), + coordinator, artifacts, SciMeshRunner(), + ) + + assert worker_a.run_once() == RunOnceOutcome(claimed=True, completed=False) + assert worker_b.run_once() == RunOnceOutcome(claimed=True, completed=True) + assert worker_a.run_once() == RunOnceOutcome(claimed=True, completed=True) + assert coordinator.claimants == ["worker-a", "worker-b", "worker-a"] + assert len(coordinator.failures) == 1 + assert coordinator.failures[0]["retryable"] is True + assert len(coordinator.submissions) == 2 + + def test_no_task_does_not_create_directory(tmp_path: Path) -> None: worker, _, _, runner, config = daemon(tmp_path, None, b"") assert worker.run_once() == RunOnceOutcome(claimed=False, completed=False) @@ -161,6 +245,7 @@ def test_interrupting_an_active_task_reports_a_sanitized_failure(tmp_path: Path) "attempt": 1, "error_code": "InterruptedError", "error_message": "worker interrupted by operator", + "retryable": True, } ] @@ -234,6 +319,7 @@ def test_bad_checksum_reports_failure_without_running(tmp_path: Path) -> None: assert worker.run_once() == RunOnceOutcome(claimed=True, completed=False) assert runner.calls == 0 assert coordinator.failures[0]["error_code"] == "ValueError" + assert coordinator.failures[0]["retryable"] is False assert not coordinator.submissions @@ -356,30 +442,35 @@ def test_runner_maps_graph_and_smiles_search_parameters(tmp_path: Path, monkeypa runner = SciMeshRunner() graph = ClaimedTask("graph", 1, "2026-07-30T00:00:00Z", "similarity-graph", InputArtifact("https://example/input", "x"), {"threshold": 0.2, "threshold_direction": "less", "block_size": 42, "max_rows": 7, "progress_every": 0}) search = ClaimedTask("search", 1, "2026-07-30T00:00:00Z", "similarity-search", InputArtifact("https://example/input", "x"), {"query_smiles": "CCO", "top_k": 3}) + search_dir = tmp_path / "search" + search_dir.mkdir() + (search_dir / "input").write_text( + "chembl_id\tcanonical_smiles\nA\tCCO\nB\tCCCO\n", encoding="utf-8" + ) runner.run(graph, tmp_path / "graph") - runner.run(search, tmp_path / "search") + result = runner.run(search, search_dir) assert "--threshold-direction" in commands[0] and "less" in commands[0] assert "--block-size" in commands[0] and "42" in commands[0] assert "--max-rows" in commands[0] and "7" in commands[0] - assert "--query-smiles" in commands[1] and "CCO" in commands[1] + assert len(commands) == 1 + assert result.metrics == { + "scanned_rows": 2, "valid_molecules": 2, "invalid_smiles": 0, "matches_emitted": 1, + } def test_runner_accepts_coordinator_workload_names(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: - commands: list[list[str]] = [] - - def fake_run(command: list[str], **_: object) -> None: - commands.append(command) - output = Path(command[command.index("--output") + 1]) - output.parent.mkdir(parents=True, exist_ok=True) - output.write_text("a,b\\n", encoding="utf-8") - - monkeypatch.setattr("scimesh.worker.runners.subprocess.run", fake_run) + task_dir = tmp_path / "search" + task_dir.mkdir() + (task_dir / "input").write_text( + "chembl_id\tcanonical_smiles\nA\tCCO\nB\tCCCO\n", encoding="utf-8" + ) task = ClaimedTask( "search", 1, "2026-07-30T00:00:00Z", "similarity_search", InputArtifact("https://example/input", "a" * 64), {"query_smiles": "CCO"}, ) - SciMeshRunner().run(task, tmp_path / "search") - assert commands[0][3] == "similarity-search" + result = SciMeshRunner().run(task, task_dir) + assert result.metrics["matches_emitted"] == 1 + assert (task_dir / "result.csv").is_file() def test_claimed_task_rejects_path_traversal_and_invalid_metadata() -> None: @@ -457,21 +548,15 @@ def test_relative_work_dir_is_normalized_for_runner_subprocesses( task_dir = config.work_dir / "task" / "1" task_dir.mkdir(parents=True) - (task_dir / "input").write_text("fixture", encoding="utf-8") - command: list[str] = [] - - def fake_run(args: list[str], **_: object) -> None: - command.extend(args) - output = Path(args[args.index("--output") + 1]) - output.write_text("id,score\n", encoding="utf-8") - - monkeypatch.setattr("scimesh.worker.runners.subprocess.run", fake_run) + (task_dir / "input").write_text( + "chembl_id\tcanonical_smiles\nA\tCCO\nB\tCCCO\n", encoding="utf-8" + ) task = ClaimedTask( "task", 1, "2026-07-30T00:00:00Z", "similarity-search", InputArtifact("https://example.test/input", "a" * 64), {"query_smiles": "CCO"}, ) SciMeshRunner().run(task, task_dir) - assert command[4] == str(task_dir / "input") + assert (task_dir / "result.csv").is_file() def test_worker_registration_sets_returned_identity(tmp_path: Path) -> None: