Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e0ee95cbab |
@@ -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`.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
@@ -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
@@ -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, ...]
|
||||
|
||||
@@ -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
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user