Add MapReduceWorkload scaffold and generic workload execution

This commit is contained in:
Emil
2026-08-02 01:02:24 +03:00
parent 19fbb8e926
commit bc76f386e5
21 changed files with 2241 additions and 1176 deletions
+4 -2
View File
@@ -77,7 +77,9 @@ def main(argv: list[str] | None = None) -> int:
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")
logging.basicConfig(
level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s"
)
# One shared token strategy backs both clients: a worker key (exchanged and
# refreshed) or a static bearer token, decided by what the config carries.
tokens = provider_from_config(
@@ -95,7 +97,7 @@ def main(argv: list[str] | None = None) -> int:
HttpArtifactClient(
config.coordinator_url, config.request_timeout, token_provider=tokens
),
SciMeshRunner(),
SciMeshRunner.for_worker(config),
).run_forever()
return 0 if completed_without_interruption else 130
+78 -13
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
import json
from dataclasses import dataclass
from math import isfinite
from pathlib import Path
@@ -10,6 +11,8 @@ import socket
from typing import Mapping
from urllib.parse import urlsplit
from scimesh.sdk.registry import AllowedPackage, workload_allowlist_from_json
def _clean_url(value: object | None) -> str | None:
"""Normalise an optional URL: drop a blank one, strip a trailing slash."""
@@ -31,6 +34,24 @@ def _positive_number(value: object, name: str, *, allow_zero: bool = False) -> N
raise ValueError(f"{name} must be {qualifier}")
def _capabilities(value: object) -> tuple[str, ...]:
"""Parse a comma-separated capability list into unique non-empty names."""
if value is None:
return ("similarity-search", "similarity_search")
if not isinstance(value, str) or not value.strip():
raise ValueError("capabilities must be a comma-separated list")
names = tuple(
dict.fromkeys(item.strip() for item in value.split(",") if item.strip())
)
if not names:
raise ValueError("capabilities cannot be empty")
return names
def _workload_allowlist(value: object) -> tuple[AllowedPackage, ...]:
return workload_allowlist_from_json(value)
@dataclass(frozen=True)
class WorkerConfig:
coordinator_url: str
@@ -60,6 +81,11 @@ class WorkerConfig:
"similarity-search",
"similarity_search",
)
# Optional allowlist of installed SDK workload packages to execute. When
# empty, the worker runs the built-in similarity-search only. Entries are
# ``{distribution, name, version, digest}`` JSON objects matching the
# installed ``scimesh.workloads`` entry points.
workload_allowlist: tuple[AllowedPackage, ...] = ()
def __post_init__(self) -> None:
parsed = urlsplit(self.coordinator_url)
@@ -72,8 +98,14 @@ class WorkerConfig:
if us.scheme not in {"http", "https"} or not us.hostname:
raise ValueError("userservice_url must be an absolute HTTP(S) URL")
if self.worker_key is not None and not self.userservice_url:
raise ValueError("worker_key requires userservice_url (SCIMESH_USERSERVICE_URL)")
if isinstance(self.cpu_count, bool) or not isinstance(self.cpu_count, int) or self.cpu_count < 1:
raise ValueError(
"worker_key requires userservice_url (SCIMESH_USERSERVICE_URL)"
)
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")
@@ -87,7 +119,9 @@ class WorkerConfig:
_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)
_positive_number(
self.cleanup_after_seconds, "cleanup_after_seconds", allow_zero=True
)
if self.max_tasks is not None:
if (
isinstance(self.max_tasks, bool)
@@ -99,6 +133,18 @@ class WorkerConfig:
raise ValueError("exit_when_idle must be a boolean")
if not self.capabilities:
raise ValueError("capabilities cannot be empty")
if any(
not isinstance(capability, str) or not capability.strip()
for capability in self.capabilities
):
raise ValueError("capabilities must contain non-empty names")
if len(self.capabilities) != len(set(self.capabilities)):
raise ValueError("capabilities must be unique")
if any(
not isinstance(package, AllowedPackage)
for package in self.workload_allowlist
):
raise ValueError("workload_allowlist must contain AllowedPackage values")
# Runner subprocesses use a task directory as their cwd. Keep the
# configured root absolute so input/output paths remain valid there
# even when the CLI received a convenient relative --work-dir value.
@@ -111,7 +157,9 @@ class 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:
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)
@@ -122,20 +170,37 @@ class WorkerConfig:
cpu_count = value("cpu_count", "SCIMESH_CPU_COUNT", os.cpu_count() or 1)
memory_mb = value("memory_mb", "SCIMESH_MEMORY_MB")
max_tasks = value("max_tasks", "SCIMESH_MAX_TASKS")
capabilities = value("capabilities", "SCIMESH_CAPABILITIES")
allowlist = value("workload_allowlist", "SCIMESH_WORKLOAD_ALLOWLIST")
worker_id = value("worker_id", "SCIMESH_WORKER_ID")
work_dir = value("work_dir", "SCIMESH_WORK_DIR", "./scimesh-worker-data")
worker_name = value("worker_name", "SCIMESH_WORKER_NAME", socket.gethostname())
poll_interval = value("poll_interval", "SCIMESH_POLL_INTERVAL", "2")
request_timeout = value("request_timeout", "SCIMESH_REQUEST_TIMEOUT", "30")
heartbeat_interval = value(
"heartbeat_interval", "SCIMESH_HEARTBEAT_INTERVAL", "15"
)
bearer_token = value("bearer_token", "SCIMESH_BEARER_TOKEN")
worker_key = value("worker_key", "SCIMESH_WORKER_KEY")
userservice_url = _clean_url(
value("userservice_url", "SCIMESH_USERSERVICE_URL")
)
return cls(
coordinator_url=url.rstrip("/"),
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())),
worker_id=str(worker_id) if worker_id is not None else None,
work_dir=Path(str(work_dir)),
worker_name=str(worker_name),
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"),
worker_key=value("worker_key", "SCIMESH_WORKER_KEY"),
userservice_url=_clean_url(value("userservice_url", "SCIMESH_USERSERVICE_URL")),
poll_interval=float(poll_interval),
request_timeout=float(request_timeout),
heartbeat_interval=float(heartbeat_interval),
bearer_token=str(bearer_token) if bearer_token is not None else None,
worker_key=str(worker_key) if worker_key is not None else None,
userservice_url=userservice_url,
cleanup_after_seconds=float(cleanup) if cleanup else None,
max_tasks=int(max_tasks) if max_tasks is not None else None,
exit_when_idle=bool(values.get("exit_when_idle", False)),
capabilities=_capabilities(capabilities),
workload_allowlist=_workload_allowlist(allowlist),
)
+91 -47
View File
@@ -3,13 +3,20 @@
The runner is a v1-wire bridge: the coordinator still claims flat tasks and
the worker still uploads one partial CSV, but execution goes through the
SDK-built workload's own Runner handler with a real ``TaskSpec``,
provenance, resource reservation, and a content-addressed local store. No
legacy distributed-protocol code is involved.
provenance, resource reservation, and a content-addressed local store.
The runner is workload-generic: it loads definitions by name (from an
explicit mapping, built-in defaults, or installed-package discovery through
an administrator allowlist) and executes any workload whose map stage has a
single ``input`` port and a single ``partial`` output. Anything else fails
closed with a clear message, so adding a workload never requires touching
worker code.
"""
from __future__ import annotations
import hashlib
import platform
from datetime import datetime, timezone
from pathlib import Path
from typing import Mapping, Protocol
@@ -28,28 +35,21 @@ from scimesh.sdk.conformance import (
ScopedArtifactSink,
)
from scimesh.sdk._validation import canonical_json
from scimesh.sdk.identity import SDK_API_VERSION
from scimesh.sdk.manifest import TrustMode
from scimesh.sdk.plans import TaskSpec
from scimesh.sdk.registry import WorkloadDefinition
from scimesh.sdk.resources import ResourceAllocation, ResourcePool
from scimesh.sdk.registry import WorkloadDefinition, WorkloadRegistry
from scimesh.sdk.resources import ResourceAllocation, ResourceInventory, ResourcePool
from scimesh.sdk.runtime import (
NegotiatedWorkload,
RuntimeCapabilities,
negotiate_manifest,
)
from scimesh.sdk.workflow import StageKind
from scimesh.workloads.library import default_sdk_runtime
from scimesh.workloads.search import similarity_search_sdk_definition
from .config import WorkerConfig
from .models import ClaimedTask, ProducedArtifact, RunResult
#: Parameters the worker may hand to a map task. ``max_rows`` is a plan-time
#: option applied before sharding and is intentionally rejected here.
_RUNNER_PARAMETERS = frozenset(
{"query_smiles", "top_k", "threshold", "threshold_direction", "progress_every"}
)
def _utc_now() -> str:
return datetime.now(timezone.utc).isoformat().replace("+00:00", "Z")
@@ -58,22 +58,93 @@ class Runner(Protocol):
def run(self, task: ClaimedTask, task_dir: Path) -> RunResult: ...
def _inventory_for(
definitions: Mapping[str, WorkloadDefinition],
*,
cpu_cores: int,
memory_mb: int,
) -> ResourceInventory:
return ResourceInventory(
cpu_cores=cpu_cores,
memory_mb=memory_mb,
scratch_mb=memory_mb,
architecture=platform.machine().lower() or "unknown",
environment_digests=tuple(
dict.fromkeys(
definition.manifest.environment.digest
for definition in definitions.values()
)
),
)
def _runtime_for(
definitions: Mapping[str, WorkloadDefinition], inventory: ResourceInventory
) -> RuntimeCapabilities:
return RuntimeCapabilities(
sdk_api_version=SDK_API_VERSION,
protocol_version="1.0.0",
profiles=("core-batch-v1",),
features={"artifact-collections": "1.0.0", "exact-verifier": "1.0.0"},
workload_capabilities=tuple(sorted(definitions)),
inventory=inventory,
)
class SciMeshRunner:
"""Execute claimed coordinator tasks through the SDK-built workloads."""
"""Execute claimed coordinator tasks through SDK-built workloads."""
def __init__(
self,
definitions: Mapping[str, WorkloadDefinition] | None = None,
*,
inventory: ResourceInventory | None = None,
runtime: RuntimeCapabilities | None = None,
) -> None:
self._definitions = dict(definitions or {})
if "similarity-search" not in self._definitions:
from scimesh.workloads.search import similarity_search_sdk_definition
self._definitions["similarity-search"] = (
similarity_search_sdk_definition().definition()
)
self._runtime = runtime or default_sdk_runtime()
self._inventory = inventory or _inventory_for(
self._definitions,
cpu_cores=1,
memory_mb=1024,
)
self._runtime = runtime or _runtime_for(self._definitions, self._inventory)
self._pool = ResourcePool(self._runtime.inventory, max_concurrency=1)
@classmethod
def for_worker(cls, config: WorkerConfig) -> "SciMeshRunner":
"""Build a runner for one worker: discover allowlisted workloads or use built-ins."""
definitions: dict[str, WorkloadDefinition] = {}
if config.workload_allowlist:
registry = WorkloadRegistry()
registry.discover_installed(config.workload_allowlist)
for description in registry.descriptions():
definition, _ = registry.require(
description.workload.name,
description.workload.version,
description.package_digest,
)
definitions[description.workload.name] = definition
if not definitions:
raise ValueError("workload_allowlist discovered no workloads")
else:
from scimesh.workloads.search import similarity_search_sdk_definition
definitions["similarity-search"] = (
similarity_search_sdk_definition().definition()
)
inventory = _inventory_for(
definitions,
cpu_cores=config.cpu_count,
memory_mb=config.memory_mb or 1024,
)
return cls(definitions=definitions, inventory=inventory)
def run(self, task: ClaimedTask, task_dir: Path) -> RunResult:
task_dir = task_dir.resolve()
workload = task.workload.replace("_", "-")
@@ -85,11 +156,14 @@ class SciMeshRunner:
map_stage = next(
stage for stage in manifest.workflow.stages if stage.kind is StageKind.MAP
)
if set(map_stage.inputs) != {"input"} or len(map_stage.outputs) != 1:
raise ValueError(
f"workload {workload} is not executable through the v1 single-input contract"
)
assert map_stage.verifier is not None
input_path = task_dir / "input"
if not input_path.is_file():
raise ValueError("claimed task input is missing")
parameters = self._resolve_parameters(task, input_path)
store = LocalArtifactStore(task_dir / "sdk-store")
input_ref = store.import_file(
input_path,
@@ -110,7 +184,7 @@ class SciMeshRunner:
optional_fallbacks=negotiated.optional_fallbacks,
task_key="map/00000000",
stage_id=map_stage.stage_id,
parameters=parameters,
parameters=task.parameters,
inputs={"input": ArtifactCollection.single(input_ref)},
expected_outputs=map_stage.outputs,
resources=map_stage.resources,
@@ -149,28 +223,6 @@ class SciMeshRunner:
finally:
self._pool.release(allocation.allocation_id)
@staticmethod
def _resolve_parameters(task: ClaimedTask, input_path: Path) -> dict[str, object]:
"""Resolve ``query_id`` once per task and reject plan-time options."""
parameters = dict(task.parameters)
query_id = parameters.get("query_id")
query_smiles = parameters.get("query_smiles")
if isinstance(query_id, str) and not isinstance(query_smiles, str):
from scimesh.chemistry.dataset import find_molecule_by_id
from rdkit import Chem
record = find_molecule_by_id(input_path, query_id)
parameters["query_smiles"] = Chem.MolToSmiles(
record.molecule, canonical=True
)
del parameters["query_id"]
unknown = set(parameters) - _RUNNER_PARAMETERS
if unknown:
raise ValueError(
"unsupported runner parameters: " + ", ".join(sorted(unknown))
)
return parameters
def _provenance(
self,
definition: WorkloadDefinition,
@@ -241,14 +293,6 @@ class SciMeshRunner:
spec.expected_outputs,
max_output_bytes=max_output_bytes,
)
if output.provenance != provenance:
raise ValueError(
"SDK workload output provenance does not match its context"
)
output.validate_against(
spec.expected_outputs,
max_output_bytes=spec.resources.max_duration_seconds, # replaced below
)
sink = context.sink
if not isinstance(sink, ScopedArtifactSink):
raise ValueError("SDK execution requires a scoped artifact sink")