Compare commits

..
Author SHA1 Message Date
Emil e0ee95cbab Harden worker contract and transport 2026-07-23 20:49:50 +03:00
15 changed files with 546 additions and 106 deletions
+6 -1
View File
@@ -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`.
+2
View File
@@ -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
+33 -4
View File
@@ -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: `<work-dir>/<task-id>/<attempt>/`.
- 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
+29 -39
View File
@@ -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 {}
+24 -4
View File
@@ -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
+67 -19
View File
@@ -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,
)
+41 -4
View File
@@ -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
+66 -20
View File
@@ -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), "<worker-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)
+101 -6
View File
@@ -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, ...]
+51
View File
@@ -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
+10 -4
View File
@@ -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()
+1
View File
@@ -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,
+17
View File
@@ -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")
+11 -1
View File
@@ -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")
+87 -4
View File
@@ -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