Add MapReduceWorkload scaffold and generic workload execution
This commit is contained in:
@@ -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
@@ -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
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user