diff --git a/README.md b/README.md index 9ddef99..be03abe 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,11 @@ # SciMesh -SciMesh is a small local framework for scientific workloads on molecular datasets. It currently provides exact molecular similarity search and exact sparse similarity-graph construction. It runs in one local Python process: there is no network service, multiprocessing, coordinator, database, or dense similarity matrix. +SciMesh is a scientific-workload framework for molecular datasets. Its public CLI +currently runs exact similarity search and sparse similarity-graph construction +locally in one Python process; it creates no dense similarity matrix. A Python +Worker client and the planned Go/PostgreSQL coordinator contract are tracked in +the repository, but distributed execution is not available yet; see +[`STATUS.md`](STATUS.md). The ChEMBL TSV database is intentionally not included in this repository. Download it separately and pass its path to the commands below. The expected columns are `chembl_id` and `canonical_smiles`. diff --git a/docs/api-contract.md b/docs/api-contract.md index 33acbd4..e8f7834 100644 --- a/docs/api-contract.md +++ b/docs/api-contract.md @@ -14,6 +14,8 @@ pull request as both implementation and contract tests. - A task becomes `completed` only after a coordinator-owned artifact is durable. - Identical repeated completion is successful; a different result for the same attempt is a conflict. +- Coordinator API calls do not follow redirects. Artifact downloads may follow + redirects only after removing the coordinator bearer token on origin change. ## Worker registration diff --git a/docs/worker-daemon-task.md b/docs/worker-daemon-task.md index 82c7538..c5e7445 100644 --- a/docs/worker-daemon-task.md +++ b/docs/worker-daemon-task.md @@ -27,7 +27,9 @@ inside the daemon. `scimesh-worker`. 2. Configuration via environment variables and CLI overrides: - `SCIMESH_COORDINATOR_URL` (required); - - `SCIMESH_WORKER_ID` (required, stable UUID or hostname-derived value); + - `SCIMESH_WORKER_NAME` (optional; defaults to the hostname); + - `SCIMESH_WORKER_ID` (optional legacy/test override; production identity is + returned by registration); - working directory for downloaded inputs and generated outputs; - poll interval and request timeout; - optional bearer token. @@ -43,6 +45,28 @@ inside the daemon. Use JSON over HTTPS. Claiming a task changes its state, so use `POST`, even if the initial diagram labels the endpoint as `GET /get_task`. +`docs/api-contract.md` is the authoritative API schema. This document explains +the daemon workflow and must not introduce a different request or response +shape. + +### Register worker + +At daemon startup, register the worker capabilities before claiming tasks: + +```http +POST /workers/register +Content-Type: application/json + +{ + "name": "lab-worker-01", + "capabilities": ["similarity-search", "similarity-graph"], + "cpu_count": 8, + "memory_mb": 16384 +} +``` + +The `worker_id` returned by this endpoint is used for the daemon lifetime. + ### Claim a task ```http @@ -90,8 +114,8 @@ Content-Type: application/json { "worker_id": "worker-01", "attempt": 1, - "status": "completed", "result": { + "artifact_id": "0d2d5a53-4c7e-467e-93d2-45ed2dc18e46", "uri": "https://coordinator.example/tasks/0d2d/result.csv", "sha256": "...", "content_type": "text/csv" @@ -121,7 +145,10 @@ The coordinator streams the artifact to its configured storage and responds: ```json { - "uri": "https://coordinator.example/tasks/0d2d/artifacts/result.csv" + "artifact_id": "0d2d5a53-4c7e-467e-93d2-45ed2dc18e46", + "uri": "https://coordinator.example/tasks/0d2d/artifacts/result.csv", + "sha256": "...", + "size_bytes": 1234 } ``` @@ -142,7 +169,9 @@ idle -> claiming -> downloading -> running -> uploading -> submitting -> idle - Verify the input checksum before running. - Create one isolated task directory: `///`. - Invoke the runner with an explicit argument list, never `shell=True`. -- Upload/submit exactly the produced result files listed by the runner. +- Upload the produced result artifact before submitting its manifest. +- Version 1 produces exactly one CSV partial result. Multi-artifact manifests + require an explicit future API-contract change. - Do not mark a task completed until every submitted artifact has a durable coordinator-provided URI. - A timeout, network error, or rejected submission must leave the local task diff --git a/scimesh/worker/artifacts.py b/scimesh/worker/artifacts.py index b075c42..b18e7f8 100644 --- a/scimesh/worker/artifacts.py +++ b/scimesh/worker/artifacts.py @@ -8,37 +8,22 @@ import json from pathlib import Path from typing import Protocol from urllib.parse import quote, urlsplit -from urllib.request import HTTPRedirectHandler, Request, build_opener +from urllib.request import Request, build_opener -from .models import ClaimedTask, ProducedArtifact +from .coordinator import CoordinatorConflictError +from .models import ClaimedTask, ProducedArtifact, UploadedArtifact +from .transport import SameOriginAuthRedirectHandler, origin - -def _origin(uri: str) -> tuple[str, str, int | None]: - parsed = urlsplit(uri) - scheme = parsed.scheme.lower() - default_port = {"http": 80, "https": 443}.get(scheme) - return scheme, (parsed.hostname or "").lower(), parsed.port or default_port - - -class _SameOriginAuthRedirectHandler(HTTPRedirectHandler): - """Do not forward the coordinator token when a download changes origin.""" - - def __init__(self, coordinator_origin: tuple[str, str, int | None]) -> None: - super().__init__() - self.coordinator_origin = coordinator_origin - - def redirect_request(self, req: Request, fp: object, code: int, msg: str, headers: object, newurl: str) -> Request | None: - redirected = super().redirect_request(req, fp, code, msg, headers, newurl) - if redirected and _origin(newurl) != self.coordinator_origin: - redirected.remove_header("Authorization") - return redirected +# Compatibility aliases for focused transport tests. +_SameOriginAuthRedirectHandler = SameOriginAuthRedirectHandler +_origin = origin class ArtifactClient(Protocol): def download(self, uri: str, destination: Path) -> None: ... def upload( self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact - ) -> str: ... + ) -> UploadedArtifact: ... class HttpArtifactClient: @@ -48,8 +33,8 @@ class HttpArtifactClient: self.coordinator_url = coordinator_url.rstrip("/") self.timeout = timeout self.bearer_token = bearer_token - self.coordinator_origin = _origin(coordinator_url) - self._opener = build_opener(_SameOriginAuthRedirectHandler(self.coordinator_origin)) + self.coordinator_origin = origin(coordinator_url) + self._opener = build_opener(SameOriginAuthRedirectHandler(self.coordinator_origin)) def download(self, uri: str, destination: Path) -> None: destination.parent.mkdir(parents=True, exist_ok=True) @@ -58,8 +43,10 @@ class HttpArtifactClient: while chunk := response.read(1024 * 1024): target.write(chunk) - def upload(self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact) -> str: - """Stream one result artifact to the coordinator and return its stable URI.""" + def upload( + self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact + ) -> UploadedArtifact: + """Stream an artifact and require durable coordinator-owned metadata.""" url = ( f"{self.coordinator_url}/tasks/{quote(task.task_id, safe='')}/artifacts/" f"{quote(artifact.path.name, safe='')}" @@ -71,11 +58,13 @@ class HttpArtifactClient: http.client.HTTPSConnection if parsed.scheme == "https" else http.client.HTTPConnection ) connection = connection_class(parsed.hostname, parsed.port, timeout=self.timeout) + local_size = artifact.path.stat().st_size + local_sha256 = sha256_file(artifact.path) try: path = parsed.path + (f"?{parsed.query}" if parsed.query else "") connection.putrequest("PUT", path) connection.putheader("Content-Type", artifact.content_type) - connection.putheader("Content-Length", str(artifact.path.stat().st_size)) + connection.putheader("Content-Length", str(local_size)) connection.putheader("X-Worker-ID", worker_id) connection.putheader("X-Task-Attempt", str(task.attempt)) for name, value in self._auth_headers_for(url).items(): @@ -86,23 +75,24 @@ class HttpArtifactClient: connection.send(chunk) response = connection.getresponse() body = response.read() - if not 200 <= response.status < 300: + if response.status == 409: + raise CoordinatorConflictError("artifact upload rejected because the task lease was lost") + if response.status != 201: raise RuntimeError(f"artifact upload rejected with status {response.status}") - if body: - try: - response_data = json.loads(body) - except json.JSONDecodeError as error: - raise RuntimeError("artifact upload returned invalid JSON") from error - response_uri = response_data.get("uri") if isinstance(response_data, dict) else None - if isinstance(response_uri, str) and response_uri: - return response_uri - return url + try: + response_data = json.loads(body) + uploaded = UploadedArtifact.from_json(response_data) + except (ValueError, json.JSONDecodeError) as error: + raise RuntimeError("artifact upload returned invalid metadata") from error + if uploaded.sha256 != local_sha256 or uploaded.size_bytes != local_size: + raise RuntimeError("artifact upload metadata does not match local artifact") + return uploaded finally: connection.close() def _auth_headers_for(self, uri: str) -> dict[str, str]: """Only coordinator-owned URLs receive the coordinator bearer token.""" - if self.bearer_token and _origin(uri) == self.coordinator_origin: + if self.bearer_token and origin(uri) == self.coordinator_origin: return {"Authorization": f"Bearer {self.bearer_token}"} return {} diff --git a/scimesh/worker/cli.py b/scimesh/worker/cli.py index 7d3bf1f..13e1fb5 100644 --- a/scimesh/worker/cli.py +++ b/scimesh/worker/cli.py @@ -14,22 +14,42 @@ from .runners import SciMeshRunner def main(argv: list[str] | None = None) -> int: - parser = argparse.ArgumentParser(prog="scimesh-worker") + parser = argparse.ArgumentParser( + prog="scimesh-worker", + epilog=( + "Environment: SCIMESH_COORDINATOR_URL, SCIMESH_WORK_DIR, " + "SCIMESH_WORKER_NAME, SCIMESH_CPU_COUNT, SCIMESH_MEMORY_MB, " + "SCIMESH_POLL_INTERVAL, SCIMESH_REQUEST_TIMEOUT, " + "SCIMESH_HEARTBEAT_INTERVAL, SCIMESH_CLEANUP_AFTER_SECONDS, and " + "SCIMESH_BEARER_TOKEN. SCIMESH_WORKER_ID is a legacy/test override." + ), + ) parser.add_argument("--coordinator-url") parser.add_argument("--worker-id") parser.add_argument("--work-dir") + parser.add_argument("--worker-name") + parser.add_argument("--cpu-count", type=int) + parser.add_argument("--memory-mb", type=int) parser.add_argument("--poll-interval", type=float) parser.add_argument("--request-timeout", type=float) parser.add_argument("--heartbeat-interval", type=float) + parser.add_argument("--cleanup-after-seconds", type=float) args = parser.parse_args(argv) - config = WorkerConfig.from_environment() overrides = {key: value for key, value in vars(args).items() if value is not None} if "work_dir" in overrides: overrides["work_dir"] = Path(overrides["work_dir"]) - config = WorkerConfig(**{**config.__dict__, **overrides}) + try: + config = WorkerConfig.from_environment(overrides) + except (TypeError, ValueError) as error: + parser.error(str(error)) logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") client = HttpCoordinatorClient(config.coordinator_url, config.request_timeout, config.bearer_token) - WorkerDaemon(config, client, HttpArtifactClient(config.coordinator_url, config.request_timeout, config.bearer_token), SciMeshRunner()).run_forever() + WorkerDaemon( + config, + client, + HttpArtifactClient(config.coordinator_url, config.request_timeout, config.bearer_token), + SciMeshRunner(), + ).run_forever() return 0 diff --git a/scimesh/worker/config.py b/scimesh/worker/config.py index 0b5ca91..4a172cb 100644 --- a/scimesh/worker/config.py +++ b/scimesh/worker/config.py @@ -3,15 +3,34 @@ from __future__ import annotations from dataclasses import dataclass +from math import isfinite from pathlib import Path import os +import socket +from typing import Mapping +from urllib.parse import urlsplit + + +def _positive_number(value: object, name: str, *, allow_zero: bool = False) -> None: + if ( + isinstance(value, bool) + or not isinstance(value, (int, float)) + or not isfinite(value) + or value < 0 + or (not allow_zero and value == 0) + ): + qualifier = "non-negative" if allow_zero else "positive" + raise ValueError(f"{name} must be {qualifier}") @dataclass(frozen=True) class WorkerConfig: coordinator_url: str - worker_id: str + worker_id: str | None work_dir: Path + worker_name: str = "scimesh-worker" + cpu_count: int = 1 + memory_mb: int | None = None poll_interval: float = 2.0 request_timeout: float = 30.0 heartbeat_interval: float = 15.0 @@ -20,27 +39,56 @@ class WorkerConfig: capabilities: tuple[str, ...] = ("similarity-search", "similarity-graph") def __post_init__(self) -> None: - if self.poll_interval <= 0: - raise ValueError("poll_interval must be positive") - if self.request_timeout <= 0: - raise ValueError("request_timeout must be positive") - if self.heartbeat_interval <= 0: - raise ValueError("heartbeat_interval must be positive") + parsed = urlsplit(self.coordinator_url) + if parsed.scheme not in {"http", "https"} or not parsed.hostname: + raise ValueError("coordinator_url must be an absolute HTTP(S) URL") + if not isinstance(self.worker_name, str) or not self.worker_name.strip(): + raise ValueError("worker_name must be non-empty") + if isinstance(self.cpu_count, bool) or not isinstance(self.cpu_count, int) or self.cpu_count < 1: + raise ValueError("cpu_count must be positive") + if self.worker_id is not None and not isinstance(self.worker_id, str): + raise ValueError("worker_id must be a string when set") + if self.memory_mb is not None and ( + isinstance(self.memory_mb, bool) + or not isinstance(self.memory_mb, int) + or self.memory_mb < 1 + ): + raise ValueError("memory_mb must be positive when set") + _positive_number(self.poll_interval, "poll_interval") + _positive_number(self.request_timeout, "request_timeout") + _positive_number(self.heartbeat_interval, "heartbeat_interval") + if self.cleanup_after_seconds is not None: + _positive_number(self.cleanup_after_seconds, "cleanup_after_seconds", allow_zero=True) + if not self.capabilities: + raise ValueError("capabilities cannot be empty") @classmethod - def from_environment(cls) -> "WorkerConfig": - url = os.getenv("SCIMESH_COORDINATOR_URL") - worker_id = os.getenv("SCIMESH_WORKER_ID") - if not url or not worker_id: - raise ValueError("SCIMESH_COORDINATOR_URL and SCIMESH_WORKER_ID are required") - cleanup = os.getenv("SCIMESH_CLEANUP_AFTER_SECONDS") + def from_environment( + cls, overrides: Mapping[str, object] | None = None + ) -> "WorkerConfig": + """Build config from environment, allowing typed CLI values to override it.""" + values = overrides or {} + + def value(name: str, environment: str, default: object | None = None) -> object | None: + override = values.get(name) + return override if override is not None else os.getenv(environment, default) + + url = value("coordinator_url", "SCIMESH_COORDINATOR_URL") + if not isinstance(url, str) or not url: + raise ValueError("SCIMESH_COORDINATOR_URL or --coordinator-url is required") + cleanup = value("cleanup_after_seconds", "SCIMESH_CLEANUP_AFTER_SECONDS") + cpu_count = value("cpu_count", "SCIMESH_CPU_COUNT", os.cpu_count() or 1) + memory_mb = value("memory_mb", "SCIMESH_MEMORY_MB") return cls( coordinator_url=url.rstrip("/"), - worker_id=worker_id, - work_dir=Path(os.getenv("SCIMESH_WORK_DIR", "./scimesh-worker-data")), - poll_interval=float(os.getenv("SCIMESH_POLL_INTERVAL", "2")), - request_timeout=float(os.getenv("SCIMESH_REQUEST_TIMEOUT", "30")), - heartbeat_interval=float(os.getenv("SCIMESH_HEARTBEAT_INTERVAL", "15")), - bearer_token=os.getenv("SCIMESH_BEARER_TOKEN"), + worker_id=value("worker_id", "SCIMESH_WORKER_ID"), + work_dir=Path(value("work_dir", "SCIMESH_WORK_DIR", "./scimesh-worker-data")), + worker_name=str(value("worker_name", "SCIMESH_WORKER_NAME", socket.gethostname())), + cpu_count=int(cpu_count), + memory_mb=int(memory_mb) if memory_mb is not None else None, + poll_interval=float(value("poll_interval", "SCIMESH_POLL_INTERVAL", "2")), + request_timeout=float(value("request_timeout", "SCIMESH_REQUEST_TIMEOUT", "30")), + heartbeat_interval=float(value("heartbeat_interval", "SCIMESH_HEARTBEAT_INTERVAL", "15")), + bearer_token=value("bearer_token", "SCIMESH_BEARER_TOKEN"), cleanup_after_seconds=float(cleanup) if cleanup else None, ) diff --git a/scimesh/worker/coordinator.py b/scimesh/worker/coordinator.py index 5fccd33..7f3d688 100644 --- a/scimesh/worker/coordinator.py +++ b/scimesh/worker/coordinator.py @@ -5,9 +5,10 @@ from __future__ import annotations import json from typing import Any, Protocol from urllib.error import HTTPError, URLError -from urllib.request import Request, urlopen +from urllib.request import Request, build_opener -from .models import ClaimedTask +from .models import ClaimedTask, RegisteredWorker +from .transport import NoRedirectHandler class CoordinatorError(RuntimeError): @@ -18,7 +19,15 @@ class CoordinatorTransientError(CoordinatorError): """A timeout, connection error, or 5xx coordinator response.""" +class CoordinatorConflictError(CoordinatorError): + """The worker no longer owns the task lease or attempted a conflicting mutation.""" + + class CoordinatorClient(Protocol): + def register( + self, name: str, capabilities: tuple[str, ...], cpu_count: int, memory_mb: int | None + ) -> RegisteredWorker: ... + def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None: ... def submit(self, task: ClaimedTask, payload: dict[str, Any]) -> None: ... @@ -33,6 +42,25 @@ class HttpCoordinatorClient: self.base_url = base_url.rstrip("/") self.timeout = timeout self.bearer_token = bearer_token + self._opener = build_opener(NoRedirectHandler()) + + def register( + self, name: str, capabilities: tuple[str, ...], cpu_count: int, memory_mb: int | None + ) -> RegisteredWorker: + payload: dict[str, Any] = { + "name": name, + "capabilities": list(capabilities), + "cpu_count": cpu_count, + } + if memory_mb is not None: + payload["memory_mb"] = memory_mb + status, body = self._request("POST", "/workers/register", payload) + if status != 200: + raise CoordinatorError(f"worker registration rejected with status {status}") + try: + return RegisteredWorker.from_json(body) + except ValueError as error: + raise CoordinatorError("invalid worker registration response") from error def claim(self, worker_id: str, capabilities: tuple[str, ...]) -> ClaimedTask | None: status, body = self._request("POST", "/tasks/claim", { @@ -48,11 +76,15 @@ class HttpCoordinatorClient: status, _ = self._request("POST", f"/tasks/{task.task_id}/result", payload) # 200/201/202 include a successful or idempotent duplicate result response. if status not in (200, 201, 202): + if status == 409: + raise CoordinatorConflictError("result rejected because the task lease was lost") raise CoordinatorError(f"result rejected with status {status}") def fail(self, task: ClaimedTask, payload: dict[str, Any]) -> None: status, _ = self._request("POST", f"/tasks/{task.task_id}/failure", payload) if status not in (200, 201, 202): + if status == 409: + raise CoordinatorConflictError("failure rejected because the task lease was lost") raise CoordinatorError(f"failure report rejected with status {status}") def heartbeat(self, task: ClaimedTask, worker_id: str) -> str: @@ -61,6 +93,8 @@ class HttpCoordinatorClient: {"worker_id": worker_id, "attempt": task.attempt}, ) if status != 200: + if status == 409: + raise CoordinatorConflictError("heartbeat rejected because the task lease was lost") raise CoordinatorError(f"heartbeat rejected with status {status}") lease_expires_at = body.get("lease_expires_at") if not isinstance(lease_expires_at, str): @@ -73,9 +107,12 @@ class HttpCoordinatorClient: headers={"Content-Type": "application/json", **self._auth_header()}, ) try: - with urlopen(request, timeout=self.timeout) as response: + with self._opener.open(request, timeout=self.timeout) as response: raw = response.read() - return response.status, json.loads(raw) if raw else {} + try: + return response.status, json.loads(raw) if raw else {} + except json.JSONDecodeError as error: + raise CoordinatorError("coordinator returned invalid JSON") from error except HTTPError as error: if error.code >= 500: raise CoordinatorTransientError(f"coordinator returned {error.code}") from error diff --git a/scimesh/worker/daemon.py b/scimesh/worker/daemon.py index 3200f6b..121a206 100644 --- a/scimesh/worker/daemon.py +++ b/scimesh/worker/daemon.py @@ -3,6 +3,7 @@ from __future__ import annotations import logging +from dataclasses import replace from pathlib import Path import random import shutil @@ -12,8 +13,8 @@ from datetime import datetime, timezone from .artifacts import ArtifactClient, sha256_file from .config import WorkerConfig -from .coordinator import CoordinatorClient, CoordinatorTransientError -from .models import ClaimedTask +from .coordinator import CoordinatorClient, CoordinatorConflictError, CoordinatorTransientError +from .models import ClaimedTask, UploadedArtifact from .runners import Runner @@ -32,6 +33,7 @@ class LeaseHeartbeat: self._lease_expires_at = self.coordinator.heartbeat( self.task, self.config.worker_id ) + self._next_delay() self._thread = threading.Thread(target=self._run, name=f"lease-{self.task.task_id}", daemon=True) self._thread.start() @@ -45,18 +47,19 @@ class LeaseHeartbeat: raise self._error def _run(self) -> None: - delay = min(self.config.heartbeat_interval, self._seconds_until_expiry() / 2) + delay = self._next_delay() while not self._stop.wait(max(delay, 0.01)): try: self._lease_expires_at = self.coordinator.heartbeat( self.task, self.config.worker_id ) + delay = self._next_delay() except Exception as error: # Surface the lease loss in the main state machine. self._error = error return - delay = min( - self.config.heartbeat_interval, self._seconds_until_expiry() / 2 - ) + + def _next_delay(self) -> float: + return min(self.config.heartbeat_interval, self._seconds_until_expiry() / 2) def _seconds_until_expiry(self) -> float: try: @@ -72,12 +75,16 @@ class LeaseHeartbeat: class WorkerDaemon: def __init__(self, config: WorkerConfig, coordinator: CoordinatorClient, artifacts: ArtifactClient, runner: Runner) -> None: self.config, self.coordinator, self.artifacts, self.runner = config, coordinator, artifacts, runner + self.worker_id = config.worker_id + self._registered = False self.log = logging.getLogger("scimesh.worker") def run_forever(self) -> None: failures = 0 while True: try: + if not self._registered: + self._register_worker() self._cleanup_expired_directories() claimed = self.run_once() failures = 0 @@ -89,16 +96,17 @@ class WorkerDaemon: self._sleep(min(self.config.poll_interval * 2 ** min(failures, 6), 60.0)) def run_once(self) -> bool: + worker_id = self._worker_id() self._log("claiming") - task = self.coordinator.claim(self.config.worker_id, self.config.capabilities) + task = self.coordinator.claim(worker_id, self.config.capabilities) if task is None: self._log("idle") return False started = time.monotonic() task_dir = self.config.work_dir / task.task_id / str(task.attempt) - task_dir.mkdir(parents=True, exist_ok=False) heartbeat = LeaseHeartbeat(task, self.coordinator, self.config) try: + task_dir.mkdir(parents=True, exist_ok=False) heartbeat.start() self._log("downloading", task) input_path = task_dir / "input" @@ -108,20 +116,28 @@ class WorkerDaemon: self._log("running", task) result = self.runner.run(task, task_dir) heartbeat.raise_if_failed() - manifests = [ - { - "uri": self.artifacts.upload(task, self.config.worker_id, artifact), - "sha256": sha256_file(artifact.path), - "content_type": artifact.content_type, - } - for artifact in result.artifacts - ] - if not manifests: - raise ValueError("runner produced no artifacts") + if len(result.artifacts) != 1: + raise ValueError("runner must produce exactly one result artifact") + artifact = result.artifacts[0] + uploaded = self.artifacts.upload(task, worker_id, artifact) + manifest = self._result_manifest(uploaded, artifact.content_type) self._log("submitting", task) heartbeat.raise_if_failed() - self.coordinator.submit(task, {"worker_id": self.config.worker_id, "attempt": task.attempt, "status": "completed", "result": manifests[0], "artifacts": manifests, "metrics": {**result.metrics, "elapsed_seconds": round(time.monotonic() - started, 3)}}) + self.coordinator.submit( + task, + { + "worker_id": worker_id, + "attempt": task.attempt, + "result": manifest, + "metrics": { + **result.metrics, + "elapsed_seconds": round(time.monotonic() - started, 3), + }, + }, + ) self._log("idle", task, elapsed_seconds=round(time.monotonic() - started, 3)) + except CoordinatorConflictError as error: + self._log("lease_lost", task, error_type=type(error).__name__) except Exception as error: self._log("failed", task, error_type=type(error).__name__) self._report_failure(task, error) @@ -132,12 +148,42 @@ class WorkerDaemon: def _report_failure(self, task: ClaimedTask, error: Exception) -> None: message = str(error).replace(str(self.config.work_dir), "")[:300] try: - self.coordinator.fail(task, {"worker_id": self.config.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}) except CoordinatorTransientError: raise except Exception: self._log("failed", task, error_type="FailureReportError") + def _register_worker(self) -> None: + registered = self.coordinator.register( + self.config.worker_name, + self.config.capabilities, + self.config.cpu_count, + self.config.memory_mb, + ) + self.worker_id = registered.worker_id + self.config = replace( + self.config, + worker_id=registered.worker_id, + heartbeat_interval=registered.heartbeat_interval_seconds, + ) + self._registered = True + self._log("registered") + + def _worker_id(self) -> str: + if not self.worker_id: + raise ValueError("worker is not registered") + return self.worker_id + + @staticmethod + def _result_manifest(uploaded: UploadedArtifact, content_type: str) -> dict[str, object]: + return { + "artifact_id": uploaded.artifact_id, + "uri": uploaded.uri, + "sha256": uploaded.sha256, + "content_type": content_type, + } + def _log(self, state: str, task: ClaimedTask | None = None, **extra: object) -> None: fields = {"worker_id": self.config.worker_id, "task_id": task.task_id if task else None, "attempt": task.attempt if task else None, "state": state, **extra} self.log.info("worker_event %s", fields) diff --git a/scimesh/worker/models.py b/scimesh/worker/models.py index 97de6a1..f1540e2 100644 --- a/scimesh/worker/models.py +++ b/scimesh/worker/models.py @@ -3,8 +3,33 @@ from __future__ import annotations from dataclasses import dataclass +from datetime import datetime +from math import isfinite from pathlib import Path from typing import Any +from urllib.parse import urlsplit +from uuid import UUID + + +def _required_string(value: object, field: str) -> str: + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"{field} must be a non-empty string") + return value + + +def _http_uri(value: object, field: str) -> str: + uri = _required_string(value, field) + parsed = urlsplit(uri) + if parsed.scheme not in {"http", "https"} or not parsed.hostname: + raise ValueError(f"{field} must be an absolute HTTP(S) URL") + return uri + + +def _sha256(value: object, field: str) -> str: + digest = _required_string(value, field).lower() + if len(digest) != 64 or any(character not in "0123456789abcdef" for character in digest): + raise ValueError(f"{field} must be a SHA-256 hex digest") + return digest @dataclass(frozen=True) @@ -26,13 +51,28 @@ class ClaimedTask: def from_json(cls, data: dict[str, Any]) -> "ClaimedTask": try: input_data = data["input"] + if not isinstance(input_data, dict): + raise ValueError("input must be an object") + raw_attempt = data["attempt"] + if isinstance(raw_attempt, bool) or not isinstance(raw_attempt, int) or raw_attempt < 1: + raise ValueError("attempt must be a positive integer") + task_id = str(UUID(_required_string(data["task_id"], "task_id"))) + lease_expires_at = _required_string(data["lease_expires_at"], "lease_expires_at") + if datetime.fromisoformat(lease_expires_at.replace("Z", "+00:00")).tzinfo is None: + raise ValueError("lease_expires_at must include a timezone") + parameters = data.get("parameters", {}) + if not isinstance(parameters, dict): + raise ValueError("parameters must be an object") return cls( - task_id=str(data["task_id"]), - attempt=int(data["attempt"]), - lease_expires_at=str(data["lease_expires_at"]), - workload=str(data["workload"]), - input=InputArtifact(uri=str(input_data["uri"]), sha256=str(input_data["sha256"])), - parameters=dict(data.get("parameters", {})), + task_id=task_id, + attempt=raw_attempt, + lease_expires_at=lease_expires_at, + workload=_required_string(data["workload"], "workload"), + input=InputArtifact( + uri=_http_uri(input_data["uri"], "input.uri"), + sha256=_sha256(input_data["sha256"], "input.sha256"), + ), + parameters=parameters, ) except (KeyError, TypeError, ValueError) as error: raise ValueError("invalid claimed-task response") from error @@ -44,6 +84,61 @@ class ProducedArtifact: content_type: str +@dataclass(frozen=True) +class UploadedArtifact: + """Coordinator-owned artifact metadata returned after a successful upload.""" + + artifact_id: str + uri: str + sha256: str + size_bytes: int + + @classmethod + def from_json(cls, data: object) -> "UploadedArtifact": + if not isinstance(data, dict): + raise ValueError("artifact upload response must be an object") + raw_size = data.get("size_bytes") + if isinstance(raw_size, bool) or not isinstance(raw_size, int) or raw_size < 0: + raise ValueError("artifact size_bytes must be a non-negative integer") + try: + return cls( + artifact_id=str(UUID(_required_string(data.get("artifact_id"), "artifact_id"))), + uri=_http_uri(data.get("uri"), "uri"), + sha256=_sha256(data.get("sha256"), "sha256"), + size_bytes=raw_size, + ) + except ValueError as error: + raise ValueError("invalid artifact upload response") from error + + +@dataclass(frozen=True) +class RegisteredWorker: + """Identity and heartbeat policy returned by worker registration.""" + + worker_id: str + heartbeat_interval_seconds: float + + @classmethod + def from_json(cls, data: object) -> "RegisteredWorker": + if not isinstance(data, dict): + raise ValueError("worker registration response must be an object") + raw_interval = data.get("heartbeat_interval_seconds") + if ( + isinstance(raw_interval, bool) + or not isinstance(raw_interval, (int, float)) + or not isfinite(raw_interval) + or raw_interval <= 0 + ): + raise ValueError("heartbeat_interval_seconds must be positive") + try: + return cls( + worker_id=str(UUID(_required_string(data.get("worker_id"), "worker_id"))), + heartbeat_interval_seconds=float(raw_interval), + ) + except ValueError as error: + raise ValueError("invalid worker registration response") from error + + @dataclass(frozen=True) class RunResult: artifacts: tuple[ProducedArtifact, ...] diff --git a/scimesh/worker/transport.py b/scimesh/worker/transport.py new file mode 100644 index 0000000..3f975f8 --- /dev/null +++ b/scimesh/worker/transport.py @@ -0,0 +1,51 @@ +"""Small HTTP transport helpers shared by coordinator and artifact clients.""" + +from __future__ import annotations + +from urllib.request import HTTPRedirectHandler, Request +from urllib.parse import urlsplit + + +def origin(uri: str) -> tuple[str, str, int | None]: + """Return a normalized HTTP origin for authorization decisions.""" + parsed = urlsplit(uri) + scheme = parsed.scheme.lower() + default_port = {"http": 80, "https": 443}.get(scheme) + return scheme, (parsed.hostname or "").lower(), parsed.port or default_port + + +class SameOriginAuthRedirectHandler(HTTPRedirectHandler): + """Strip coordinator authorization when an artifact redirect changes origin.""" + + def __init__(self, coordinator_origin: tuple[str, str, int | None]) -> None: + super().__init__() + self.coordinator_origin = coordinator_origin + + def redirect_request( + self, + req: Request, + fp: object, + code: int, + msg: str, + headers: object, + newurl: str, + ) -> Request | None: + redirected = super().redirect_request(req, fp, code, msg, headers, newurl) + if redirected and origin(newurl) != self.coordinator_origin: + redirected.remove_header("Authorization") + return redirected + + +class NoRedirectHandler(HTTPRedirectHandler): + """Reject redirects for mutating coordinator API calls.""" + + def redirect_request( + self, + req: Request, + fp: object, + code: int, + msg: str, + headers: object, + newurl: str, + ) -> Request | None: + return None diff --git a/scimesh/workloads/similarity_graph.py b/scimesh/workloads/similarity_graph.py index 4c6fa0d..08df9a2 100644 --- a/scimesh/workloads/similarity_graph.py +++ b/scimesh/workloads/similarity_graph.py @@ -47,10 +47,15 @@ def _fingerprinted_molecules( tsv_path: Path, max_rows: int | None ) -> tuple[list[GraphMolecule], DatasetStats]: stats = DatasetStats() - molecules = [ - GraphMolecule(record.molecule_id, fingerprint(record.molecule)) - for record in iter_valid_molecules(tsv_path, stats, max_rows=max_rows) - ] + molecules: list[GraphMolecule] = [] + seen_ids: set[str] = set() + for record in iter_valid_molecules(tsv_path, stats, max_rows=max_rows): + if not record.molecule_id: + raise ValueError("Dataset contains an empty chembl_id") + if record.molecule_id in seen_ids: + raise ValueError(f"Dataset contains a duplicate chembl_id: {record.molecule_id}") + seen_ids.add(record.molecule_id) + molecules.append(GraphMolecule(record.molecule_id, fingerprint(record.molecule))) return molecules, stats @@ -119,6 +124,7 @@ def build_similarity_graph( def write_graph_edges(output_path: Path, edges: list[SimilarityEdge]) -> None: """Write a deterministic sparse edge list CSV.""" + 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=["source_id", "target_id", "similarity"]) writer.writeheader() diff --git a/scimesh/workloads/similarity_search.py b/scimesh/workloads/similarity_search.py index d3c92de..574749a 100644 --- a/scimesh/workloads/similarity_search.py +++ b/scimesh/workloads/similarity_search.py @@ -135,6 +135,7 @@ def search_similar( def write_search_results(output_path: Path, matches: list[SimilarityMatch]) -> None: """Write ranked matches to a deterministic CSV file.""" + output_path.parent.mkdir(parents=True, exist_ok=True) with output_path.open("w", encoding="utf-8", newline="") as destination: writer = csv.DictWriter( destination, diff --git a/tests/test_similarity_graph.py b/tests/test_similarity_graph.py index 9341e01..2399ac2 100644 --- a/tests/test_similarity_graph.py +++ b/tests/test_similarity_graph.py @@ -2,6 +2,7 @@ from __future__ import annotations from pathlib import Path +import pytest from rdkit import DataStructs from scimesh.chemistry.dataset import DatasetStats, iter_valid_molecules @@ -61,3 +62,19 @@ def test_graph_supports_less_than_threshold_direction(small_dataset: Path) -> No ) assert all(edge.similarity <= 0.15 for edge in result.edges) + + +def test_graph_rejects_duplicate_identifiers(tmp_path: Path) -> None: + dataset = tmp_path / "duplicate_ids.tsv" + dataset.write_text( + "chembl_id\tcanonical_smiles\nDUP\tCCO\nDUP\tCCC\n", encoding="utf-8" + ) + + with pytest.raises(ValueError, match="duplicate chembl_id"): + build_similarity_graph(dataset, threshold=0.1, block_size=1) + + +def test_graph_writer_creates_missing_output_directory(tmp_path: Path) -> None: + output = tmp_path / "nested" / "edges.csv" + write_graph_edges(output, []) + assert output.read_text(encoding="utf-8").startswith("source_id,target_id") diff --git a/tests/test_similarity_search.py b/tests/test_similarity_search.py index 4af1c62..66eb5e2 100644 --- a/tests/test_similarity_search.py +++ b/tests/test_similarity_search.py @@ -6,7 +6,11 @@ from rdkit import Chem, DataStructs from scimesh.chemistry.dataset import DatasetStats, find_molecule_by_id, iter_valid_molecules from scimesh.chemistry.fingerprints import fingerprint -from scimesh.workloads.similarity_search import SimilarityMatch, search_similar +from scimesh.workloads.similarity_search import ( + SimilarityMatch, + search_similar, + write_search_results, +) def test_search_matches_full_sorting_and_skips_query_and_invalid( @@ -56,3 +60,9 @@ def test_search_can_rank_and_filter_least_similar_molecules( assert result.matches == sorted( result.matches, key=lambda match: match.sort_key("less") ) + + +def test_search_writer_creates_missing_output_directory(tmp_path: Path) -> None: + output = tmp_path / "nested" / "results.csv" + write_search_results(output, []) + assert output.read_text(encoding="utf-8").startswith("rank,chembl_id") diff --git a/tests/test_worker_daemon.py b/tests/test_worker_daemon.py index 583ee93..bd2b7af 100644 --- a/tests/test_worker_daemon.py +++ b/tests/test_worker_daemon.py @@ -11,9 +11,17 @@ import pytest from scimesh.worker.config import WorkerConfig from scimesh.worker.coordinator import CoordinatorTransientError from scimesh.worker.daemon import LeaseHeartbeat, WorkerDaemon -from scimesh.worker.models import ClaimedTask, InputArtifact, ProducedArtifact, RunResult +from scimesh.worker.models import ( + ClaimedTask, + InputArtifact, + ProducedArtifact, + RegisteredWorker, + RunResult, + UploadedArtifact, +) from scimesh.worker.artifacts import HttpArtifactClient, _SameOriginAuthRedirectHandler, _origin from scimesh.worker.runners import SciMeshRunner +from scimesh.worker.transport import NoRedirectHandler class FakeCoordinator: @@ -24,6 +32,11 @@ class FakeCoordinator: task, self.task = self.task, None return task + def register( + self, name: str, capabilities: tuple[str, ...], cpu_count: int, memory_mb: int | None + ) -> RegisteredWorker: + return RegisteredWorker("11111111-1111-4111-8111-111111111111", 15) + def submit(self, task: ClaimedTask, payload: dict) -> None: self.submissions.append(payload) @@ -42,9 +55,17 @@ class FakeArtifacts: def download(self, uri: str, destination: Path) -> None: destination.write_bytes(self.content) - def upload(self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact) -> str: + def upload( + self, task: ClaimedTask, worker_id: str, artifact: ProducedArtifact + ) -> UploadedArtifact: self.uploaded.append((task.task_id, worker_id, artifact.path)) - return f"https://example.test/tasks/{task.task_id}/artifacts/{artifact.path.name}" + content = artifact.path.read_bytes() + return UploadedArtifact( + "22222222-2222-4222-8222-222222222222", + f"https://example.test/tasks/{task.task_id}/artifacts/{artifact.path.name}", + hashlib.sha256(content).hexdigest(), + len(content), + ) class FakeRunner: def __init__(self) -> None: @@ -75,8 +96,9 @@ def test_claims_runs_uploads_and_submits_csv(tmp_path: Path) -> None: assert runner.calls == 1 assert len(artifacts.uploaded) == 1 assert coordinator.heartbeats == [("task-1", 1, "worker-1")] - assert coordinator.submissions[0]["status"] == "completed" + assert "status" not in coordinator.submissions[0] assert coordinator.submissions[0]["result"]["content_type"] == "text/csv" + assert coordinator.submissions[0]["result"]["artifact_id"] == "22222222-2222-4222-8222-222222222222" assert coordinator.submissions[0]["result"]["uri"].startswith("https://example.test/tasks/task-1/artifacts/") @@ -95,6 +117,14 @@ def test_bad_checksum_reports_failure_without_running(tmp_path: Path) -> None: assert not coordinator.submissions +def test_directory_creation_failure_is_reported(tmp_path: Path) -> None: + content = b"input fixture" + worker, coordinator, _, _, config = daemon(tmp_path, make_task(content), content) + (config.work_dir / "task-1" / "1").mkdir(parents=True) + assert worker.run_once() is True + assert coordinator.failures[0]["error_code"] == "FileExistsError" + + 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: @@ -133,6 +163,12 @@ def test_redirect_to_external_storage_strips_authorization() -> 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 + + def test_lease_is_renewed_while_a_runner_is_still_working(tmp_path: Path) -> None: content = b"input fixture" worker, coordinator, _, _, config = daemon(tmp_path, make_task(content), content) @@ -184,3 +220,50 @@ def test_runner_maps_graph_and_smiles_search_parameters(tmp_path: Path, monkeypa 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] + + +def test_claimed_task_rejects_path_traversal_and_invalid_metadata() -> None: + payload = { + "task_id": "../outside", + "attempt": 1, + "lease_expires_at": "2026-07-30T00:00:00Z", + "workload": "similarity-search", + "input": {"uri": "https://example.test/input", "sha256": "a" * 64}, + "parameters": {}, + } + with pytest.raises(ValueError, match="invalid claimed-task response"): + ClaimedTask.from_json(payload) + + +def test_uploaded_artifact_requires_complete_durable_metadata() -> None: + artifact = UploadedArtifact.from_json( + { + "artifact_id": "22222222-2222-4222-8222-222222222222", + "uri": "https://coordinator.example/artifacts/222/download", + "sha256": "a" * 64, + "size_bytes": 12, + } + ) + assert artifact.size_bytes == 12 + with pytest.raises(ValueError, match="artifact size_bytes"): + UploadedArtifact.from_json({"artifact_id": "missing"}) + + +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( + { + "coordinator_url": "https://coordinator.example", + "work_dir": tmp_path, + "worker_name": "test-worker", + } + ) + assert config.coordinator_url == "https://coordinator.example" + assert config.worker_id is None + + +def test_worker_registration_sets_returned_identity(tmp_path: Path) -> None: + worker, _, _, _, _ = daemon(tmp_path, None, b"") + worker._register_worker() + assert worker.worker_id == "11111111-1111-4111-8111-111111111111" + assert worker.config.heartbeat_interval == 15