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
View File
@@ -19,6 +19,7 @@ from .artifacts import (
PortSpec,
Provenance,
)
from .batch import MapReduceWorkload
from .conformance import (
CancellationFlag,
LocalArtifactStore,
@@ -77,6 +78,7 @@ from .registry import (
WorkloadDefinition,
WorkloadDescription,
WorkloadRegistry,
workload_allowlist_from_json,
)
from .resources import (
AcceleratorDevice,
@@ -155,6 +157,7 @@ __all__ = [
"LocalTaskContext",
"LoopSpec",
"MANIFEST_SCHEMA_VERSION",
"MapReduceWorkload",
"NegotiatedWorkload",
"NetworkPolicy",
"NumericTolerance",
@@ -207,6 +210,7 @@ __all__ = [
"WorkloadLimits",
"WorkloadManifest",
"WorkloadRegistry",
"workload_allowlist_from_json",
"assert_manifest_round_trip",
"installed_distribution_digest",
"negotiate_manifest",
+554
View File
@@ -0,0 +1,554 @@
"""High-level core-batch-v1 authoring scaffold for map/reduce workloads.
``MapReduceWorkload`` is the SDK's primary authoring surface for the
``core-batch-v1`` profile: a subclass declares its identity, parameter
schema, artifact ports, and three scientific hooks (partitioning, per-shard
computation, partial merging), and the base class assembles the immutable
manifest, the map/reduce stages, the workflow DAG, the digest-pinned
planner/runner/reducer handlers, and the exact-artifact verifier.
The base class deliberately supports only the static byte-exact map/reduce
shape with a single external input and one output port. Workloads that need
a different DAG (for example block-pair tasks with two map inputs) override
the ``map_stage_inputs``, ``plan_tasks``, ``parse_partial_key``, and
``validate_partial_keys`` hooks; anything outside the model must use the
lower-level SDK value objects directly.
"""
from __future__ import annotations
import hashlib
import shutil
from pathlib import Path
from typing import Any, Mapping, Sequence
from .artifacts import (
ArtifactCollection,
ArtifactItem,
ArtifactRef,
ArtifactSchema,
Cardinality,
CollectionKind,
OutputManifest,
PortSpec,
)
from .execution import (
CheckpointPolicy,
ExecutionProfile,
NetworkPolicy,
RetryPolicy,
)
from .identity import ComponentRef, VersionRange, WorkloadId
from .manifest import (
DeterminismProfile,
EnvironmentSpec,
PackageSpec,
TrustMode,
VerifierSpec,
WorkloadLimits,
WorkloadManifest,
)
from .plans import JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
from .protocols import PlanningContext, ReduceContext, TaskContext
from .registry import WorkloadDefinition
from .resources import ResourceRequirements
from .verification import ExactArtifactVerifier
from .workflow import ArtifactEdge, PortRef, StageKind, StageSpec, WorkflowSpec
_EXACT_ARTIFACT = ComponentRef("exact-artifact", 1)
_EXACT_VERIFIER = ExactArtifactVerifier()
def _sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as source:
for block in iter(lambda: source.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def _default_entry_point(module: str, kind: str) -> str:
return f"{module}:{kind}@v1"
class MapReduceWorkload:
"""Base class for static byte-exact map/reduce workloads.
Required class attributes:
- ``workload_id``: ``WorkloadId`` identity (name + version);
- ``parameters_schema``: strict JSON object schema (the registry validates
``additionalProperties: false`` for the top level);
- ``input_port``, ``partial_port``, ``output_port``: typed ``PortSpec``
values for the external input, one map partial, and the final result.
Scientific hooks (override as needed):
- ``domain_validate(parameters)``: extra job-parameter validation;
- ``resolved_parameters(request)``: values carried into the plan;
- ``partition_input(input_path, parameters, workspace)``: deterministic
shard splitting; returns one TSV/CSV file per map task;
- ``plan_tasks(shard_paths, resolved, job, negotiated, map_stage, context)``:
task construction (default: one task per shard, ``map/<index>``);
- ``compute_shard(inputs, parameters, output_path)``: one map task;
``inputs`` maps every ``map_stage_inputs`` port to a materialized file;
- ``parse_partial_key(key)`` and ``validate_partial_keys(parsed)``:
partial-key policy for the reducer;
- ``reduce_partials(partial_paths, parameters, output_path)``: the
deterministic merge.
Optional class attributes: ``map_stage_inputs`` (default one ``input``
port), ``map_parameter_names``, ``reduce_parameter_names``, ``capabilities``,
``trust_modes``, ``workflow_id``, ``limits``, ``resources``, ``execution``,
``map_entry_point``, ``reduce_entry_point``.
"""
workload_id: WorkloadId
description: str = ""
parameters_schema: Mapping[str, Any] = {}
input_port: PortSpec
partial_port: PortSpec
output_port: PortSpec
map_stage_inputs: Mapping[str, PortSpec] | None = None
map_parameter_names: tuple[str, ...] = ()
reduce_parameter_names: tuple[str, ...] = ()
capabilities: tuple[str, ...] | None = None
trust_modes: tuple[TrustMode, ...] = (TrustMode.TRUSTED, TrustMode.UNTRUSTED_QUORUM)
workflow_id: str | None = None
limits: WorkloadLimits | None = None
resources: ResourceRequirements | None = None
execution: ExecutionProfile | None = None
map_entry_point: str | None = None
reduce_entry_point: str | None = None
def __init__(
self,
*,
package_digest: str,
environment_digest: str,
) -> None:
workload_id = self.workload_id
if not isinstance(workload_id, WorkloadId):
raise ValueError("workload_id must be a WorkloadId")
if not isinstance(self.input_port, PortSpec):
raise ValueError("input_port must be a PortSpec")
if not isinstance(self.partial_port, PortSpec):
raise ValueError("partial_port must be a PortSpec")
if not isinstance(self.output_port, PortSpec):
raise ValueError("output_port must be a PortSpec")
if not self.parameters_schema:
raise ValueError("parameters_schema must be provided")
module = type(self).__module__
self.map_entry_point = self.map_entry_point or _default_entry_point(
module, "map"
)
self.reduce_entry_point = self.reduce_entry_point or _default_entry_point(
module, "reduce"
)
self.entry_point = self.map_entry_point
self.map_stage_inputs = dict(
self.map_stage_inputs or {"input": self.input_port}
)
if set(self.map_stage_inputs) != {"input"} and any(
port.schema != self.input_port.schema
for port in self.map_stage_inputs.values()
):
raise ValueError(
"additional map inputs must share the external input schema"
)
self.map_parameter_names = tuple(self.map_parameter_names)
self.reduce_parameter_names = tuple(
self.reduce_parameter_names or self.map_parameter_names
)
self.capabilities = tuple(self.capabilities or (workload_id.name,))
if workload_id.name not in self.capabilities:
raise ValueError("capabilities must include the workload name")
name = workload_id.name
resources = self.resources or ResourceRequirements(
profile=f"{name}-cpu-v1",
cpu_cores=1,
memory_mb=1024,
scratch_mb=1024,
max_duration_seconds=3600,
)
execution = self.execution or ExecutionProfile(
profile=f"{name}-python-process-v1",
network=NetworkPolicy.TRUSTED,
timeout_seconds=3600,
checkpoint=CheckpointPolicy(),
)
limits = self.limits or WorkloadLimits(
max_input_bytes=self.input_port.schema.max_bytes,
max_tasks=10_000,
max_output_bytes=self.output_port.schema.max_bytes,
)
trust_values = tuple(mode.value for mode in self.trust_modes)
map_stage = StageSpec(
stage_id="map",
kind=StageKind.MAP,
entry_point=self.map_entry_point,
needs=(),
inputs=self.map_stage_inputs,
outputs={"partial": self.partial_port},
parameter_names=self.map_parameter_names,
resources=resources,
execution=execution,
retry=RetryPolicy(),
verifier=_EXACT_ARTIFACT,
trust_modes=trust_values,
max_fan_out=limits.max_tasks,
cacheable=True,
)
reduce_input = PortSpec(
schema=self.partial_port.schema,
cardinality=Cardinality.MANY,
collection=CollectionKind.KEYED,
)
reduce_stage = StageSpec(
stage_id="reduce",
kind=StageKind.REDUCE,
entry_point=self.reduce_entry_point,
needs=("map",),
inputs={"partials": reduce_input},
outputs={"result": self.output_port},
parameter_names=self.reduce_parameter_names,
resources=resources,
execution=execution,
retry=RetryPolicy(),
verifier=_EXACT_ARTIFACT,
trust_modes=trust_values,
max_fan_out=1,
cacheable=True,
)
edges = [
ArtifactEdge(PortRef("input"), PortRef(port, "map"))
for port in self.map_stage_inputs
] + [
ArtifactEdge(PortRef("partial", "map"), PortRef("partials", "reduce")),
]
workflow = WorkflowSpec(
workflow_id=self.workflow_id or f"{name}-map-reduce-v1",
inputs={"input": self.input_port},
stages=(map_stage, reduce_stage),
edges=tuple(edges),
outputs={"result": PortRef("result", "reduce")},
max_tasks=limits.max_tasks,
max_output_bytes=limits.max_output_bytes,
)
self.manifest = WorkloadManifest(
sdk_api=VersionRange(">=1.0,<2.0"),
protocol=VersionRange(">=1,<2"),
workload=workload_id,
description=self.description,
package=PackageSpec("scimesh", package_digest),
environment=EnvironmentSpec(
"python-process",
environment_digest,
{"adapter": "sdk-native"},
),
parameters_schema=self.parameters_schema,
workflow=workflow,
inputs={"input": self.input_port},
outputs={"result": self.output_port},
determinism=DeterminismProfile.BYTE_EXACT,
trust_modes=self.trust_modes,
verifier=VerifierSpec(_EXACT_ARTIFACT, {}),
limits=limits,
capabilities=self.capabilities,
conformance_profiles=("core-batch-v1",),
)
self._exact_verifier = _EXACT_VERIFIER
self._resources = resources
self._execution = execution
self._limits = limits
# ------------------------------------------------------------------
# Public assembly
# ------------------------------------------------------------------
def definition(self) -> WorkloadDefinition:
map_entry_point = self.map_entry_point
reduce_entry_point = self.reduce_entry_point
assert map_entry_point is not None and reduce_entry_point is not None
return WorkloadDefinition(
manifest=self.manifest,
planner=self,
runners={map_entry_point: self},
reducers={reduce_entry_point: self},
verifiers={_EXACT_ARTIFACT.canonical: self._exact_verifier},
)
# ------------------------------------------------------------------
# Scientific hooks
# ------------------------------------------------------------------
def domain_validate(self, parameters: Mapping[str, Any]) -> None:
"""Extra job-parameter validation beyond the JSON schema."""
def resolved_parameters(self, request: JobRequest) -> dict[str, Any]:
"""Values persisted as the plan's resolved parameters."""
return dict(request.parameters)
def partition_input(
self,
input_path: Path,
parameters: Mapping[str, Any],
workspace: Path,
) -> list[Path]:
"""Split the materialized input into deterministic shard files."""
raise NotImplementedError("partition_input must be implemented")
def plan_tasks(
self,
shard_paths: Sequence[Path],
resolved: Mapping[str, Any],
job: ValidatedJob,
negotiated: Any,
map_stage: StageSpec,
context: PlanningContext,
) -> list[TaskSpec]:
"""Build one map TaskSpec per planned shard."""
task_parameters = self.task_parameters(resolved)
tasks: list[TaskSpec] = []
for index, path in enumerate(shard_paths):
sealed = context.sink.seal(
path,
declaration=self.input_port.schema,
)
tasks.append(
self.task_spec(
map_stage,
job,
negotiated,
f"map/{index:08d}",
task_parameters,
{"input": ArtifactCollection.single(sealed)},
)
)
return tasks
def task_parameters(self, resolved: Mapping[str, Any]) -> dict[str, Any]:
"""Project resolved parameters onto the map stage projection."""
return {
key: value
for key, value in resolved.items()
if key in set(self.map_parameter_names)
}
def compute_shard(
self,
inputs: Mapping[str, Path],
parameters: Mapping[str, Any],
output_path: Path,
) -> Mapping[str, int | float]:
"""Compute one map task and write its partial CSV."""
raise NotImplementedError("compute_shard must be implemented")
def parse_partial_key(self, key: str) -> Any:
"""Parse one ``map.<key>`` partial key into an orderable identity."""
prefix = "map."
if not key.startswith(prefix):
raise ValueError("partial key must use map.<eight-digit-index>")
raw = key[len(prefix) :]
if len(raw) != 8 or not raw.isdigit():
raise ValueError("partial key must use map.<eight-digit-index>")
return int(raw)
def validate_partial_keys(self, parsed: Sequence[Any]) -> None:
"""Enforce the reducer's partial-set invariant (default: contiguous)."""
indices = sorted(int(value) for value in parsed)
if indices != list(range(len(indices))):
raise ValueError("partial keys must be complete and contiguous")
def reduce_partials(
self,
partial_paths: Sequence[Path],
parameters: Mapping[str, Any],
output_path: Path,
) -> Mapping[str, int | float]:
"""Merge materialized partials into one deterministic final CSV."""
raise NotImplementedError("reduce_partials must be implemented")
# ------------------------------------------------------------------
# Framework handlers
# ------------------------------------------------------------------
def validate(self, request: JobRequest) -> ValidatedJob:
if request.workload != self.manifest.workload:
raise ValueError("workload received a request for another workload")
self.domain_validate(request.parameters)
return ValidatedJob(request, self.resolved_parameters(request))
def plan(self, job: ValidatedJob, context: PlanningContext) -> WorkflowPlan:
if not isinstance(job, ValidatedJob):
raise ValueError("job must be a ValidatedJob")
collection = job.request.inputs.get("input")
if collection is None:
raise ValueError("workload requires the input port")
self.input_port.validate_collection(collection, "job input")
input_path = context.catalog.materialize(collection.items[0].artifact)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
resolved = dict(job.resolved_parameters)
resolved = self.resolved_parameters_for_plan(job, input_path, resolved)
shard_paths = self.partition_input(input_path, resolved, workspace)
negotiated = context.negotiated
map_stage = self.manifest.workflow.stages[0]
tasks = self.plan_tasks(
shard_paths, resolved, job, negotiated, map_stage, context
)
return self.workflow_plan(job, negotiated, resolved, tasks)
def resolved_parameters_for_plan(
self,
job: ValidatedJob,
input_path: Path,
resolved: dict[str, Any],
) -> dict[str, Any]:
"""Hook to enrich resolved parameters at plan time (query resolution)."""
return resolved
def workflow_plan(
self,
job: ValidatedJob,
negotiated: Any,
resolved: Mapping[str, Any],
tasks: Sequence[TaskSpec],
) -> WorkflowPlan:
return WorkflowPlan(
workload=self.manifest.workload,
package_digest=self.manifest.package.digest,
manifest_digest=self.manifest.digest,
trust_mode=job.request.trust_mode,
sdk_api_version=negotiated.sdk_api_version,
protocol_version=negotiated.protocol_version,
manifest_schema_version=self.manifest.manifest_schema_version,
workflow_schema_version=self.manifest.workflow.schema_version,
environment_digest=self.manifest.environment.digest,
verifier=self.manifest.verifier.verifier,
selected_features=negotiated.selected_features,
optional_fallbacks=negotiated.optional_fallbacks,
workflow_id=self.manifest.workflow.workflow_id,
resolved_parameters=dict(resolved),
tasks=tuple(tasks),
)
def task_spec(
self,
map_stage: StageSpec,
job: ValidatedJob,
negotiated: Any,
task_key: str,
parameters: Mapping[str, Any],
inputs: Mapping[str, ArtifactCollection],
) -> TaskSpec:
assert map_stage.verifier is not None
return TaskSpec(
workload=self.manifest.workload,
package_digest=self.manifest.package.digest,
manifest_digest=self.manifest.digest,
trust_mode=job.request.trust_mode,
sdk_api_version=negotiated.sdk_api_version,
protocol_version=negotiated.protocol_version,
manifest_schema_version=self.manifest.manifest_schema_version,
workflow_schema_version=self.manifest.workflow.schema_version,
environment_digest=self.manifest.environment.digest,
verifier=map_stage.verifier,
selected_features=negotiated.selected_features,
optional_fallbacks=negotiated.optional_fallbacks,
task_key=task_key,
stage_id=map_stage.stage_id,
parameters=parameters,
inputs=inputs,
expected_outputs=map_stage.outputs,
resources=map_stage.resources,
execution=map_stage.execution,
).validate_stage(map_stage)
def run(self, context: TaskContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
inputs: dict[str, Path] = {}
for name, port in self.map_stage_inputs.items():
collection = context.task.inputs.get(name)
if collection is None:
raise ValueError(f"map task requires the {name} input")
port.validate_collection(collection, f"map input {name}")
assert collection.items
inputs[name] = context.catalog.materialize(collection.items[0].artifact)
output_path = workspace / "result.csv"
metrics = self.compute_shard(
inputs,
context.task.parameters,
output_path,
)
context.cancellation.raise_if_cancelled()
sealed = context.sink.seal(
output_path,
declaration=self.partial_port.schema,
)
return OutputManifest(
context.task.task_key,
{"partial": ArtifactCollection.single(sealed)},
metrics,
context.provenance,
).validate_against(
context.task.expected_outputs,
max_output_bytes=self._limits.max_output_bytes,
)
def reduce(self, context: ReduceContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
collection = context.accepted_inputs.get("partials")
if collection is None:
raise ValueError("reducer requires a non-empty keyed partial collection")
if collection.kind is not CollectionKind.KEYED or not collection.items:
raise ValueError("reducer requires a non-empty keyed partial collection")
self.manifest.workflow.stages[1].inputs["partials"].validate_collection(
collection,
"reducer partials",
)
parsed = [self.parse_partial_key(item.key or "") for item in collection.items]
self.validate_partial_keys(parsed)
expected_keys = context.task.expected_input_keys.get("partials")
if expected_keys is None or {item.key for item in collection.items} != set(
expected_keys
):
raise ValueError("partial keys do not match the coordinator expected set")
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
partial_paths: list[Path] = []
for index, item in sorted(
enumerate(collection.items),
key=lambda entry: entry[1].key or "",
):
artifact: ArtifactRef = item.artifact
source = context.catalog.materialize(artifact)
target = workspace / artifact.artifact_id
if source.resolve() != target.resolve():
shutil.copyfile(source, target)
if _sha256_file(target) != artifact.sha256:
raise ValueError("materialized partial checksum does not match")
partial_paths.append(target)
result_path = workspace / "result.csv"
metrics = self.reduce_partials(
partial_paths,
context.task.parameters,
result_path,
)
context.cancellation.raise_if_cancelled()
sealed = context.sink.seal(
result_path,
declaration=self.output_port.schema,
)
return OutputManifest(
context.task.task_key,
{"result": ArtifactCollection.single(sealed)},
metrics,
context.provenance,
).validate_against(
context.task.expected_outputs,
max_output_bytes=self._limits.max_output_bytes,
)
+216 -56
View File
@@ -24,8 +24,20 @@ from .identity import ComponentRef, WorkloadId
from .integrity import installed_distribution_digest
from .manifest import WorkloadManifest
from .plans import JobRequest, ValidatedJob, WorkflowPlan
from .protocols import Planner, PlanningContext, PlanningResources, Reducer, Runner, Verifier
from .runtime import CompatibilityError, NegotiatedWorkload, RuntimeCapabilities, negotiate_manifest
from .protocols import (
Planner,
PlanningContext,
PlanningResources,
Reducer,
Runner,
Verifier,
)
from .runtime import (
CompatibilityError,
NegotiatedWorkload,
RuntimeCapabilities,
negotiate_manifest,
)
from .schema import validate_parameter_instance
from .workflow import StageKind
@@ -51,11 +63,11 @@ def _validate_entry_point_ownership(entry_point: metadata.EntryPoint) -> None:
raise ValueError("workload entry point has an invalid module path")
raw_top_level = distribution.read_text("top_level.txt")
declared = {
line.strip()
for line in raw_top_level.splitlines()
if line.strip()
} if raw_top_level is not None else set()
declared = (
{line.strip() for line in raw_top_level.splitlines() if line.strip()}
if raw_top_level is not None
else set()
)
root_name = parts[0]
if root_name not in declared:
raise ValueError("workload entry point module is outside its distribution")
@@ -81,9 +93,18 @@ def _validate_entry_point_ownership(entry_point: metadata.EntryPoint) -> None:
ownership_root = package_root.resolve()
candidates = [
*(Path(str(module_base) + suffix) for suffix in machinery.SOURCE_SUFFIXES),
*(Path(str(module_base) + suffix) for suffix in machinery.EXTENSION_SUFFIXES),
*(module_base / ("__init__" + suffix) for suffix in machinery.SOURCE_SUFFIXES),
*(module_base / ("__init__" + suffix) for suffix in machinery.EXTENSION_SUFFIXES),
*(
Path(str(module_base) + suffix)
for suffix in machinery.EXTENSION_SUFFIXES
),
*(
module_base / ("__init__" + suffix)
for suffix in machinery.SOURCE_SUFFIXES
),
*(
module_base / ("__init__" + suffix)
for suffix in machinery.EXTENSION_SUFFIXES
),
]
else:
if len(parts) != 1:
@@ -92,7 +113,10 @@ def _validate_entry_point_ownership(entry_point: metadata.EntryPoint) -> None:
module_base = Path(distribution.locate_file(root_name))
candidates = [
*(Path(str(module_base) + suffix) for suffix in machinery.SOURCE_SUFFIXES),
*(Path(str(module_base) + suffix) for suffix in machinery.EXTENSION_SUFFIXES),
*(
Path(str(module_base) + suffix)
for suffix in machinery.EXTENSION_SUFFIXES
),
]
existing = tuple(candidate for candidate in candidates if candidate.is_file())
if len(existing) != 1 or not existing[0].resolve().is_relative_to(ownership_root):
@@ -126,7 +150,9 @@ class WorkloadDefinition:
for name, handler in values.items():
canonical = require_string(name, f"{field} entry point", max_length=256)
if not callable(getattr(handler, method, None)):
raise ValueError(f"definition {field} handler must implement {method}")
raise ValueError(
f"definition {field} handler must implement {method}"
)
copied[canonical] = handler
object.__setattr__(self, field, MappingProxyType(copied))
for stage in self.manifest.workflow.stages:
@@ -149,23 +175,39 @@ class WorkloadDefinition:
)
verifier_key = self.manifest.verifier.verifier.canonical
if verifier_key not in self.verifiers:
raise ValueError(f"definition has no installed manifest verifier: {verifier_key}")
raise ValueError(
f"definition has no installed manifest verifier: {verifier_key}"
)
for key, verifier in self.verifiers.items():
try:
declared_identity = ComponentRef.from_dict(key)
except ValueError as error:
raise ValueError("definition verifier keys must be component identities") from error
if declared_identity.canonical != key or getattr(verifier, "identity", None) != declared_identity:
raise ValueError("definition verifier handler identity does not match its key")
raise ValueError(
"definition verifier keys must be component identities"
) from error
if (
declared_identity.canonical != key
or getattr(verifier, "identity", None) != declared_identity
):
raise ValueError(
"definition verifier handler identity does not match its key"
)
manifest_verifier = self.verifiers[verifier_key]
handler_configuration = getattr(manifest_verifier, "configuration", None)
if handler_configuration is None:
if self.manifest.verifier.configuration:
raise ValueError("manifest verifier configuration is not bound by its handler")
raise ValueError(
"manifest verifier configuration is not bound by its handler"
)
elif dict(handler_configuration) != dict(self.manifest.verifier.configuration):
raise ValueError("manifest verifier configuration does not match its handler")
raise ValueError(
"manifest verifier configuration does not match its handler"
)
for stage in self.manifest.workflow.stages:
if stage.verifier is not None and stage.verifier.canonical not in self.verifiers:
if (
stage.verifier is not None
and stage.verifier.canonical not in self.verifiers
):
raise ValueError(
f"definition has no installed stage verifier: {stage.verifier.canonical}"
)
@@ -178,13 +220,54 @@ class AllowedPackage:
digest: str
def __post_init__(self) -> None:
distribution = require_string(self.distribution, "distribution", max_length=128).lower()
distribution = require_string(
self.distribution, "distribution", max_length=128
).lower()
if not re.fullmatch(r"[a-z0-9]+(?:[-_.][a-z0-9]+)*", distribution):
raise ValueError("distribution must be a canonical Python distribution name")
raise ValueError(
"distribution must be a canonical Python distribution name"
)
object.__setattr__(self, "distribution", distribution.replace("_", "-"))
if not isinstance(self.workload, WorkloadId):
raise ValueError("allowed workload must be a WorkloadId")
object.__setattr__(self, "digest", require_sha256(self.digest, "allowed digest", prefixed=True))
object.__setattr__(
self, "digest", require_sha256(self.digest, "allowed digest", prefixed=True)
)
def workload_allowlist_from_json(value: object) -> tuple[AllowedPackage, ...]:
"""Parse a JSON array of ``{distribution, name, version, digest}`` allowlist entries."""
import json
if value is None or value == "":
return ()
if not isinstance(value, str):
raise ValueError("workload allowlist must be a JSON array")
try:
decoded = json.loads(value)
except (TypeError, json.JSONDecodeError, RecursionError) as error:
raise ValueError("workload allowlist must be valid JSON") from error
if not isinstance(decoded, list):
raise ValueError("workload allowlist must be a JSON array")
entries: list[AllowedPackage] = []
for item in decoded:
if not isinstance(item, dict) or not {
"distribution",
"name",
"version",
"digest",
}.issubset(item):
raise ValueError(
"workload allowlist entries need distribution, name, version, and digest"
)
entries.append(
AllowedPackage(
str(item["distribution"]),
WorkloadId(str(item["name"]), str(item["version"])),
str(item["digest"]),
)
)
return tuple(entries)
@dataclass(frozen=True, slots=True)
@@ -223,14 +306,18 @@ class WorkloadRegistry:
self._enabled: set[tuple[str, str, str]] = set()
self._lock = RLock()
def register(self, definition: WorkloadDefinition, *, enabled: bool = False) -> None:
def register(
self, definition: WorkloadDefinition, *, enabled: bool = False
) -> None:
if not isinstance(definition, WorkloadDefinition):
raise ValueError("definition must be a WorkloadDefinition")
workload = definition.manifest.workload
key = (workload.name, workload.version)
with self._lock:
if key in self._definitions:
raise ValueError(f"workload version already registered: {workload.name}@{workload.version}")
raise ValueError(
f"workload version already registered: {workload.name}@{workload.version}"
)
self._definitions[key] = definition
if enabled:
self._enabled.add((*key, definition.manifest.package.digest))
@@ -240,7 +327,9 @@ class WorkloadRegistry:
with self._lock:
definition = self._registered(name, version)
if digest != definition.manifest.package.digest:
raise ValueError("package digest does not match the registered manifest")
raise ValueError(
"package digest does not match the registered manifest"
)
self._enabled.add((definition.manifest.workload.name, version, digest))
def disable(self, name: str, version: str, package_digest: str) -> None:
@@ -257,7 +346,9 @@ class WorkloadRegistry:
try:
return self._definitions[(canonical, version)]
except KeyError as error:
raise ValueError(f"unknown workload version: {canonical}@{version}") from error
raise ValueError(
f"unknown workload version: {canonical}@{version}"
) from error
def require(
self,
@@ -270,10 +361,21 @@ class WorkloadRegistry:
digest = require_sha256(package_digest, "package_digest", prefixed=True)
with self._lock:
definition = self._registered(name, version)
identity = (definition.manifest.workload.name, definition.manifest.workload.version, digest)
if digest != definition.manifest.package.digest or identity not in self._enabled:
identity = (
definition.manifest.workload.name,
definition.manifest.workload.version,
digest,
)
if (
digest != definition.manifest.package.digest
or identity not in self._enabled
):
raise ValueError("workload package digest is not enabled")
negotiated = negotiate_manifest(definition.manifest, runtime) if runtime is not None else None
negotiated = (
negotiate_manifest(definition.manifest, runtime)
if runtime is not None
else None
)
return definition, negotiated
def plan(
@@ -293,31 +395,41 @@ class WorkloadRegistry:
runtime=runtime,
)
assert negotiated is not None
self._validate_request_compatibility(request, definition.manifest, runtime, negotiated)
self._validate_request_compatibility(
request, definition.manifest, runtime, negotiated
)
self._validate_request_shape(request, definition.manifest)
validated = definition.planner.validate(request)
if not isinstance(validated, ValidatedJob) or validated.request != request:
raise ValueError("planner.validate must return a ValidatedJob for the same request")
raise ValueError(
"planner.validate must return a ValidatedJob for the same request"
)
plan = definition.planner.plan(
validated,
_NegotiatedPlanningContext(context, negotiated),
)
if not isinstance(plan, WorkflowPlan) or plan.workload != request.workload:
raise ValueError("planner.plan must return a WorkflowPlan for the requested workload")
raise ValueError(
"planner.plan must return a WorkflowPlan for the requested workload"
)
if (
plan.package_digest != definition.manifest.package.digest
or plan.manifest_digest != definition.manifest.digest
or plan.trust_mode is not request.trust_mode
or plan.sdk_api_version != runtime.sdk_api_version
or plan.protocol_version != runtime.protocol_version
or plan.manifest_schema_version != definition.manifest.manifest_schema_version
or plan.workflow_schema_version != definition.manifest.workflow.schema_version
or plan.manifest_schema_version
!= definition.manifest.manifest_schema_version
or plan.workflow_schema_version
!= definition.manifest.workflow.schema_version
or plan.environment_digest != definition.manifest.environment.digest
or plan.verifier != definition.manifest.verifier.verifier
or plan.selected_features != negotiated.selected_features
or plan.optional_fallbacks != negotiated.optional_fallbacks
):
raise ValueError("planner plan does not carry the selected immutable workload pin")
raise ValueError(
"planner plan does not carry the selected immutable workload pin"
)
plan.validate_workflow(definition.manifest.workflow)
self._validate_plan_limits(request, plan, definition.manifest)
return WorkflowPlan.from_json(plan.to_json())
@@ -369,7 +481,9 @@ class WorkloadRegistry:
)
@staticmethod
def _validate_request_shape(request: JobRequest, manifest: WorkloadManifest) -> None:
def _validate_request_shape(
request: JobRequest, manifest: WorkloadManifest
) -> None:
if set(request.inputs) != set(manifest.inputs):
raise ValueError("job input ports do not match the manifest")
total_bytes = 0
@@ -380,7 +494,9 @@ class WorkloadRegistry:
for item in request.inputs[name].items:
existing = artifact_references.get(item.artifact.artifact_id)
if existing is not None and existing != item.artifact:
raise ValueError("job reuses an artifact ID with conflicting metadata")
raise ValueError(
"job reuses an artifact ID with conflicting metadata"
)
artifact_references[item.artifact.artifact_id] = item.artifact
if total_bytes > manifest.limits.max_input_bytes:
raise ValueError("job inputs exceed the manifest byte limit")
@@ -388,7 +504,15 @@ class WorkloadRegistry:
raise ValueError("job inputs exceed the manifest artifact limit")
import json
from ._validation import thaw_json
if len(json.dumps(thaw_json(request.parameters), allow_nan=False).encode("utf-8")) > manifest.limits.max_parameter_bytes:
if (
len(
json.dumps(thaw_json(request.parameters), allow_nan=False).encode(
"utf-8"
)
)
> manifest.limits.max_parameter_bytes
):
raise ValueError("job parameters exceed the manifest byte limit")
validate_parameter_instance(request.parameters, manifest.parameters_schema)
@@ -408,13 +532,23 @@ class WorkloadRegistry:
for item in collection.items:
existing = references.get(item.artifact.artifact_id)
if existing is not None and existing != item.artifact:
raise ValueError("workflow plan reuses an artifact ID with conflicting metadata")
raise ValueError(
"workflow plan reuses an artifact ID with conflicting metadata"
)
references[item.artifact.artifact_id] = item.artifact
if len(canonical_json(task.parameters).encode("utf-8")) > manifest.limits.max_parameter_bytes:
raise ValueError("planned task parameters exceed the manifest byte limit")
if (
len(canonical_json(task.parameters).encode("utf-8"))
> manifest.limits.max_parameter_bytes
):
raise ValueError(
"planned task parameters exceed the manifest byte limit"
)
if len(references) > manifest.limits.max_artifacts:
raise ValueError("workflow plan exceeds the manifest artifact limit")
if len(canonical_json(plan.resolved_parameters).encode("utf-8")) > manifest.limits.max_parameter_bytes:
if (
len(canonical_json(plan.resolved_parameters).encode("utf-8"))
> manifest.limits.max_parameter_bytes
):
raise ValueError("resolved parameters exceed the manifest byte limit")
def descriptions(self) -> tuple[WorkloadDescription, ...]:
@@ -455,7 +589,10 @@ class WorkloadRegistry:
for key, approval in allowed.items():
if _normalized_distribution_name(key[0]) != distribution:
continue
if entry_point.name != f"{approval.workload.name}@{approval.workload.version}":
if (
entry_point.name
!= f"{approval.workload.name}@{approval.workload.version}"
):
continue
_validate_entry_point_ownership(entry_point)
# Import policy is process-global, so installed discovery is
@@ -463,9 +600,12 @@ class WorkloadRegistry:
# cache prefix prevents pre-existing package pyc files from
# being consumed, while dont_write_bytecode keeps the
# measured source tree unchanged during both load and factory.
with _DISCOVERY_IMPORT_LOCK, TemporaryDirectory(
prefix="scimesh-discovery-cache-"
) as cache_prefix:
with (
_DISCOVERY_IMPORT_LOCK,
TemporaryDirectory(
prefix="scimesh-discovery-cache-"
) as cache_prefix,
):
measured_before = installed_distribution_digest(entry_point.dist)
if measured_before != approval.digest:
raise ValueError(
@@ -479,44 +619,64 @@ class WorkloadRegistry:
loaded = entry_point.load()
definition = (
loaded()
if callable(loaded) and not isinstance(loaded, WorkloadDefinition)
if callable(loaded)
and not isinstance(loaded, WorkloadDefinition)
else loaded
)
finally:
sys.pycache_prefix = previous_cache_prefix
sys.dont_write_bytecode = previous_bytecode_policy
if installed_distribution_digest(entry_point.dist) != measured_before:
if (
installed_distribution_digest(entry_point.dist)
!= measured_before
):
raise ValueError(
"installed package content changed while loading its entry point"
)
if not isinstance(definition, WorkloadDefinition):
raise ValueError("workload entry point must provide a WorkloadDefinition")
raise ValueError(
"workload entry point must provide a WorkloadDefinition"
)
if definition.manifest.workload != approval.workload:
raise ValueError("discovered workload identity does not match its allowlist entry")
raise ValueError(
"discovered workload identity does not match its allowlist entry"
)
if definition.manifest.package.distribution != approval.distribution:
raise ValueError("discovered package identity does not match its allowlist entry")
raise ValueError(
"discovered package identity does not match its allowlist entry"
)
if definition.manifest.package.digest != approval.digest:
raise ValueError("discovered package digest does not match its allowlist entry")
raise ValueError(
"discovered package digest does not match its allowlist entry"
)
if key in discovered:
raise ValueError("multiple installed entry points match one allowlist entry")
raise ValueError(
"multiple installed entry points match one allowlist entry"
)
pending.append(definition)
discovered.add(key)
break
missing = sorted(set(allowed) - discovered)
if missing:
identities = ", ".join(f"{name}@{version}" for _, name, version in missing)
raise ValueError("allowlisted workload entry points were not installed: " + identities)
raise ValueError(
"allowlisted workload entry points were not installed: " + identities
)
pending_keys = [
(definition.manifest.workload.name, definition.manifest.workload.version)
for definition in pending
]
if len(pending_keys) != len(set(pending_keys)):
raise ValueError("multiple allowlisted distributions provide one workload version")
raise ValueError(
"multiple allowlisted distributions provide one workload version"
)
with self._lock:
conflicts = [key for key in pending_keys if key in self._definitions]
if conflicts:
name, version = conflicts[0]
raise ValueError(f"workload version already registered: {name}@{version}")
raise ValueError(
f"workload version already registered: {name}@{version}"
)
definitions = dict(self._definitions)
enabled = set(self._enabled)
for key, definition in zip(pending_keys, pending):