"""SDK-built ``similarity-search`` workload definition and handlers. A thin subclass of ``MapReduceWorkload``: the SDK assembles the manifest, stages, workflow, and digest-pinned handlers. This module declares the search scientific contract (plan-time query resolution, deterministic sharding, local top-k per shard, bounded heap merge) and the hooks that partition, compute, and merge. """ from __future__ import annotations from pathlib import Path from typing import Any, Mapping, Sequence from rdkit import Chem from scimesh.chemistry.dataset import find_molecule_by_id, parse_smiles from scimesh.chemistry.fingerprints import FP_RADIUS, FP_SIZE from scimesh.sdk.artifacts import ArtifactSchema, ComponentRef, PortSpec from scimesh.sdk.batch import MapReduceWorkload from scimesh.sdk.identity import SchemaRef, WorkloadId from scimesh.sdk.plans import JobRequest, ValidatedJob from scimesh.sdk.registry import WorkloadDefinition from ..environment import current_environment_digest, current_scimesh_package_digest from .core import merge_search_partials, run_search_shard, write_search_shards MAP_ENTRY_POINT = "scimesh.workloads.search.definition:map_search@v1" REDUCE_ENTRY_POINT = "scimesh.workloads.search.definition:reduce_search@v1" _MAP_PARAMETERS = ( "query_id", "query_smiles", "top_k", "threshold", "threshold_direction", "progress_every", ) def _parameters_schema() -> dict[str, Any]: return { "type": "object", "additionalProperties": False, "properties": { "query_id": {"type": "string", "minLength": 1, "maxLength": 200}, "query_smiles": {"type": "string", "minLength": 1, "maxLength": 200}, "top_k": {"type": "integer", "minimum": 1}, "threshold": {"type": "number", "minimum": 0, "maximum": 1}, "threshold_direction": {"enum": ["greater", "less"]}, "max_rows": {"type": "integer", "minimum": 1}, "progress_every": {"type": "integer", "minimum": 0}, }, "oneOf": [ {"required": ["query_id"], "not": {"required": ["query_smiles"]}}, {"required": ["query_smiles"], "not": {"required": ["query_id"]}}, ], } def _dataset_schema() -> ArtifactSchema: return ArtifactSchema( SchemaRef("molecule-table", 1), "text/tab-separated-values", "utf-8", max_bytes=10 * 1024 * 1024 * 1024, validator=ComponentRef("delimited-table", 1), validator_configuration={ "required_columns": ["canonical_smiles", "chembl_id"], }, max_records=100_000_000, canonicalizer="scimesh-tsv-v1", ) def _search_table_schema(ref: SchemaRef, canonicalizer: str) -> ArtifactSchema: return ArtifactSchema( ref, "text/csv", "utf-8", max_bytes=1024 * 1024 * 1024, validator=ComponentRef("delimited-table", 1), validator_configuration={ "columns": ["rank", "chembl_id", "canonical_smiles", "similarity"], }, max_records=100_000, canonicalizer=canonicalizer, ) class SimilaritySearchSDKWorkload(MapReduceWorkload): """Exact top-k Tanimoto search over deterministic shards with a bounded merge.""" workload_id = WorkloadId("similarity-search", "1.0.0") description = ( "Exact top-k Tanimoto molecular similarity search over " "deterministic TSV shards with a bounded merge." ) parameters_schema = _parameters_schema() input_port = PortSpec(_dataset_schema()) partial_port = PortSpec( _search_table_schema( SchemaRef("similarity-search-partial", 1), "scimesh-search-partial-v1" ) ) output_port = PortSpec( _search_table_schema( SchemaRef("similarity-search-result", 1), "scimesh-search-result-v1" ) ) map_parameter_names = _MAP_PARAMETERS reduce_parameter_names = _MAP_PARAMETERS + ("query_source", "fingerprint") map_entry_point = MAP_ENTRY_POINT reduce_entry_point = REDUCE_ENTRY_POINT def __init__( self, *, shard_rows: int, package_digest: str, environment_digest: str, ) -> None: if ( isinstance(shard_rows, bool) or not isinstance(shard_rows, int) or shard_rows < 1 ): raise ValueError("shard_rows must be a positive integer") self.shard_rows = shard_rows super().__init__( package_digest=package_digest, environment_digest=environment_digest, ) # ------------------------------------------------------------------ # Scientific hooks # ------------------------------------------------------------------ @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)): raise ValueError(f"{name} must be a number between 0 and 1") return float(value) def domain_validate(self, parameters: Mapping[str, Any]) -> None: unknown = set(parameters) - { "query_id", "query_smiles", "top_k", "threshold", "threshold_direction", "max_rows", "progress_every", } if unknown: raise ValueError( "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 resolved_parameters(self, request: JobRequest) -> dict[str, Any]: parameters = request.parameters query_id = parameters.get("query_id") if isinstance(query_id, str): query_source: dict[str, str] = {"kind": "chembl_id", "value": query_id} else: query_source = { "kind": "smiles", "value": self._string(parameters.get("query_smiles"), "query_smiles"), } resolved: dict[str, Any] = { "query_source": 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 resolved_parameters_for_plan( self, job: ValidatedJob, input_path: Path, resolved: dict[str, Any], ) -> dict[str, Any]: query_id = job.request.parameters.get("query_id") if isinstance(query_id, str): record = find_molecule_by_id(input_path, query_id) resolved["query_smiles"] = Chem.MolToSmiles(record.molecule, canonical=True) else: supplied = job.request.parameters["query_smiles"] assert isinstance(supplied, str) molecule = parse_smiles(supplied) if molecule is None: raise ValueError("query_smiles is invalid") resolved["query_smiles"] = Chem.MolToSmiles(molecule, canonical=True) return resolved def partition_input( self, input_path: Path, parameters: Mapping[str, Any], workspace: Path, ) -> list[Path]: max_rows = parameters.get("max_rows") return write_search_shards( input_path, workspace, self.shard_rows, int(max_rows) if isinstance(max_rows, int) else None, ) def compute_shard( self, inputs: Mapping[str, Path], parameters: Mapping[str, Any], output_path: Path, ) -> Mapping[str, int | float]: return run_search_shard(inputs["input"], parameters, output_path) def reduce_partials( self, partial_paths: Sequence[Path], parameters: Mapping[str, Any], output_path: Path, ) -> Mapping[str, int | float]: return merge_search_partials(partial_paths, parameters, output_path) def similarity_search_sdk_definition( *, shard_rows: int = 10_000, package_digest: str | None = None, environment_digest: str | None = None, ) -> SimilaritySearchSDKWorkload: """Build the SDK-built similarity-search definition for tests.""" return SimilaritySearchSDKWorkload( shard_rows=shard_rows, package_digest=package_digest or current_scimesh_package_digest(), environment_digest=environment_digest or current_environment_digest(), ) def workload_definition() -> WorkloadDefinition: """Installed entry-point factory for the SDK-built similarity-search.""" return similarity_search_sdk_definition().definition()