Replace legacy distributed protocol with SDK-built workloads
This commit is contained in:
@@ -1,186 +0,0 @@
|
||||
"""Contract tests for the coordinator-independent distributed workload boundary."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Mapping, Sequence
|
||||
from uuid import NAMESPACE_URL, uuid5
|
||||
|
||||
import pytest
|
||||
|
||||
from scimesh.distributed import (
|
||||
ArtifactReference,
|
||||
CompletedPartial,
|
||||
DistributedPlan,
|
||||
DistributedWorkloadRegistry,
|
||||
FinalResult,
|
||||
PlannedTask,
|
||||
PlanningService,
|
||||
)
|
||||
|
||||
|
||||
def artifact(seed: str, content_type: str = "text/tab-separated-values") -> ArtifactReference:
|
||||
return ArtifactReference(
|
||||
artifact_id=str(uuid5(NAMESPACE_URL, seed)),
|
||||
sha256=(seed.encode("utf-8").hex() * 64)[:64],
|
||||
content_type=content_type,
|
||||
)
|
||||
|
||||
|
||||
class DummyWorkload:
|
||||
"""A deterministic fake workload used to test the generic CTX-07 bridge."""
|
||||
|
||||
name = "dummy-workload"
|
||||
description = "A deterministic test workload."
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.plan_calls = 0
|
||||
self.received_partials: tuple[CompletedPartial, ...] = ()
|
||||
|
||||
def validate_job(self, parameters: Mapping[str, object]) -> None:
|
||||
if parameters != {"mode": "valid"}:
|
||||
raise ValueError("mode must be valid")
|
||||
|
||||
def plan(
|
||||
self,
|
||||
input_path: Path,
|
||||
input_artifact_id: str,
|
||||
parameters: Mapping[str, object],
|
||||
shard_rows: int,
|
||||
workspace: Path,
|
||||
) -> DistributedPlan:
|
||||
self.plan_calls += 1
|
||||
assert input_path.name == "input.tsv"
|
||||
assert workspace.name == "workspace"
|
||||
return DistributedPlan(
|
||||
workload=self.name,
|
||||
resolved_parameters={"mode": parameters["mode"], "source": input_artifact_id},
|
||||
tasks=(
|
||||
PlannedTask(0, artifact(f"{input_artifact_id}:0"), {"mode": "valid"}),
|
||||
PlannedTask(1, artifact(f"{input_artifact_id}:1"), {"mode": "valid"}),
|
||||
),
|
||||
)
|
||||
|
||||
def reduce(
|
||||
self,
|
||||
partial_results: Sequence[CompletedPartial],
|
||||
parameters: Mapping[str, object],
|
||||
workspace: Path,
|
||||
) -> FinalResult:
|
||||
self.received_partials = tuple(partial_results)
|
||||
return FinalResult(artifact("final", "text/csv"), {"partial_count": len(partial_results)})
|
||||
|
||||
|
||||
def service() -> tuple[PlanningService, DummyWorkload]:
|
||||
workload = DummyWorkload()
|
||||
registry = DistributedWorkloadRegistry()
|
||||
registry.register(workload)
|
||||
return PlanningService(registry), workload
|
||||
|
||||
|
||||
def test_unknown_workload_is_rejected_before_a_plan_is_written(tmp_path: Path) -> None:
|
||||
planner, workload = service()
|
||||
|
||||
with pytest.raises(ValueError, match="unknown distributed workload"):
|
||||
planner.plan(
|
||||
"unknown-workload", tmp_path / "input.tsv", artifact("input").artifact_id,
|
||||
{"mode": "valid"}, 10, tmp_path / "workspace",
|
||||
)
|
||||
|
||||
assert workload.plan_calls == 0
|
||||
|
||||
|
||||
def test_invalid_job_is_rejected_before_the_planner_runs(tmp_path: Path) -> None:
|
||||
planner, workload = service()
|
||||
|
||||
with pytest.raises(ValueError, match="mode must be valid"):
|
||||
planner.plan(
|
||||
"dummy-workload", tmp_path / "input.tsv", artifact("input").artifact_id,
|
||||
{"mode": "invalid"}, 10, tmp_path / "workspace",
|
||||
)
|
||||
|
||||
assert workload.plan_calls == 0
|
||||
|
||||
|
||||
def test_two_shard_plan_is_deterministic_and_json_serializable(tmp_path: Path) -> None:
|
||||
planner, _ = service()
|
||||
input_artifact_id = artifact("input").artifact_id
|
||||
first = planner.plan(
|
||||
"dummy-workload", tmp_path / "input.tsv", input_artifact_id,
|
||||
{"mode": "valid"}, 10, tmp_path / "workspace",
|
||||
)
|
||||
second = planner.plan(
|
||||
"dummy-workload", tmp_path / "input.tsv", input_artifact_id,
|
||||
{"mode": "valid"}, 10, tmp_path / "workspace",
|
||||
)
|
||||
|
||||
assert first.to_json() == second.to_json()
|
||||
payload = json.loads(first.to_json())
|
||||
assert [task["chunk_index"] for task in payload["tasks"]] == [0, 1]
|
||||
assert all(set(task) == {"chunk_index", "input_artifact", "parameters"} for task in payload["tasks"])
|
||||
assert DistributedPlan.from_json(first.to_json()) == first
|
||||
|
||||
|
||||
def test_plan_rejects_unsafe_or_non_deterministic_task_payloads() -> None:
|
||||
with pytest.raises(ValueError, match="unique, ascending"):
|
||||
DistributedPlan(
|
||||
workload="dummy-workload",
|
||||
resolved_parameters={},
|
||||
tasks=(
|
||||
PlannedTask(1, artifact("one"), {}),
|
||||
PlannedTask(0, artifact("zero"), {}),
|
||||
),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="JSON-compatible"):
|
||||
PlannedTask(0, artifact("bad"), {"path": Path("not-serializable")})
|
||||
|
||||
with pytest.raises(ValueError, match="URI or local path"):
|
||||
PlannedTask(0, artifact("uri"), {"input": "file:///tmp/input.tsv"})
|
||||
|
||||
with pytest.raises(ValueError, match="canonical hyphenated"):
|
||||
DistributedPlan("dummy_workload", {}, (PlannedTask(0, artifact("one"), {}),))
|
||||
|
||||
|
||||
def test_reducer_receives_completed_partials_in_chunk_order(tmp_path: Path) -> None:
|
||||
planner, workload = service()
|
||||
result = planner.reduce(
|
||||
"dummy-workload",
|
||||
(
|
||||
CompletedPartial(3, artifact("three", "text/csv"), {"scanned_rows": 10}),
|
||||
CompletedPartial(1, artifact("one", "text/csv"), {"scanned_rows": 10}),
|
||||
),
|
||||
{"mode": "valid"},
|
||||
tmp_path / "workspace",
|
||||
)
|
||||
|
||||
assert [partial.chunk_index for partial in workload.received_partials] == [1, 3]
|
||||
assert result.metrics == {"partial_count": 2}
|
||||
|
||||
|
||||
def test_reducer_rejects_duplicate_chunk_indexes_before_invocation(tmp_path: Path) -> None:
|
||||
planner, workload = service()
|
||||
duplicate = CompletedPartial(0, artifact("partial", "text/csv"), {"scanned_rows": 1})
|
||||
|
||||
with pytest.raises(ValueError, match="unique chunk_index"):
|
||||
planner.reduce("dummy-workload", (duplicate, duplicate), {"mode": "valid"}, tmp_path)
|
||||
|
||||
assert workload.received_partials == ()
|
||||
|
||||
|
||||
def test_artifact_references_never_accept_paths_or_uris() -> None:
|
||||
with pytest.raises(ValueError, match="UUID"):
|
||||
ArtifactReference("file:///tmp/input.tsv", "a" * 64, "text/csv")
|
||||
with pytest.raises(ValueError, match="lowercase SHA-256"):
|
||||
ArtifactReference(str(uuid5(NAMESPACE_URL, "input")), "A" * 64, "text/csv")
|
||||
|
||||
|
||||
def test_registry_descriptions_are_stable_and_duplicate_names_are_rejected() -> None:
|
||||
registry = DistributedWorkloadRegistry()
|
||||
first, second = DummyWorkload(), DummyWorkload()
|
||||
registry.register(first)
|
||||
|
||||
assert registry.descriptions()[0].name == "dummy-workload"
|
||||
with pytest.raises(ValueError, match="already registered"):
|
||||
registry.register(second)
|
||||
@@ -1,184 +0,0 @@
|
||||
"""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",
|
||||
)
|
||||
@@ -36,10 +36,9 @@ from scimesh.sdk import (
|
||||
WorkloadDefinition,
|
||||
WorkloadRegistry,
|
||||
assert_manifest_round_trip,
|
||||
default_sdk_registry,
|
||||
default_sdk_runtime,
|
||||
similarity_search_sdk_adapter,
|
||||
)
|
||||
from scimesh.workloads.library import default_sdk_registry, default_sdk_runtime
|
||||
from scimesh.workloads.search import similarity_search_sdk_definition
|
||||
from scimesh.workloads.similarity_search import search_similar, write_search_results
|
||||
|
||||
|
||||
@@ -60,8 +59,10 @@ def _registered_similarity_search(shard_rows: int = 2):
|
||||
registry = default_sdk_registry(shard_rows=shard_rows)
|
||||
runtime = default_sdk_runtime()
|
||||
descriptions = registry.descriptions()
|
||||
assert len(descriptions) == 1
|
||||
description = descriptions[0]
|
||||
assert len(descriptions) == 3
|
||||
description = next(
|
||||
item for item in descriptions if item.workload.name == "similarity-search"
|
||||
)
|
||||
definition, negotiated = registry.require(
|
||||
description.workload.name,
|
||||
description.workload.version,
|
||||
@@ -128,7 +129,10 @@ def test_local_sdk_executor_matches_similarity_search_reference(tmp_path: Path)
|
||||
reference = search_similar(dataset, query, top_k=3, progress_every=0)
|
||||
write_search_results(reference_path, reference.matches)
|
||||
|
||||
assert artifact_store.materialize(result_artifact).read_bytes() == reference_path.read_bytes()
|
||||
assert (
|
||||
artifact_store.materialize(result_artifact).read_bytes()
|
||||
== reference_path.read_bytes()
|
||||
)
|
||||
assert result.task_key == "reduce/final"
|
||||
assert result.metrics == {"matches_emitted": 3, "partial_count": 3}
|
||||
|
||||
@@ -187,7 +191,9 @@ def test_legacy_adapter_planning_is_deterministic_ordered_and_path_free(
|
||||
shard_ids: list[list[str]] = []
|
||||
for task in first.tasks:
|
||||
artifact = task.inputs["input"].items[0].artifact
|
||||
with artifact_store.materialize(artifact).open(encoding="utf-8", newline="") as source:
|
||||
with artifact_store.materialize(artifact).open(
|
||||
encoding="utf-8", newline=""
|
||||
) as source:
|
||||
shard_ids.append(
|
||||
[row["chembl_id"] for row in csv.DictReader(source, delimiter="\t")]
|
||||
)
|
||||
@@ -250,7 +256,12 @@ def test_local_store_rejects_malformed_content_before_publishing(
|
||||
|
||||
with pytest.raises(ValueError, match="not a valid bounded document"):
|
||||
store.import_file(malformed, declaration=declaration)
|
||||
assert tuple(path for path in store.root.iterdir() if not path.name.startswith(".seal-")) == ()
|
||||
assert (
|
||||
tuple(
|
||||
path for path in store.root.iterdir() if not path.name.startswith(".seal-")
|
||||
)
|
||||
== ()
|
||||
)
|
||||
|
||||
|
||||
def test_delimited_validator_rejects_headerless_data_and_enforces_record_limit(
|
||||
@@ -334,7 +345,11 @@ def test_local_executor_rejects_handler_forged_outputs(
|
||||
dataset = tmp_path / "molecules.tsv"
|
||||
_write_tiny_dataset(dataset)
|
||||
_, runtime, description, original, _ = _registered_similarity_search()
|
||||
map_stage = next(stage for stage in original.manifest.workflow.stages if stage.kind is StageKind.MAP)
|
||||
map_stage = next(
|
||||
stage
|
||||
for stage in original.manifest.workflow.stages
|
||||
if stage.kind is StageKind.MAP
|
||||
)
|
||||
inner = original.runners[map_stage.entry_point]
|
||||
|
||||
class ForgingRunner:
|
||||
@@ -372,7 +387,9 @@ def test_local_executor_rejects_handler_forged_outputs(
|
||||
)
|
||||
|
||||
|
||||
def test_local_executor_rejects_profiles_that_claim_network_isolation(tmp_path: Path) -> None:
|
||||
def test_local_executor_rejects_profiles_that_claim_network_isolation(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
dataset = tmp_path / "molecules.tsv"
|
||||
_write_tiny_dataset(dataset)
|
||||
_, runtime, description, original, _ = _registered_similarity_search()
|
||||
@@ -409,7 +426,9 @@ def test_local_executor_rejects_aliased_terminal_outputs_before_planning(
|
||||
_write_tiny_dataset(dataset)
|
||||
_, runtime, description, original, _ = _registered_similarity_search()
|
||||
reducer = next(
|
||||
stage for stage in original.manifest.workflow.stages if stage.kind is StageKind.REDUCE
|
||||
stage
|
||||
for stage in original.manifest.workflow.stages
|
||||
if stage.kind is StageKind.REDUCE
|
||||
)
|
||||
internal_name = next(iter(reducer.outputs))
|
||||
workflow = replace(
|
||||
@@ -515,8 +534,7 @@ def _advanced_execution_manifest(
|
||||
elif case == "retries":
|
||||
features = ("retries",)
|
||||
changed = tuple(
|
||||
replace(stage, retry=RetryPolicy(max_attempts=2))
|
||||
for stage in stages
|
||||
replace(stage, retry=RetryPolicy(max_attempts=2)) for stage in stages
|
||||
)
|
||||
elif case == "secrets":
|
||||
features = ("secret-injection",)
|
||||
@@ -604,7 +622,9 @@ def test_local_executor_rejects_profiles_it_cannot_enforce(
|
||||
)
|
||||
|
||||
|
||||
def test_local_executor_fails_when_the_declared_verifier_rejects(tmp_path: Path) -> None:
|
||||
def test_local_executor_fails_when_the_declared_verifier_rejects(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
dataset = tmp_path / "molecules.tsv"
|
||||
_write_tiny_dataset(dataset)
|
||||
_, runtime, _, original, _ = _registered_similarity_search()
|
||||
@@ -639,19 +659,21 @@ def test_local_executor_fails_when_the_declared_verifier_rejects(tmp_path: Path)
|
||||
)
|
||||
|
||||
|
||||
def test_local_executor_enforces_the_declared_output_byte_budget(tmp_path: Path) -> None:
|
||||
def test_local_executor_enforces_the_declared_output_byte_budget(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
dataset = tmp_path / "molecules.tsv"
|
||||
_write_tiny_dataset(dataset)
|
||||
runtime = default_sdk_runtime()
|
||||
adapter = similarity_search_sdk_adapter(shard_rows=2)
|
||||
workload = similarity_search_sdk_definition(shard_rows=2)
|
||||
# The planner pins its own manifest into every task, so the budget cut must
|
||||
# be applied to the adapter's manifest for plan and definition to agree.
|
||||
adapter.manifest = replace(
|
||||
adapter.manifest,
|
||||
workflow=replace(adapter.manifest.workflow, max_output_bytes=256),
|
||||
limits=replace(adapter.manifest.limits, max_output_bytes=256),
|
||||
# be applied to the workload's manifest for plan and definition to agree.
|
||||
workload.manifest = replace(
|
||||
workload.manifest,
|
||||
workflow=replace(workload.manifest.workflow, max_output_bytes=256),
|
||||
limits=replace(workload.manifest.limits, max_output_bytes=256),
|
||||
)
|
||||
definition = adapter.definition()
|
||||
definition = workload.definition()
|
||||
registry = WorkloadRegistry()
|
||||
registry.register(definition, enabled=True)
|
||||
store = LocalArtifactStore(tmp_path / "artifacts")
|
||||
|
||||
@@ -25,13 +25,13 @@ from scimesh.sdk import (
|
||||
VerifyContext,
|
||||
WorkloadRegistry,
|
||||
assert_manifest_round_trip,
|
||||
default_sdk_runtime,
|
||||
)
|
||||
from scimesh.sdk.descriptors import (
|
||||
from scimesh.workloads.descriptors import (
|
||||
DESCRIPTOR_COLUMNS,
|
||||
descriptor_batch_sdk_definition,
|
||||
compute_descriptor_batch,
|
||||
)
|
||||
from scimesh.workloads.library import default_sdk_runtime
|
||||
|
||||
|
||||
def _write_tiny_dataset(path: Path) -> None:
|
||||
@@ -428,7 +428,7 @@ def test_descriptor_batch_discovery_imports_an_allowlisted_installed_entry_point
|
||||
class EntryPoint:
|
||||
name = "descriptor-batch@1.0.0"
|
||||
dist = metadata.distribution("scimesh")
|
||||
value = "scimesh.sdk.descriptors.definition:workload_definition"
|
||||
value = "scimesh.workloads.descriptors:workload_definition"
|
||||
|
||||
@property
|
||||
def module(self) -> str:
|
||||
|
||||
@@ -0,0 +1,295 @@
|
||||
"""Tests for the SDK-built similarity-graph workload."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from scimesh.sdk import (
|
||||
ArtifactCollection,
|
||||
DeterminismProfile,
|
||||
JobRequest,
|
||||
LocalArtifactStore,
|
||||
LocalCoreBatchExecutor,
|
||||
LocalPlanningContext,
|
||||
StageKind,
|
||||
assert_manifest_round_trip,
|
||||
)
|
||||
from scimesh.workloads.graph import (
|
||||
check_pair_coverage,
|
||||
merge_edge_partials,
|
||||
similarity_graph_sdk_definition,
|
||||
)
|
||||
from scimesh.workloads.library import default_sdk_registry, default_sdk_runtime
|
||||
from scimesh.workloads.similarity_graph import (
|
||||
build_similarity_graph,
|
||||
write_graph_edges,
|
||||
)
|
||||
|
||||
|
||||
def _write_tiny_dataset(path: Path) -> None:
|
||||
path.write_text(
|
||||
"chembl_id\tcanonical_smiles\n"
|
||||
"A\tCCO\n"
|
||||
"B\tCCCC\n"
|
||||
"C\tCCN\n"
|
||||
"D\tCCCCCC\n"
|
||||
"E\tnot-a-smiles\n"
|
||||
"F\tCCOCC\n"
|
||||
"G\tc1ccccc1\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
|
||||
def _registered_similarity_graph():
|
||||
registry = default_sdk_registry()
|
||||
runtime = default_sdk_runtime()
|
||||
description = next(
|
||||
item
|
||||
for item in registry.descriptions()
|
||||
if item.workload.name == "similarity-graph"
|
||||
)
|
||||
definition, negotiated = registry.require(
|
||||
description.workload.name,
|
||||
description.workload.version,
|
||||
description.package_digest,
|
||||
runtime=runtime,
|
||||
)
|
||||
return registry, runtime, description, definition, negotiated
|
||||
|
||||
|
||||
def _request_for(
|
||||
dataset: Path,
|
||||
artifact_store: LocalArtifactStore,
|
||||
definition,
|
||||
*,
|
||||
threshold: float = 0.3,
|
||||
threshold_direction: str = "greater",
|
||||
block_size: int = 2,
|
||||
) -> JobRequest:
|
||||
input_port = definition.manifest.inputs["input"]
|
||||
dataset_artifact = artifact_store.import_file(
|
||||
dataset,
|
||||
declaration=input_port.schema,
|
||||
)
|
||||
return JobRequest(
|
||||
workload=definition.manifest.workload,
|
||||
parameters={
|
||||
"threshold": threshold,
|
||||
"threshold_direction": threshold_direction,
|
||||
"block_size": block_size,
|
||||
},
|
||||
inputs={"input": ArtifactCollection.single(dataset_artifact)},
|
||||
)
|
||||
|
||||
|
||||
def test_similarity_graph_manifest_is_registered_and_negotiable() -> None:
|
||||
_, runtime, description, definition, negotiated = _registered_similarity_graph()
|
||||
manifest = definition.manifest
|
||||
|
||||
assert description.enabled is True
|
||||
assert manifest.workload.name == "similarity-graph"
|
||||
assert manifest.workload.version == "1.0.0"
|
||||
assert manifest.determinism is DeterminismProfile.BYTE_EXACT
|
||||
assert manifest.verifier.verifier.canonical == "exact-artifact@1"
|
||||
assert set(mode.value for mode in manifest.trust_modes) == {
|
||||
"trusted",
|
||||
"untrusted_quorum",
|
||||
}
|
||||
assert [stage.kind for stage in manifest.workflow.stages] == [
|
||||
StageKind.MAP,
|
||||
StageKind.REDUCE,
|
||||
]
|
||||
assert set(definition.runners) == {manifest.workflow.stages[0].entry_point}
|
||||
assert set(definition.reducers) == {manifest.workflow.stages[1].entry_point}
|
||||
assert negotiated is not None
|
||||
assert_manifest_round_trip(manifest)
|
||||
assert runtime is not None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("threshold_direction", ("greater", "less"))
|
||||
def test_local_sdk_executor_matches_similarity_graph_reference(
|
||||
tmp_path: Path,
|
||||
threshold_direction: str,
|
||||
) -> None:
|
||||
dataset = tmp_path / "molecules.tsv"
|
||||
_write_tiny_dataset(dataset)
|
||||
registry, runtime, description, definition, _ = _registered_similarity_graph()
|
||||
artifact_store = LocalArtifactStore(tmp_path / "artifacts")
|
||||
request = _request_for(
|
||||
dataset,
|
||||
artifact_store,
|
||||
definition,
|
||||
threshold_direction=threshold_direction,
|
||||
)
|
||||
|
||||
result = LocalCoreBatchExecutor(
|
||||
registry,
|
||||
runtime,
|
||||
artifact_store,
|
||||
tmp_path / "sdk-work",
|
||||
).execute(request, description.package_digest)
|
||||
result_artifact = result.outputs["result"].items[0].artifact
|
||||
|
||||
reference_path = tmp_path / "reference.csv"
|
||||
reference = build_similarity_graph(
|
||||
dataset,
|
||||
threshold=0.3,
|
||||
block_size=1_000,
|
||||
threshold_direction=threshold_direction,
|
||||
)
|
||||
write_graph_edges(reference_path, reference.edges)
|
||||
|
||||
assert (
|
||||
artifact_store.materialize(result_artifact).read_bytes()
|
||||
== reference_path.read_bytes()
|
||||
)
|
||||
assert result.task_key == "reduce/final"
|
||||
assert result.metrics["partial_count"] == 6
|
||||
assert result.metrics["edges_emitted"] == len(reference.edges)
|
||||
|
||||
|
||||
def test_similarity_graph_result_is_invariant_to_block_size(tmp_path: Path) -> None:
|
||||
dataset = tmp_path / "molecules.tsv"
|
||||
_write_tiny_dataset(dataset)
|
||||
registry, runtime, _, definition, _ = _registered_similarity_graph()
|
||||
artifact_store = LocalArtifactStore(tmp_path / "artifacts")
|
||||
|
||||
outputs = []
|
||||
for block_size in (2, 3):
|
||||
request = _request_for(
|
||||
dataset,
|
||||
artifact_store,
|
||||
definition,
|
||||
block_size=block_size,
|
||||
)
|
||||
result = LocalCoreBatchExecutor(
|
||||
registry,
|
||||
runtime,
|
||||
artifact_store,
|
||||
tmp_path / f"sdk-work-{block_size}",
|
||||
).execute(request, definition.manifest.package.digest)
|
||||
artifact = result.outputs["result"].items[0].artifact
|
||||
outputs.append(artifact_store.materialize(artifact).read_bytes())
|
||||
assert result.metrics["partial_count"] == {2: 6, 3: 3}[block_size]
|
||||
|
||||
assert outputs[0] == outputs[1]
|
||||
|
||||
|
||||
def test_similarity_graph_planning_covers_each_block_pair_once(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
dataset = tmp_path / "molecules.tsv"
|
||||
_write_tiny_dataset(dataset)
|
||||
registry, runtime, description, definition, _ = _registered_similarity_graph()
|
||||
artifact_store = LocalArtifactStore(tmp_path / "artifacts")
|
||||
request = _request_for(dataset, artifact_store, definition)
|
||||
input_artifact = request.inputs["input"].items[0].artifact
|
||||
|
||||
plan = registry.plan(
|
||||
request,
|
||||
description.package_digest,
|
||||
runtime,
|
||||
LocalPlanningContext(
|
||||
artifact_store,
|
||||
artifact_store,
|
||||
tmp_path / "plan",
|
||||
allowed_artifacts=(input_artifact,),
|
||||
),
|
||||
)
|
||||
|
||||
assert [task.task_key for task in plan.tasks] == [
|
||||
"map/0000x0000",
|
||||
"map/0000x0001",
|
||||
"map/0000x0002",
|
||||
"map/0001x0001",
|
||||
"map/0001x0002",
|
||||
"map/0002x0002",
|
||||
]
|
||||
for task in plan.tasks:
|
||||
assert set(task.inputs) == {"left", "right"}
|
||||
assert task.inputs["left"].items[0].artifact is not None
|
||||
pairs = {
|
||||
(int(left), int(right))
|
||||
for task in plan.tasks
|
||||
for left, right in (task.task_key[len("map/") :].split("x"),)
|
||||
}
|
||||
check_pair_coverage(tuple(sorted(pairs)))
|
||||
|
||||
wire_payload = plan.to_json()
|
||||
assert str(tmp_path) not in wire_payload
|
||||
assert "file://" not in wire_payload
|
||||
assert "worker://" not in wire_payload
|
||||
assert "workspace" not in wire_payload
|
||||
|
||||
|
||||
def test_similarity_graph_rejects_duplicate_molecule_ids(tmp_path: Path) -> None:
|
||||
dataset = tmp_path / "duplicates.tsv"
|
||||
dataset.write_text(
|
||||
"chembl_id\tcanonical_smiles\nA\tCCO\nB\tCCCC\nA\tCCN\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
registry, runtime, _, definition, _ = _registered_similarity_graph()
|
||||
artifact_store = LocalArtifactStore(tmp_path / "artifacts")
|
||||
request = _request_for(dataset, artifact_store, definition)
|
||||
|
||||
with pytest.raises(ValueError, match="duplicate chembl_id"):
|
||||
LocalCoreBatchExecutor(
|
||||
registry,
|
||||
runtime,
|
||||
artifact_store,
|
||||
tmp_path / "work",
|
||||
).execute(request, definition.manifest.package.digest)
|
||||
|
||||
|
||||
def test_similarity_graph_reducer_rejects_duplicate_unordered_pairs(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
first = tmp_path / "first.csv"
|
||||
first.write_text(
|
||||
"source_id,target_id,similarity\nA,B,0.500000\nC,D,0.100000\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
second = tmp_path / "second.csv"
|
||||
second.write_text(
|
||||
"source_id,target_id,similarity\nB,A,0.500000\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="duplicate unordered pair"):
|
||||
merge_edge_partials((first, second), tmp_path / "result.csv")
|
||||
|
||||
|
||||
def test_similarity_graph_pair_coverage_rejects_missing_block_pair() -> None:
|
||||
with pytest.raises(ValueError, match="do not cover the full block pair set"):
|
||||
check_pair_coverage(((0, 0), (0, 1)))
|
||||
|
||||
with pytest.raises(ValueError, match="duplicate block pair"):
|
||||
check_pair_coverage(((0, 0), (0, 0), (0, 1), (1, 1)))
|
||||
|
||||
|
||||
def test_similarity_graph_merge_is_deterministically_sorted(tmp_path: Path) -> None:
|
||||
first = tmp_path / "first.csv"
|
||||
first.write_text(
|
||||
"source_id,target_id,similarity\nC,A,0.200000\nB,C,0.400000\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
second = tmp_path / "second.csv"
|
||||
second.write_text(
|
||||
"source_id,target_id,similarity\nA,B,0.900000\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
result_path = tmp_path / "result.csv"
|
||||
metrics = merge_edge_partials((first, second), result_path)
|
||||
|
||||
assert metrics == {"partial_count": 2, "edges_emitted": 3}
|
||||
rows = list(csv.DictReader(result_path.open(encoding="utf-8", newline="")))
|
||||
assert [
|
||||
(row["source_id"], row["target_id"], row["similarity"]) for row in rows
|
||||
] == [
|
||||
("A", "B", "0.900000"),
|
||||
("B", "C", "0.400000"),
|
||||
("C", "A", "0.200000"),
|
||||
]
|
||||
+32
-20
@@ -25,20 +25,22 @@ from scimesh.sdk import (
|
||||
WorkloadDefinition,
|
||||
WorkloadId,
|
||||
WorkloadRegistry,
|
||||
current_scimesh_package_digest,
|
||||
default_sdk_runtime,
|
||||
installed_distribution_digest,
|
||||
similarity_search_sdk_adapter,
|
||||
)
|
||||
from scimesh.sdk.schema import (
|
||||
ParameterValidationError,
|
||||
validate_parameter_instance,
|
||||
validate_schema_definition,
|
||||
)
|
||||
from scimesh.workloads.environment import current_scimesh_package_digest
|
||||
from scimesh.workloads.library import default_sdk_runtime
|
||||
from scimesh.workloads.search import similarity_search_sdk_definition
|
||||
|
||||
|
||||
def _definition(*, version: str = "1.0.0", digest_character: str = "a") -> WorkloadDefinition:
|
||||
original = similarity_search_sdk_adapter(shard_rows=2).definition()
|
||||
def _definition(
|
||||
*, version: str = "1.0.0", digest_character: str = "a"
|
||||
) -> WorkloadDefinition:
|
||||
original = similarity_search_sdk_definition(shard_rows=2).definition()
|
||||
manifest = replace(
|
||||
original.manifest,
|
||||
workload=WorkloadId("similarity-search", version),
|
||||
@@ -72,13 +74,16 @@ def test_registry_requires_an_explicit_enabled_version_and_digest() -> None:
|
||||
registry.register(first)
|
||||
|
||||
registry.enable("similarity-search", "2.0.0", "sha256:" + "b" * 64)
|
||||
assert [item.workload.version for item in registry.descriptions()] == ["1.0.0", "2.0.0"]
|
||||
assert [item.workload.version for item in registry.descriptions()] == [
|
||||
"1.0.0",
|
||||
"2.0.0",
|
||||
]
|
||||
|
||||
|
||||
def test_compatibility_failure_occurs_before_planner_invocation(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
original = similarity_search_sdk_adapter(shard_rows=2).definition()
|
||||
original = similarity_search_sdk_definition(shard_rows=2).definition()
|
||||
|
||||
class CountingPlanner:
|
||||
calls = 0
|
||||
@@ -140,7 +145,7 @@ def test_job_selected_features_and_trust_mode_fail_closed_before_planning(
|
||||
request_changes: dict[str, object],
|
||||
error_code: str,
|
||||
) -> None:
|
||||
definition = similarity_search_sdk_adapter(shard_rows=2).definition()
|
||||
definition = similarity_search_sdk_definition(shard_rows=2).definition()
|
||||
registry = WorkloadRegistry()
|
||||
registry.register(definition, enabled=True)
|
||||
input_port = definition.manifest.inputs["input"]
|
||||
@@ -180,7 +185,7 @@ class _EntryPoints(tuple):
|
||||
def test_discovery_imports_only_an_exact_allowlisted_installed_entry_point(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
definition = similarity_search_sdk_adapter().definition()
|
||||
definition = similarity_search_sdk_definition().definition()
|
||||
loaded: list[str] = []
|
||||
|
||||
class EntryPoint:
|
||||
@@ -191,7 +196,7 @@ def test_discovery_imports_only_an_exact_allowlisted_installed_entry_point(
|
||||
if distribution == "scimesh"
|
||||
else SimpleNamespace(name=distribution)
|
||||
)
|
||||
self.value = "scimesh.sdk.builtins:similarity_search_sdk_adapter"
|
||||
self.value = "scimesh.workloads.search:similarity_search_sdk_definition"
|
||||
|
||||
@property
|
||||
def module(self) -> str:
|
||||
@@ -232,14 +237,14 @@ def test_discovery_imports_only_an_exact_allowlisted_installed_entry_point(
|
||||
def test_discovery_measures_package_before_importing_entry_point(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
definition = similarity_search_sdk_adapter().definition()
|
||||
definition = similarity_search_sdk_definition().definition()
|
||||
loaded = False
|
||||
|
||||
class EntryPoint:
|
||||
name = "similarity-search@1.0.0"
|
||||
dist = metadata.distribution("scimesh")
|
||||
value = "scimesh.sdk.builtins:similarity_search_sdk_adapter"
|
||||
module = "scimesh.sdk.builtins"
|
||||
value = "scimesh.workloads.search:similarity_search_sdk_definition"
|
||||
module = "scimesh.workloads.search"
|
||||
|
||||
def load(self):
|
||||
nonlocal loaded
|
||||
@@ -333,7 +338,7 @@ def test_discovery_rejects_entry_point_module_owned_by_another_distribution(
|
||||
"scimesh.sdk.registry.metadata.entry_points",
|
||||
lambda: _EntryPoints((EntryPoint(),)),
|
||||
)
|
||||
definition = similarity_search_sdk_adapter().definition()
|
||||
definition = similarity_search_sdk_definition().definition()
|
||||
with pytest.raises(ValueError, match="outside its distribution"):
|
||||
WorkloadRegistry().discover_installed(
|
||||
(
|
||||
@@ -463,14 +468,14 @@ def test_disabled_workload_is_not_resolvable_until_re_enabled() -> None:
|
||||
def test_discovery_rechecks_the_package_digest_after_loading(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
definition = similarity_search_sdk_adapter().definition()
|
||||
definition = similarity_search_sdk_definition().definition()
|
||||
digests = iter((definition.manifest.package.digest, "sha256:" + "e" * 64))
|
||||
|
||||
class EntryPoint:
|
||||
name = "similarity-search@1.0.0"
|
||||
dist = metadata.distribution("scimesh")
|
||||
value = "scimesh.sdk.builtins:similarity_search_sdk_adapter"
|
||||
module = "scimesh.sdk.builtins"
|
||||
value = "scimesh.workloads.search:similarity_search_sdk_definition"
|
||||
module = "scimesh.workloads.search"
|
||||
|
||||
def load(self):
|
||||
return lambda: definition
|
||||
@@ -498,10 +503,17 @@ def test_discovery_rechecks_the_package_digest_after_loading(
|
||||
assert registry.descriptions() == ()
|
||||
|
||||
|
||||
def test_request_trust_mode_must_be_enforceable_by_runtime_and_stages(tmp_path: Path) -> None:
|
||||
original = similarity_search_sdk_adapter(shard_rows=2).definition()
|
||||
def test_request_trust_mode_must_be_enforceable_by_runtime_and_stages(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
original = similarity_search_sdk_definition(shard_rows=2).definition()
|
||||
stages = tuple(
|
||||
replace(stage, trust_modes=("trusted",))
|
||||
for stage in original.manifest.workflow.stages
|
||||
)
|
||||
manifest = replace(
|
||||
original.manifest,
|
||||
workflow=replace(original.manifest.workflow, stages=stages),
|
||||
trust_modes=(TrustMode.TRUSTED, TrustMode.VERIFIED),
|
||||
)
|
||||
definition = WorkloadDefinition(
|
||||
@@ -554,7 +566,7 @@ def test_request_trust_mode_must_be_enforceable_by_runtime_and_stages(tmp_path:
|
||||
|
||||
|
||||
def test_job_cannot_require_a_feature_outside_the_runtime(tmp_path: Path) -> None:
|
||||
original = similarity_search_sdk_adapter(shard_rows=2).definition()
|
||||
original = similarity_search_sdk_definition(shard_rows=2).definition()
|
||||
manifest = replace(
|
||||
original.manifest,
|
||||
optional_features=(
|
||||
|
||||
@@ -0,0 +1,252 @@
|
||||
"""Tests for the SDK-built similarity-search workload."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from scimesh.sdk import (
|
||||
ArtifactCollection,
|
||||
DeterminismProfile,
|
||||
JobRequest,
|
||||
LocalArtifactStore,
|
||||
LocalCoreBatchExecutor,
|
||||
LocalPlanningContext,
|
||||
StageKind,
|
||||
WorkloadRegistry,
|
||||
assert_manifest_round_trip,
|
||||
)
|
||||
from scimesh.workloads.library import default_sdk_registry, default_sdk_runtime
|
||||
from scimesh.workloads.search import similarity_search_sdk_definition
|
||||
from scimesh.workloads.similarity_search import (
|
||||
find_molecule_by_id,
|
||||
search_similar,
|
||||
write_search_results,
|
||||
)
|
||||
|
||||
|
||||
def _write_tiny_dataset(path: Path) -> None:
|
||||
path.write_text(
|
||||
"chembl_id\tcanonical_smiles\textra\n"
|
||||
"QUERY\tCCO\tquery\n"
|
||||
"ALCOHOL\tCCCO\talcohol\n"
|
||||
"ALKANE\tCCCC\talkane\n"
|
||||
"BROKEN\tnot-a-smiles\tinvalid\n"
|
||||
"DUPLICATE\tCCO\tduplicate\n"
|
||||
"AMINE\tCCN\tamine\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
|
||||
def _registered_similarity_search(shard_rows: int = 2):
|
||||
registry = default_sdk_registry(shard_rows=shard_rows)
|
||||
runtime = default_sdk_runtime()
|
||||
description = next(
|
||||
item
|
||||
for item in registry.descriptions()
|
||||
if item.workload.name == "similarity-search"
|
||||
)
|
||||
definition, negotiated = registry.require(
|
||||
description.workload.name,
|
||||
description.workload.version,
|
||||
description.package_digest,
|
||||
runtime=runtime,
|
||||
)
|
||||
return registry, runtime, description, definition, negotiated
|
||||
|
||||
|
||||
def _request_for(
|
||||
dataset: Path,
|
||||
artifact_store: LocalArtifactStore,
|
||||
definition,
|
||||
) -> JobRequest:
|
||||
input_port = definition.manifest.inputs["input"]
|
||||
dataset_artifact = artifact_store.import_file(
|
||||
dataset,
|
||||
declaration=input_port.schema,
|
||||
)
|
||||
return JobRequest(
|
||||
workload=definition.manifest.workload,
|
||||
parameters={"query_id": "QUERY", "top_k": 3, "progress_every": 0},
|
||||
inputs={"input": ArtifactCollection.single(dataset_artifact)},
|
||||
)
|
||||
|
||||
|
||||
def test_similarity_search_manifest_is_registered_and_negotiable() -> None:
|
||||
_, runtime, description, definition, negotiated = _registered_similarity_search()
|
||||
manifest = definition.manifest
|
||||
|
||||
assert description.enabled is True
|
||||
assert manifest.workload.name == "similarity-search"
|
||||
assert manifest.workload.version == "1.0.0"
|
||||
assert manifest.determinism is DeterminismProfile.BYTE_EXACT
|
||||
assert manifest.verifier.verifier.canonical == "exact-artifact@1"
|
||||
assert set(mode.value for mode in manifest.trust_modes) == {
|
||||
"trusted",
|
||||
"untrusted_quorum",
|
||||
}
|
||||
assert [stage.kind for stage in manifest.workflow.stages] == [
|
||||
StageKind.MAP,
|
||||
StageKind.REDUCE,
|
||||
]
|
||||
assert set(definition.runners) == {manifest.workflow.stages[0].entry_point}
|
||||
assert set(definition.reducers) == {manifest.workflow.stages[1].entry_point}
|
||||
assert negotiated is not None
|
||||
assert negotiated.manifest == manifest
|
||||
assert_manifest_round_trip(manifest)
|
||||
assert runtime is not None
|
||||
|
||||
|
||||
def test_local_sdk_executor_matches_similarity_search_reference(tmp_path: Path) -> None:
|
||||
dataset = tmp_path / "molecules.tsv"
|
||||
_write_tiny_dataset(dataset)
|
||||
registry, runtime, description, definition, _ = _registered_similarity_search()
|
||||
artifact_store = LocalArtifactStore(tmp_path / "artifacts")
|
||||
request = _request_for(dataset, artifact_store, definition)
|
||||
|
||||
result = LocalCoreBatchExecutor(
|
||||
registry,
|
||||
runtime,
|
||||
artifact_store,
|
||||
tmp_path / "sdk-work",
|
||||
).execute(request, description.package_digest)
|
||||
result_artifact = result.outputs["result"].items[0].artifact
|
||||
|
||||
reference_path = tmp_path / "reference.csv"
|
||||
query = find_molecule_by_id(dataset, "QUERY")
|
||||
reference = search_similar(dataset, query, top_k=3, progress_every=0)
|
||||
write_search_results(reference_path, reference.matches)
|
||||
|
||||
assert (
|
||||
artifact_store.materialize(result_artifact).read_bytes()
|
||||
== reference_path.read_bytes()
|
||||
)
|
||||
assert result.task_key == "reduce/final"
|
||||
assert dict(result.metrics) == {"matches_emitted": 3, "partial_count": 3}
|
||||
|
||||
|
||||
def test_similarity_search_planning_is_deterministic_ordered_and_path_free(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
dataset = tmp_path / "molecules.tsv"
|
||||
_write_tiny_dataset(dataset)
|
||||
registry, runtime, description, definition, _ = _registered_similarity_search()
|
||||
artifact_store = LocalArtifactStore(tmp_path / "artifacts")
|
||||
request = _request_for(dataset, artifact_store, definition)
|
||||
input_artifact = request.inputs["input"].items[0].artifact
|
||||
|
||||
first = registry.plan(
|
||||
request,
|
||||
description.package_digest,
|
||||
runtime,
|
||||
LocalPlanningContext(
|
||||
artifact_store,
|
||||
artifact_store,
|
||||
tmp_path / "first-plan",
|
||||
allowed_artifacts=(input_artifact,),
|
||||
),
|
||||
)
|
||||
second = registry.plan(
|
||||
request,
|
||||
description.package_digest,
|
||||
runtime,
|
||||
LocalPlanningContext(
|
||||
artifact_store,
|
||||
artifact_store,
|
||||
tmp_path / "second-plan",
|
||||
allowed_artifacts=(input_artifact,),
|
||||
),
|
||||
)
|
||||
|
||||
assert first.to_json() == second.to_json()
|
||||
assert first.digest == second.digest
|
||||
assert first.package_digest == definition.manifest.package.digest
|
||||
assert first.manifest_digest == definition.manifest.digest
|
||||
assert [task.task_key for task in first.tasks] == [
|
||||
"map/00000000",
|
||||
"map/00000001",
|
||||
"map/00000002",
|
||||
]
|
||||
assert all(task.stage_id == "map" for task in first.tasks)
|
||||
assert all("query_id" not in task.parameters for task in first.tasks)
|
||||
assert all(task.parameters["query_smiles"] == "CCO" for task in first.tasks)
|
||||
assert all(task.parameters["top_k"] == 3 for task in first.tasks)
|
||||
assert all(task.package_digest == first.package_digest for task in first.tasks)
|
||||
assert all(task.manifest_digest == first.manifest_digest for task in first.tasks)
|
||||
assert first.resolved_parameters["query_source"] == {
|
||||
"kind": "chembl_id",
|
||||
"value": "QUERY",
|
||||
}
|
||||
|
||||
shard_ids: list[list[str]] = []
|
||||
for task in first.tasks:
|
||||
artifact = task.inputs["input"].items[0].artifact
|
||||
with artifact_store.materialize(artifact).open(
|
||||
encoding="utf-8", newline=""
|
||||
) as source:
|
||||
shard_ids.append(
|
||||
[row["chembl_id"] for row in csv.DictReader(source, delimiter="\t")]
|
||||
)
|
||||
assert shard_ids == [
|
||||
["QUERY", "ALCOHOL"],
|
||||
["ALKANE", "BROKEN"],
|
||||
["DUPLICATE", "AMINE"],
|
||||
]
|
||||
|
||||
wire_payload = first.to_json()
|
||||
assert str(tmp_path) not in wire_payload
|
||||
assert "file://" not in wire_payload
|
||||
assert "worker://" not in wire_payload
|
||||
assert "workspace" not in wire_payload
|
||||
|
||||
|
||||
def test_similarity_search_rejects_ambiguous_or_mistyped_parameters(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
dataset = tmp_path / "molecules.tsv"
|
||||
_write_tiny_dataset(dataset)
|
||||
registry, runtime, _, definition, _ = _registered_similarity_search()
|
||||
artifact_store = LocalArtifactStore(tmp_path / "artifacts")
|
||||
input_artifact = artifact_store.import_file(
|
||||
dataset,
|
||||
declaration=definition.manifest.inputs["input"].schema,
|
||||
)
|
||||
base = JobRequest(
|
||||
workload=definition.manifest.workload,
|
||||
parameters={"query_id": "QUERY", "top_k": 3},
|
||||
inputs={"input": ArtifactCollection.single(input_artifact)},
|
||||
)
|
||||
|
||||
for bad_parameters, message in (
|
||||
({"query_id": "QUERY", "query_smiles": "CCO"}, "oneOf did not match"),
|
||||
({"query_id": "QUERY", "top_k": 0}, "violates minimum"),
|
||||
):
|
||||
request = JobRequest(
|
||||
workload=base.workload,
|
||||
parameters=bad_parameters,
|
||||
inputs=base.inputs,
|
||||
)
|
||||
with pytest.raises(ValueError, match=message):
|
||||
registry.plan(
|
||||
request,
|
||||
definition.manifest.package.digest,
|
||||
runtime,
|
||||
LocalPlanningContext(
|
||||
artifact_store,
|
||||
artifact_store,
|
||||
tmp_path / "bad-plan",
|
||||
allowed_artifacts=(input_artifact,),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_similarity_search_workload_definition_is_discoverable() -> None:
|
||||
from scimesh.sdk import WorkloadDefinition
|
||||
from scimesh.workloads.search import workload_definition
|
||||
|
||||
definition = workload_definition()
|
||||
assert isinstance(definition, WorkloadDefinition)
|
||||
assert definition.manifest.workload.name == "similarity-search"
|
||||
assert definition.manifest.workload.version == "1.0.0"
|
||||
+229
-61
@@ -23,7 +23,11 @@ from scimesh.worker.models import (
|
||||
RunResult,
|
||||
UploadedArtifact,
|
||||
)
|
||||
from scimesh.worker.artifacts import HttpArtifactClient, _SameOriginAuthRedirectHandler, _origin
|
||||
from scimesh.worker.artifacts import (
|
||||
HttpArtifactClient,
|
||||
_SameOriginAuthRedirectHandler,
|
||||
_origin,
|
||||
)
|
||||
from scimesh.worker.runners import SciMeshRunner
|
||||
from scimesh.worker.transport import NoRedirectHandler
|
||||
|
||||
@@ -32,12 +36,18 @@ class FakeCoordinator:
|
||||
def __init__(self, task: ClaimedTask | None) -> None:
|
||||
self.task, self.submissions, self.failures, self.heartbeats = task, [], [], []
|
||||
|
||||
def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None:
|
||||
def claim(
|
||||
self, worker_id: str, capabilities: tuple[str, ...]
|
||||
) -> ClaimedTask | None:
|
||||
task, self.task = self.task, None
|
||||
return task
|
||||
|
||||
def register(
|
||||
self, name: str, capabilities: tuple[str, ...], cpu_count: int, memory_mb: int | None
|
||||
self,
|
||||
name: str,
|
||||
capabilities: tuple[str, ...],
|
||||
cpu_count: int,
|
||||
memory_mb: int | None,
|
||||
) -> RegisteredWorker:
|
||||
return RegisteredWorker("11111111-1111-4111-8111-111111111111", 15)
|
||||
|
||||
@@ -71,6 +81,7 @@ class FakeArtifacts:
|
||||
len(content),
|
||||
)
|
||||
|
||||
|
||||
class FakeRunner:
|
||||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
@@ -84,18 +95,40 @@ class FakeRunner:
|
||||
|
||||
def make_task(content: bytes, checksum: str | None = None) -> ClaimedTask:
|
||||
lease = (datetime.now(timezone.utc) + timedelta(seconds=60)).isoformat()
|
||||
return ClaimedTask("task-1", 1, lease, "similarity-search", InputArtifact("https://example.test/input", checksum or hashlib.sha256(content).hexdigest()), {"query_id": "CHEMBL1"})
|
||||
return ClaimedTask(
|
||||
"task-1",
|
||||
1,
|
||||
lease,
|
||||
"similarity-search",
|
||||
InputArtifact(
|
||||
"https://example.test/input",
|
||||
checksum or hashlib.sha256(content).hexdigest(),
|
||||
),
|
||||
{"query_id": "CHEMBL1"},
|
||||
)
|
||||
|
||||
|
||||
def daemon(tmp_path: Path, task: ClaimedTask | None, content: bytes):
|
||||
coordinator, artifacts, runner = FakeCoordinator(task), FakeArtifacts(content), FakeRunner()
|
||||
coordinator, artifacts, runner = (
|
||||
FakeCoordinator(task),
|
||||
FakeArtifacts(content),
|
||||
FakeRunner(),
|
||||
)
|
||||
config = WorkerConfig("https://example.test", "worker-1", tmp_path / "work")
|
||||
return WorkerDaemon(config, coordinator, artifacts, runner), coordinator, artifacts, runner, config
|
||||
return (
|
||||
WorkerDaemon(config, coordinator, artifacts, runner),
|
||||
coordinator,
|
||||
artifacts,
|
||||
runner,
|
||||
config,
|
||||
)
|
||||
|
||||
|
||||
def test_claims_runs_uploads_and_submits_csv(tmp_path: Path) -> None:
|
||||
content = b"input fixture"
|
||||
worker, coordinator, artifacts, runner, _ = daemon(tmp_path, make_task(content), content)
|
||||
worker, coordinator, artifacts, runner, _ = daemon(
|
||||
tmp_path, make_task(content), content
|
||||
)
|
||||
assert worker.run_once() == RunOnceOutcome(claimed=True, completed=True)
|
||||
assert runner.calls == 1
|
||||
assert len(artifacts.uploaded) == 1
|
||||
@@ -107,10 +140,16 @@ 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"
|
||||
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,
|
||||
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)
|
||||
@@ -131,11 +170,19 @@ def test_two_workers_complete_resolved_shards_after_one_retry(tmp_path: Path) ->
|
||||
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,
|
||||
"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,
|
||||
"other-task",
|
||||
1,
|
||||
first.lease_expires_at,
|
||||
"similarity-search",
|
||||
first.input,
|
||||
{"query_smiles": "CCO", "top_k": 5},
|
||||
)
|
||||
|
||||
@@ -145,16 +192,26 @@ def test_two_workers_complete_resolved_shards_after_one_retry(tmp_path: Path) ->
|
||||
self.queue = [first, second]
|
||||
self.claimants: list[str] = []
|
||||
|
||||
def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None:
|
||||
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"]:
|
||||
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.task_id,
|
||||
2,
|
||||
task.lease_expires_at,
|
||||
task.workload,
|
||||
task.input,
|
||||
task.parameters,
|
||||
)
|
||||
)
|
||||
@@ -174,11 +231,15 @@ def test_two_workers_complete_resolved_shards_after_one_retry(tmp_path: Path) ->
|
||||
artifacts = FakeArtifacts(content)
|
||||
worker_a = WorkerDaemon(
|
||||
WorkerConfig("https://example.test", "worker-a", tmp_path / "worker-a"),
|
||||
coordinator, artifacts, FailFirstAttempt(),
|
||||
coordinator,
|
||||
artifacts,
|
||||
FailFirstAttempt(),
|
||||
)
|
||||
worker_b = WorkerDaemon(
|
||||
WorkerConfig("https://example.test", "worker-b", tmp_path / "worker-b"),
|
||||
coordinator, artifacts, SciMeshRunner(),
|
||||
coordinator,
|
||||
artifacts,
|
||||
SciMeshRunner(),
|
||||
)
|
||||
|
||||
assert worker_a.run_once() == RunOnceOutcome(claimed=True, completed=False)
|
||||
@@ -197,16 +258,22 @@ def test_no_task_does_not_create_directory(tmp_path: Path) -> None:
|
||||
assert not config.work_dir.exists()
|
||||
|
||||
|
||||
def test_once_worker_exits_after_an_empty_claim(tmp_path: Path, caplog: pytest.LogCaptureFixture) -> None:
|
||||
def test_once_worker_exits_after_an_empty_claim(
|
||||
tmp_path: Path, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
caplog.set_level(logging.INFO, logger="scimesh.worker")
|
||||
worker, _, _, runner, _ = daemon(tmp_path, None, b"")
|
||||
worker.config = WorkerConfig(**{**worker.config.__dict__, "exit_when_idle": True, "max_tasks": 1})
|
||||
worker.config = WorkerConfig(
|
||||
**{**worker.config.__dict__, "exit_when_idle": True, "max_tasks": 1}
|
||||
)
|
||||
assert worker.run_forever() is True
|
||||
assert runner.calls == 0
|
||||
assert "queue_empty" in caplog.text
|
||||
|
||||
|
||||
def test_worker_stops_after_the_configured_number_of_claims(tmp_path: Path, caplog: pytest.LogCaptureFixture) -> None:
|
||||
def test_worker_stops_after_the_configured_number_of_claims(
|
||||
tmp_path: Path, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
caplog.set_level(logging.INFO, logger="scimesh.worker")
|
||||
content = b"input fixture"
|
||||
worker, _, _, runner, _ = daemon(tmp_path, make_task(content), content)
|
||||
@@ -216,10 +283,15 @@ def test_worker_stops_after_the_configured_number_of_claims(tmp_path: Path, capl
|
||||
assert "max_tasks_reached" in caplog.text
|
||||
|
||||
|
||||
def test_keyboard_interrupt_stops_worker_without_propagating(tmp_path: Path, caplog: pytest.LogCaptureFixture) -> None:
|
||||
def test_keyboard_interrupt_stops_worker_without_propagating(
|
||||
tmp_path: Path, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
caplog.set_level(logging.INFO, logger="scimesh.worker")
|
||||
|
||||
class InterruptingCoordinator(FakeCoordinator):
|
||||
def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None:
|
||||
def claim(
|
||||
self, worker_id: str, capabilities: tuple[str, ...]
|
||||
) -> ClaimedTask | None:
|
||||
raise KeyboardInterrupt
|
||||
|
||||
worker, _, _, _, _ = daemon(tmp_path, None, b"")
|
||||
@@ -228,7 +300,9 @@ def test_keyboard_interrupt_stops_worker_without_propagating(tmp_path: Path, cap
|
||||
assert "interrupted" in caplog.text
|
||||
|
||||
|
||||
def test_interrupting_an_active_task_reports_a_sanitized_failure(tmp_path: Path) -> None:
|
||||
def test_interrupting_an_active_task_reports_a_sanitized_failure(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
content = b"input fixture"
|
||||
worker, coordinator, _, _, _ = daemon(tmp_path, make_task(content), content)
|
||||
|
||||
@@ -271,12 +345,16 @@ def test_max_tasks_counts_successes_not_failed_claims(tmp_path: Path) -> None:
|
||||
),
|
||||
]
|
||||
|
||||
def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None:
|
||||
def claim(
|
||||
self, worker_id: str, capabilities: tuple[str, ...]
|
||||
) -> ClaimedTask | None:
|
||||
return self.tasks.pop(0) if self.tasks else None
|
||||
|
||||
coordinator = SequencedCoordinator()
|
||||
artifacts, runner = FakeArtifacts(successful_content), FakeRunner()
|
||||
config = WorkerConfig("https://example.test", "worker-1", tmp_path / "work", max_tasks=1)
|
||||
config = WorkerConfig(
|
||||
"https://example.test", "worker-1", tmp_path / "work", max_tasks=1
|
||||
)
|
||||
worker = WorkerDaemon(config, coordinator, artifacts, runner)
|
||||
assert worker.run_forever() is True
|
||||
assert len(coordinator.failures) == 1
|
||||
@@ -303,9 +381,12 @@ def test_worker_cli_uses_a_nonzero_exit_code_for_interruption(
|
||||
return False
|
||||
|
||||
monkeypatch.setattr(worker_cli, "WorkerDaemon", InterruptedDaemon)
|
||||
assert worker_cli.main(
|
||||
["--coordinator-url", "https://example.test", "--work-dir", str(tmp_path)]
|
||||
) == 130
|
||||
assert (
|
||||
worker_cli.main(
|
||||
["--coordinator-url", "https://example.test", "--work-dir", str(tmp_path)]
|
||||
)
|
||||
== 130
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", [0, -1, True])
|
||||
@@ -315,7 +396,9 @@ def test_max_tasks_must_be_positive(value: object, tmp_path: Path) -> None:
|
||||
|
||||
|
||||
def test_bad_checksum_reports_failure_without_running(tmp_path: Path) -> None:
|
||||
worker, coordinator, _, runner, _ = daemon(tmp_path, make_task(b"actual", "not-the-hash"), b"actual")
|
||||
worker, coordinator, _, runner, _ = daemon(
|
||||
tmp_path, make_task(b"actual", "not-the-hash"), b"actual"
|
||||
)
|
||||
assert worker.run_once() == RunOnceOutcome(claimed=True, completed=False)
|
||||
assert runner.calls == 0
|
||||
assert coordinator.failures[0]["error_code"] == "ValueError"
|
||||
@@ -323,7 +406,9 @@ def test_bad_checksum_reports_failure_without_running(tmp_path: Path) -> None:
|
||||
assert not coordinator.submissions
|
||||
|
||||
|
||||
def test_failure_reporting_removes_paths_outside_the_worker_directory(tmp_path: Path) -> None:
|
||||
def test_failure_reporting_removes_paths_outside_the_worker_directory(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
worker, coordinator, _, _, _ = daemon(tmp_path, make_task(b"input"), b"input")
|
||||
error = subprocess.CalledProcessError(
|
||||
1,
|
||||
@@ -344,9 +429,13 @@ def test_directory_creation_failure_is_reported(tmp_path: Path) -> None:
|
||||
assert coordinator.failures[0]["error_code"] == "FileExistsError"
|
||||
|
||||
|
||||
def test_transient_claim_error_is_propagated_for_bounded_backoff(tmp_path: Path) -> None:
|
||||
def test_transient_claim_error_is_propagated_for_bounded_backoff(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
class UnavailableCoordinator(FakeCoordinator):
|
||||
def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None:
|
||||
def claim(
|
||||
self, worker_id: str, capabilities: tuple[str, ...]
|
||||
) -> ClaimedTask | None:
|
||||
raise CoordinatorTransientError("temporary outage")
|
||||
|
||||
worker, _, _, _, _ = daemon(tmp_path, None, b"")
|
||||
@@ -368,7 +457,9 @@ def test_task_directories_are_retained_until_cleanup_is_enabled(tmp_path: Path)
|
||||
|
||||
def test_input_token_is_sent_only_to_the_coordinator_origin() -> None:
|
||||
client = HttpArtifactClient("https://coordinator.example/api", 10, "secret")
|
||||
assert client._auth_headers_for("https://coordinator.example/tasks/1/input") == {"Authorization": "Bearer secret"}
|
||||
assert client._auth_headers_for("https://coordinator.example/tasks/1/input") == {
|
||||
"Authorization": "Bearer secret"
|
||||
}
|
||||
assert client._auth_headers_for("https://bucket.example/presigned") == {}
|
||||
|
||||
|
||||
@@ -384,17 +475,28 @@ def test_relative_input_uri_is_resolved_against_the_coordinator() -> None:
|
||||
def test_redirect_to_external_storage_strips_authorization() -> None:
|
||||
handler = _SameOriginAuthRedirectHandler(_origin("https://coordinator.example"))
|
||||
source = Request(
|
||||
"https://coordinator.example/tasks/1/input", headers={"Authorization": "Bearer secret"}
|
||||
"https://coordinator.example/tasks/1/input",
|
||||
headers={"Authorization": "Bearer secret"},
|
||||
)
|
||||
redirected = handler.redirect_request(
|
||||
source, None, 302, "Found", {}, "https://bucket.example/presigned"
|
||||
)
|
||||
redirected = handler.redirect_request(source, None, 302, "Found", {}, "https://bucket.example/presigned")
|
||||
assert redirected is not None
|
||||
assert redirected.get_header("Authorization") is None
|
||||
|
||||
|
||||
def test_api_requests_never_follow_redirects() -> None:
|
||||
handler = NoRedirectHandler()
|
||||
request = Request("https://coordinator.example/tasks/claim", headers={"Authorization": "Bearer secret"})
|
||||
assert handler.redirect_request(request, None, 302, "Found", {}, "https://other.example") is None
|
||||
request = Request(
|
||||
"https://coordinator.example/tasks/claim",
|
||||
headers={"Authorization": "Bearer secret"},
|
||||
)
|
||||
assert (
|
||||
handler.redirect_request(
|
||||
request, None, 302, "Found", {}, "https://other.example"
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_lease_is_renewed_while_a_runner_is_still_working(tmp_path: Path) -> None:
|
||||
@@ -429,44 +531,101 @@ def test_heartbeat_reschedules_from_the_renewed_lease(tmp_path: Path) -> None:
|
||||
assert len(coordinator.heartbeats) >= 3
|
||||
|
||||
|
||||
def test_runner_maps_graph_and_smiles_search_parameters(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)
|
||||
def test_runner_executes_search_through_the_sdk_and_rejects_graph(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
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})
|
||||
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")
|
||||
with pytest.raises(ValueError, match="unsupported workload"):
|
||||
runner.run(graph, tmp_path / "graph")
|
||||
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 len(commands) == 1
|
||||
assert result.metrics == {
|
||||
"scanned_rows": 2, "valid_molecules": 2, "invalid_smiles": 0, "matches_emitted": 1,
|
||||
"scanned_rows": 2,
|
||||
"valid_molecules": 2,
|
||||
"invalid_smiles": 0,
|
||||
"matches_emitted": 1,
|
||||
}
|
||||
assert result.artifacts[0].content_type == "text/csv"
|
||||
assert (
|
||||
result.artifacts[0]
|
||||
.path.read_text(encoding="utf-8")
|
||||
.startswith("rank,chembl_id,canonical_smiles,similarity\n")
|
||||
)
|
||||
|
||||
|
||||
def test_runner_accepts_coordinator_workload_names(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_runner_resolves_query_id_from_the_shard_and_rejects_plan_time_options(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
task_dir = tmp_path / "search"
|
||||
task_dir.mkdir()
|
||||
(task_dir / "input").write_text(
|
||||
"chembl_id\tcanonical_smiles\nQUERY\tCCO\nMATCH\tCCCO\n", encoding="utf-8"
|
||||
)
|
||||
task = ClaimedTask(
|
||||
"search",
|
||||
1,
|
||||
"2026-07-30T00:00:00Z",
|
||||
"similarity-search",
|
||||
InputArtifact("https://example/input", "a" * 64),
|
||||
{"query_id": "QUERY", "top_k": 5},
|
||||
)
|
||||
result = SciMeshRunner().run(task, task_dir)
|
||||
assert result.metrics["matches_emitted"] == 1
|
||||
assert (task_dir / "result.csv").is_file()
|
||||
|
||||
with_max_rows = ClaimedTask(
|
||||
"search",
|
||||
1,
|
||||
"2026-07-30T00:00:00Z",
|
||||
"similarity-search",
|
||||
InputArtifact("https://example/input", "a" * 64),
|
||||
{"query_smiles": "CCO", "max_rows": 1},
|
||||
)
|
||||
with pytest.raises(ValueError, match="unsupported runner parameters"):
|
||||
SciMeshRunner().run(with_max_rows, task_dir)
|
||||
|
||||
|
||||
def test_runner_accepts_coordinator_workload_names(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
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"},
|
||||
"search",
|
||||
1,
|
||||
"2026-07-30T00:00:00Z",
|
||||
"similarity_search",
|
||||
InputArtifact("https://example/input", "a" * 64),
|
||||
{"query_smiles": "CCO"},
|
||||
)
|
||||
result = SciMeshRunner().run(task, task_dir)
|
||||
assert result.metrics["matches_emitted"] == 1
|
||||
@@ -502,7 +661,10 @@ def test_claimed_task_accepts_a_coordinator_relative_input_path() -> None:
|
||||
"attempt": 1,
|
||||
"lease_expires_at": "2026-07-30T00:00:00Z",
|
||||
"workload": "similarity_search",
|
||||
"input": {"uri": "/tasks/11111111-1111-4111-8111-111111111111/input", "sha256": "a" * 64},
|
||||
"input": {
|
||||
"uri": "/tasks/11111111-1111-4111-8111-111111111111/input",
|
||||
"sha256": "a" * 64,
|
||||
},
|
||||
"parameters": {},
|
||||
}
|
||||
)
|
||||
@@ -523,7 +685,9 @@ def test_uploaded_artifact_requires_complete_durable_metadata() -> None:
|
||||
UploadedArtifact.from_json({"artifact_id": "missing"})
|
||||
|
||||
|
||||
def test_environment_overrides_allow_cli_only_configuration(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None:
|
||||
def test_environment_overrides_allow_cli_only_configuration(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
|
||||
) -> None:
|
||||
monkeypatch.delenv("SCIMESH_COORDINATOR_URL", raising=False)
|
||||
config = WorkerConfig.from_environment(
|
||||
{
|
||||
@@ -552,8 +716,12 @@ def test_relative_work_dir_is_normalized_for_runner_subprocesses(
|
||||
"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"},
|
||||
"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 (task_dir / "result.csv").is_file()
|
||||
|
||||
Reference in New Issue
Block a user