Serve documentation from the operator UI

This commit is contained in:
Emil
2026-08-02 15:50:07 +03:00
parent 284aef5d6f
commit f20cc7fe00
72 changed files with 3509 additions and 698 deletions
+1 -1
View File
@@ -14,7 +14,7 @@ ID_COLUMN = "chembl_id"
SMILES_COLUMN = "canonical_smiles"
# Invalid records are expected in large datasets; suppress one RDKit error per row.
RDLogger.DisableLog("rdApp.error")
RDLogger.DisableLog("rdApp.error") # type: ignore[attr-defined]
@dataclass
+1
View File
@@ -17,3 +17,4 @@ class Workload(Protocol):
def run(self, args: argparse.Namespace) -> int:
"""Execute the workload."""
...
+2 -1
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import json
import math
import re
from enum import Enum
from types import MappingProxyType
from typing import Any, Mapping
from urllib.parse import unquote
@@ -372,7 +373,7 @@ def version_in_range(version: object, expression: str) -> bool:
return True
def enum_value(enum_type: type[Any], value: object, field: str) -> Any:
def enum_value(enum_type: type[Enum], value: object, field: str) -> Any:
try:
return enum_type(value)
except (TypeError, ValueError) as error:
+291 -70
View File
@@ -32,6 +32,13 @@ from .identity import ComponentRef, OUTPUT_SCHEMA_VERSION, SchemaRef, WorkloadId
class CollectionKind(str, Enum):
"""How a set of artifacts is structured on a port.
``SINGLE`` holds at most one unkeyed artifact; ``ORDERED`` and ``SET``
hold unkeyed artifacts with a canonical order; ``KEYED`` holds artifacts
with unique string keys.
"""
SINGLE = "single"
ORDERED = "ordered"
KEYED = "keyed"
@@ -39,6 +46,12 @@ class CollectionKind(str, Enum):
class Cardinality(str, Enum):
"""How many artifacts a port accepts.
``ONE`` requires exactly one artifact, ``OPTIONAL`` at most one, and
``MANY`` at least one (with an ordered, keyed, or set collection).
"""
ONE = "one"
OPTIONAL = "optional"
MANY = "many"
@@ -65,12 +78,22 @@ class ArtifactSchema:
def __post_init__(self) -> None:
if not isinstance(self.ref, SchemaRef):
raise ValueError("artifact schema ref must be a SchemaRef")
object.__setattr__(self, "media_type", require_string(self.media_type, "media_type", max_length=128))
if "/" not in self.media_type or any(character.isspace() for character in self.media_type):
object.__setattr__(
self,
"media_type",
require_string(self.media_type, "media_type", max_length=128),
)
if "/" not in self.media_type or any(
character.isspace() for character in self.media_type
):
raise ValueError("media_type must be a valid type/subtype token")
if self.encoding is not None:
object.__setattr__(self, "encoding", require_identifier(self.encoding, "encoding"))
object.__setattr__(self, "max_bytes", require_positive_int(self.max_bytes, "max_bytes"))
object.__setattr__(
self, "encoding", require_identifier(self.encoding, "encoding")
)
object.__setattr__(
self, "max_bytes", require_positive_int(self.max_bytes, "max_bytes")
)
if not isinstance(self.validator, ComponentRef):
raise ValueError("artifact schema validator must be a ComponentRef")
object.__setattr__(
@@ -83,7 +106,11 @@ class ArtifactSchema:
),
)
if self.max_records is not None:
object.__setattr__(self, "max_records", require_positive_int(self.max_records, "max_records"))
object.__setattr__(
self,
"max_records",
require_positive_int(self.max_records, "max_records"),
)
dimensions = tuple(self.max_dimensions)
if any(
isinstance(value, bool) or not isinstance(value, int) or value < 1
@@ -99,9 +126,19 @@ class ArtifactSchema:
"canonicalizer",
require_identifier(self.canonicalizer, "canonicalizer"),
)
object.__setattr__(self, "privacy_class", require_identifier(self.privacy_class, "privacy_class"))
object.__setattr__(self, "retention_class", require_identifier(self.retention_class, "retention_class"))
if not isinstance(self.streaming, bool) or not isinstance(self.allow_nested_collections, bool):
object.__setattr__(
self,
"privacy_class",
require_identifier(self.privacy_class, "privacy_class"),
)
object.__setattr__(
self,
"retention_class",
require_identifier(self.retention_class, "retention_class"),
)
if not isinstance(self.streaming, bool) or not isinstance(
self.allow_nested_collections, bool
):
raise ValueError("streaming and allow_nested_collections must be booleans")
def to_dict(self) -> dict[str, object]:
@@ -126,10 +163,19 @@ class ArtifactSchema:
if not isinstance(value, Mapping):
raise ValueError("artifact schema must be an object")
fields = {
"ref", "media_type", "encoding", "max_bytes", "validator",
"validator_configuration", "max_records",
"max_dimensions", "streaming", "canonicalizer", "privacy_class",
"retention_class", "allow_nested_collections",
"ref",
"media_type",
"encoding",
"max_bytes",
"validator",
"validator_configuration",
"max_records",
"max_dimensions",
"streaming",
"canonicalizer",
"privacy_class",
"retention_class",
"allow_nested_collections",
}
require_exact_keys(value, fields, "artifact schema")
dimensions = value["max_dimensions"]
@@ -154,6 +200,13 @@ class ArtifactSchema:
@dataclass(frozen=True, slots=True)
class PortSpec:
"""A typed artifact port: schema, cardinality, and collection kind.
``validate_collection`` enforces the declared shape against an
``ArtifactCollection``, including per-artifact schema, media type, byte,
record, and dimension bounds.
"""
schema: ArtifactSchema
cardinality: Cardinality = Cardinality.ONE
collection: CollectionKind = CollectionKind.SINGLE
@@ -161,14 +214,32 @@ class PortSpec:
def __post_init__(self) -> None:
if not isinstance(self.schema, ArtifactSchema):
raise ValueError("port schema must be an ArtifactSchema")
object.__setattr__(self, "cardinality", enum_value(Cardinality, self.cardinality, "cardinality"))
object.__setattr__(self, "collection", enum_value(CollectionKind, self.collection, "collection"))
if self.cardinality is Cardinality.MANY and self.collection is CollectionKind.SINGLE:
raise ValueError("many cardinality requires an ordered, keyed, or set collection")
if self.cardinality is not Cardinality.MANY and self.collection is not CollectionKind.SINGLE:
object.__setattr__(
self,
"cardinality",
enum_value(Cardinality, self.cardinality, "cardinality"),
)
object.__setattr__(
self,
"collection",
enum_value(CollectionKind, self.collection, "collection"),
)
if (
self.cardinality is Cardinality.MANY
and self.collection is CollectionKind.SINGLE
):
raise ValueError(
"many cardinality requires an ordered, keyed, or set collection"
)
if (
self.cardinality is not Cardinality.MANY
and self.collection is not CollectionKind.SINGLE
):
raise ValueError("one and optional cardinality require a single collection")
def validate_collection(self, value: "ArtifactCollection", field: str = "artifact collection") -> None:
def validate_collection(
self, value: "ArtifactCollection", field: str = "artifact collection"
) -> None:
if value.kind is not self.collection:
raise ValueError(f"{field} kind does not match its port declaration")
count = len(value.items)
@@ -183,7 +254,9 @@ class PortSpec:
if artifact.schema != self.schema.ref:
raise ValueError(f"{field} contains an artifact with the wrong schema")
if artifact.media_type != self.schema.media_type:
raise ValueError(f"{field} contains an artifact with the wrong media type")
raise ValueError(
f"{field} contains an artifact with the wrong media type"
)
if artifact.size_bytes > self.schema.max_bytes:
raise ValueError(f"{field} exceeds its per-artifact byte limit")
if self.schema.max_records is not None:
@@ -193,10 +266,14 @@ class PortSpec:
raise ValueError(f"{field} exceeds its record limit")
if self.schema.max_dimensions:
if not artifact.dimensions:
raise ValueError(f"{field} is missing its required dimension summary")
raise ValueError(
f"{field} is missing its required dimension summary"
)
if len(artifact.dimensions) != len(self.schema.max_dimensions) or any(
actual > maximum
for actual, maximum in zip(artifact.dimensions, self.schema.max_dimensions)
for actual, maximum in zip(
artifact.dimensions, self.schema.max_dimensions
)
):
raise ValueError(f"{field} exceeds its dimension limits")
@@ -211,7 +288,9 @@ class PortSpec:
def from_dict(cls, value: object) -> "PortSpec":
if not isinstance(value, Mapping):
raise ValueError("port specification must be an object")
require_exact_keys(value, {"schema", "cardinality", "collection"}, "port specification")
require_exact_keys(
value, {"schema", "cardinality", "collection"}, "port specification"
)
return cls(
schema=ArtifactSchema.from_dict(value["schema"]),
cardinality=value["cardinality"], # type: ignore[arg-type]
@@ -232,16 +311,28 @@ class ArtifactRef:
dimensions: tuple[int, ...] = ()
def __post_init__(self) -> None:
object.__setattr__(self, "artifact_id", require_uuid(self.artifact_id, "artifact_id"))
object.__setattr__(
self, "artifact_id", require_uuid(self.artifact_id, "artifact_id")
)
object.__setattr__(self, "sha256", require_sha256(self.sha256, "sha256"))
if not isinstance(self.schema, SchemaRef):
raise ValueError("artifact schema must be a SchemaRef")
object.__setattr__(self, "media_type", require_string(self.media_type, "media_type", max_length=128))
if "/" not in self.media_type or any(character.isspace() for character in self.media_type):
object.__setattr__(
self,
"media_type",
require_string(self.media_type, "media_type", max_length=128),
)
if "/" not in self.media_type or any(
character.isspace() for character in self.media_type
):
raise ValueError("media_type must be a valid type/subtype token")
object.__setattr__(self, "size_bytes", require_nonnegative_int(self.size_bytes, "size_bytes"))
object.__setattr__(
self, "size_bytes", require_nonnegative_int(self.size_bytes, "size_bytes")
)
if self.records is not None:
object.__setattr__(self, "records", require_nonnegative_int(self.records, "records"))
object.__setattr__(
self, "records", require_nonnegative_int(self.records, "records")
)
dimensions = tuple(self.dimensions)
if any(
isinstance(value, bool) or not isinstance(value, int) or value < 0
@@ -269,7 +360,15 @@ class ArtifactRef:
raise ValueError("artifact reference must be an object")
require_exact_keys(
value,
{"artifact_id", "sha256", "schema", "media_type", "size_bytes", "records", "dimensions"},
{
"artifact_id",
"sha256",
"schema",
"media_type",
"size_bytes",
"records",
"dimensions",
},
"artifact reference",
)
dimensions = value["dimensions"]
@@ -288,6 +387,12 @@ class ArtifactRef:
@dataclass(frozen=True, slots=True)
class ArtifactItem:
"""One artifact inside a collection, optionally bound to a string key.
Keys are canonical identifiers used by keyed collections (for example
shard indexes); ordered and set collections never use keys.
"""
artifact: ArtifactRef
key: str | None = None
@@ -295,7 +400,9 @@ class ArtifactItem:
if not isinstance(self.artifact, ArtifactRef):
raise ValueError("artifact item must contain an ArtifactRef")
if self.key is not None:
object.__setattr__(self, "key", require_identifier(self.key, "artifact key"))
object.__setattr__(
self, "key", require_identifier(self.key, "artifact key")
)
def to_dict(self) -> dict[str, object]:
return {"key": self.key, "artifact": self.artifact.to_dict()}
@@ -310,17 +417,29 @@ class ArtifactItem:
@dataclass(frozen=True, slots=True)
class ArtifactCollection:
"""An immutable set of artifacts with a canonical, content-addressed order.
Keyed collections sort by key; set collections sort by schema, checksum,
and size and reject duplicates; single collections hold at most one
unkeyed artifact. ``digest`` is the content pin used in provenance and
verification bindings.
"""
kind: CollectionKind
items: tuple[ArtifactItem, ...]
def __post_init__(self) -> None:
object.__setattr__(self, "kind", enum_value(CollectionKind, self.kind, "collection.kind"))
object.__setattr__(
self, "kind", enum_value(CollectionKind, self.kind, "collection.kind")
)
items = tuple(self.items)
if any(not isinstance(item, ArtifactItem) for item in items):
raise ValueError("collection items must be ArtifactItem values")
if self.kind is CollectionKind.SINGLE:
if len(items) > 1 or any(item.key is not None for item in items):
raise ValueError("single collection contains at most one unkeyed artifact")
raise ValueError(
"single collection contains at most one unkeyed artifact"
)
elif self.kind is CollectionKind.KEYED:
if any(item.key is None for item in items):
raise ValueError("keyed collection requires a key for every artifact")
@@ -333,11 +452,17 @@ class ArtifactCollection:
raise ValueError("ordered and set collections must not use keys")
if self.kind is CollectionKind.SET:
identities = [
(item.artifact.schema, item.artifact.sha256, item.artifact.size_bytes)
(
item.artifact.schema,
item.artifact.sha256,
item.artifact.size_bytes,
)
for item in items
]
if len(identities) != len(set(identities)):
raise ValueError("set collection must not contain duplicate artifacts")
raise ValueError(
"set collection must not contain duplicate artifacts"
)
items = tuple(
sorted(
items,
@@ -352,7 +477,9 @@ class ArtifactCollection:
@classmethod
def single(cls, artifact: ArtifactRef | None) -> "ArtifactCollection":
return cls(CollectionKind.SINGLE, () if artifact is None else (ArtifactItem(artifact),))
return cls(
CollectionKind.SINGLE, () if artifact is None else (ArtifactItem(artifact),)
)
@property
def size_bytes(self) -> int:
@@ -376,7 +503,10 @@ class ArtifactCollection:
return hashlib.sha256(canonical_json(payload).encode("utf-8")).hexdigest()
def to_dict(self) -> dict[str, object]:
return {"kind": self.kind.value, "items": [item.to_dict() for item in self.items]}
return {
"kind": self.kind.value,
"items": [item.to_dict() for item in self.items],
}
@classmethod
def from_dict(cls, value: object) -> "ArtifactCollection":
@@ -405,6 +535,13 @@ def _timestamp(value: object, field: str) -> str:
@dataclass(frozen=True, slots=True)
class Provenance:
"""Immutable audit record stamped onto every sealed output.
Captures the workload pin, negotiated features, input and parameter
digests, the execution contract digest, coordinator job/task identity,
and timestamps. Verification bindings are derived from it.
"""
workload: WorkloadId
sdk_api_version: str
protocol_version: str
@@ -433,34 +570,64 @@ class Provenance:
def __post_init__(self) -> None:
if not isinstance(self.workload, WorkloadId):
raise ValueError("provenance workload must be a WorkloadId")
object.__setattr__(self, "sdk_api_version", require_string(self.sdk_api_version, "sdk_api_version"))
object.__setattr__(self, "protocol_version", require_string(self.protocol_version, "protocol_version"))
object.__setattr__(
self,
"sdk_api_version",
require_string(self.sdk_api_version, "sdk_api_version"),
)
object.__setattr__(
self,
"protocol_version",
require_string(self.protocol_version, "protocol_version"),
)
parse_release(self.sdk_api_version, "sdk_api_version")
parse_release(self.protocol_version, "protocol_version")
object.__setattr__(
self,
"manifest_schema_version",
require_positive_int(self.manifest_schema_version, "manifest_schema_version"),
require_positive_int(
self.manifest_schema_version, "manifest_schema_version"
),
)
object.__setattr__(
self,
"workflow_schema_version",
require_positive_int(self.workflow_schema_version, "workflow_schema_version"),
require_positive_int(
self.workflow_schema_version, "workflow_schema_version"
),
)
if not isinstance(self.verifier, ComponentRef):
raise ValueError("provenance verifier must be a ComponentRef")
schemas = tuple(self.artifact_schemas)
if not schemas or any(not isinstance(schema, SchemaRef) for schema in schemas):
raise ValueError("provenance artifact_schemas must contain SchemaRef values")
raise ValueError(
"provenance artifact_schemas must contain SchemaRef values"
)
if len(schemas) != len(set(schemas)):
raise ValueError("provenance artifact_schemas must be unique")
if schemas != tuple(sorted(schemas, key=lambda schema: schema.canonical)):
raise ValueError("provenance artifact_schemas must be in canonical order")
object.__setattr__(self, "artifact_schemas", schemas)
object.__setattr__(self, "package_digest", require_sha256(self.package_digest, "package_digest", prefixed=True))
object.__setattr__(self, "manifest_digest", require_sha256(self.manifest_digest, "manifest_digest"))
object.__setattr__(self, "environment_digest", require_sha256(self.environment_digest, "environment_digest", prefixed=True))
runtime = freeze_json_mapping(self.worker_runtime, "worker_runtime", forbid_locations=True)
object.__setattr__(
self,
"package_digest",
require_sha256(self.package_digest, "package_digest", prefixed=True),
)
object.__setattr__(
self,
"manifest_digest",
require_sha256(self.manifest_digest, "manifest_digest"),
)
object.__setattr__(
self,
"environment_digest",
require_sha256(
self.environment_digest, "environment_digest", prefixed=True
),
)
runtime = freeze_json_mapping(
self.worker_runtime, "worker_runtime", forbid_locations=True
)
if len(canonical_json(runtime).encode("utf-8")) > 65_536:
raise ValueError("worker_runtime exceeds 64 KiB")
object.__setattr__(self, "worker_runtime", runtime)
@@ -471,8 +638,16 @@ class Provenance:
if not resource_ids or len(resource_ids) != len(set(resource_ids)):
raise ValueError("allocated_resource_ids must be non-empty and unique")
object.__setattr__(self, "allocated_resource_ids", resource_ids)
object.__setattr__(self, "parameters_digest", require_sha256(self.parameters_digest, "parameters_digest"))
object.__setattr__(self, "input_collection_digest", require_sha256(self.input_collection_digest, "input_collection_digest"))
object.__setattr__(
self,
"parameters_digest",
require_sha256(self.parameters_digest, "parameters_digest"),
)
object.__setattr__(
self,
"input_collection_digest",
require_sha256(self.input_collection_digest, "input_collection_digest"),
)
object.__setattr__(
self,
"execution_contract_digest",
@@ -488,7 +663,9 @@ class Provenance:
)
for name, version in selected_features.items():
require_identifier(name, "provenance selected feature")
require_string(version, "provenance selected feature version", max_length=32)
require_string(
version, "provenance selected feature version", max_length=32
)
parse_release(version, "provenance selected feature version")
for name, fallback in optional_fallbacks.items():
require_identifier(name, "provenance fallback feature")
@@ -497,21 +674,34 @@ class Provenance:
raise ValueError("provenance feature cannot be selected and fallbacked")
object.__setattr__(self, "selected_features", selected_features)
object.__setattr__(self, "optional_fallbacks", optional_fallbacks)
object.__setattr__(self, "job_id", require_uuid(self.job_id, "provenance.job_id"))
object.__setattr__(self, "task_id", require_uuid(self.task_id, "provenance.task_id"))
object.__setattr__(self, "started_at", _timestamp(self.started_at, "started_at"))
object.__setattr__(self, "finished_at", _timestamp(self.finished_at, "finished_at"))
if datetime.fromisoformat(self.finished_at.replace("Z", "+00:00")) < datetime.fromisoformat(
self.started_at.replace("Z", "+00:00")
):
object.__setattr__(
self, "job_id", require_uuid(self.job_id, "provenance.job_id")
)
object.__setattr__(
self, "task_id", require_uuid(self.task_id, "provenance.task_id")
)
object.__setattr__(
self, "started_at", _timestamp(self.started_at, "started_at")
)
object.__setattr__(
self, "finished_at", _timestamp(self.finished_at, "finished_at")
)
if datetime.fromisoformat(
self.finished_at.replace("Z", "+00:00")
) < datetime.fromisoformat(self.started_at.replace("Z", "+00:00")):
raise ValueError("finished_at must not precede started_at")
trust_mode = require_identifier(self.trust_mode, "provenance.trust_mode")
if trust_mode not in {"trusted", "verified", "untrusted_quorum"}:
raise ValueError("provenance.trust_mode is unsupported")
object.__setattr__(self, "trust_mode", trust_mode)
if self.random_seed is not None and (isinstance(self.random_seed, bool) or not isinstance(self.random_seed, int)):
if self.random_seed is not None and (
isinstance(self.random_seed, bool) or not isinstance(self.random_seed, int)
):
raise ValueError("random_seed must be an integer")
lineage = tuple(require_uuid(value, "checkpoint_lineage") for value in self.checkpoint_lineage)
lineage = tuple(
require_uuid(value, "checkpoint_lineage")
for value in self.checkpoint_lineage
)
if len(lineage) != len(set(lineage)):
raise ValueError("checkpoint_lineage must not contain duplicate artifacts")
object.__setattr__(self, "checkpoint_lineage", lineage)
@@ -549,21 +739,43 @@ class Provenance:
if not isinstance(value, Mapping):
raise ValueError("provenance must be an object")
fields = {
"workload", "sdk_api_version", "protocol_version", "manifest_schema_version",
"workflow_schema_version", "verifier", "artifact_schemas", "package_digest",
"manifest_digest", "environment_digest", "worker_runtime", "allocated_resource_ids",
"parameters_digest", "input_collection_digest", "execution_contract_digest",
"selected_features", "optional_fallbacks",
"job_id", "task_id",
"started_at", "finished_at",
"trust_mode", "random_seed", "checkpoint_lineage",
"workload",
"sdk_api_version",
"protocol_version",
"manifest_schema_version",
"workflow_schema_version",
"verifier",
"artifact_schemas",
"package_digest",
"manifest_digest",
"environment_digest",
"worker_runtime",
"allocated_resource_ids",
"parameters_digest",
"input_collection_digest",
"execution_contract_digest",
"selected_features",
"optional_fallbacks",
"job_id",
"task_id",
"started_at",
"finished_at",
"trust_mode",
"random_seed",
"checkpoint_lineage",
}
require_exact_keys(value, fields, "provenance")
resource_ids = value["allocated_resource_ids"]
artifact_schemas = value["artifact_schemas"]
lineage = value["checkpoint_lineage"]
if not isinstance(resource_ids, list) or not isinstance(artifact_schemas, list) or not isinstance(lineage, list):
raise ValueError("provenance resource IDs and checkpoint lineage must be arrays")
if (
not isinstance(resource_ids, list)
or not isinstance(artifact_schemas, list)
or not isinstance(lineage, list)
):
raise ValueError(
"provenance resource IDs and checkpoint lineage must be arrays"
)
return cls(
workload=WorkloadId.from_dict(value["workload"]),
sdk_api_version=value["sdk_api_version"], # type: ignore[arg-type]
@@ -571,7 +783,9 @@ class Provenance:
manifest_schema_version=value["manifest_schema_version"], # type: ignore[arg-type]
workflow_schema_version=value["workflow_schema_version"], # type: ignore[arg-type]
verifier=ComponentRef.from_dict(value["verifier"]),
artifact_schemas=tuple(SchemaRef.from_dict(item) for item in artifact_schemas),
artifact_schemas=tuple(
SchemaRef.from_dict(item) for item in artifact_schemas
),
package_digest=value["package_digest"], # type: ignore[arg-type]
manifest_digest=value["manifest_digest"], # type: ignore[arg-type]
environment_digest=value["environment_digest"], # type: ignore[arg-type]
@@ -601,7 +815,9 @@ class OutputManifest:
schema_version: int = OUTPUT_SCHEMA_VERSION
def __post_init__(self) -> None:
require_schema_version(self.schema_version, OUTPUT_SCHEMA_VERSION, "output schema_version")
require_schema_version(
self.schema_version, OUTPUT_SCHEMA_VERSION, "output schema_version"
)
object.__setattr__(self, "task_key", require_task_key(self.task_key))
if not isinstance(self.outputs, Mapping) or not self.outputs:
raise ValueError("outputs must be a non-empty object")
@@ -638,7 +854,9 @@ class OutputManifest:
details.append("missing " + ", ".join(missing))
if unexpected:
details.append("unexpected " + ", ".join(unexpected))
raise ValueError("output ports do not match the declaration: " + "; ".join(details))
raise ValueError(
"output ports do not match the declaration: " + "; ".join(details)
)
total = 0
for name, port in expected.items():
if not isinstance(port, PortSpec):
@@ -691,7 +909,10 @@ class OutputManifest:
return cls(
schema_version=value["schema_version"], # type: ignore[arg-type]
task_key=value["task_key"], # type: ignore[arg-type]
outputs={name: ArtifactCollection.from_dict(item) for name, item in outputs.items()},
outputs={
name: ArtifactCollection.from_dict(item)
for name, item in outputs.items()
},
metrics=value["metrics"], # type: ignore[arg-type]
provenance=Provenance.from_dict(value["provenance"]),
)
+10 -3
View File
@@ -342,7 +342,11 @@ class MapReduceWorkload:
"""
import csv
if isinstance(self.shard_rows, bool) or not isinstance(self.shard_rows, int) or self.shard_rows < 1:
if (
isinstance(self.shard_rows, bool)
or not isinstance(self.shard_rows, int)
or self.shard_rows < 1
):
raise ValueError("shard_rows must be a positive integer")
media_type = self.input_port.schema.media_type
if media_type == "text/tab-separated-values":
@@ -564,6 +568,7 @@ class MapReduceWorkload:
def run(self, context: TaskContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
assert self.map_stage_inputs is not None
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
inputs: dict[str, Path] = {}
@@ -572,8 +577,10 @@ class MapReduceWorkload:
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)
item = next(iter(collection.items), None)
if item is None:
raise ValueError(f"map task input {name} is empty")
inputs[name] = context.catalog.materialize(item.artifact)
output_path = workspace / "result.csv"
metrics = self.compute_shard(
inputs,
+210 -79
View File
@@ -66,7 +66,8 @@ class LocalArtifactStore:
tuple[int | None, tuple[int, ...]],
],
],
] | None = None,
]
| None = None,
) -> None:
self.root = root.resolve()
self.root.mkdir(parents=True, exist_ok=True)
@@ -96,7 +97,9 @@ class LocalArtifactStore:
try:
source_fd = os.open(path, source_flags)
except OSError as error:
raise ValueError("artifact sink could not open a regular non-symlink file") from error
raise ValueError(
"artifact sink could not open a regular non-symlink file"
) from error
try:
return self.seal_descriptor(
source_fd,
@@ -116,7 +119,11 @@ class LocalArtifactStore:
dimensions: tuple[int, ...] = (),
) -> ArtifactRef:
"""Copy and validate one already safely opened regular-file descriptor."""
if isinstance(descriptor, bool) or not isinstance(descriptor, int) or descriptor < 0:
if (
isinstance(descriptor, bool)
or not isinstance(descriptor, int)
or descriptor < 0
):
raise ValueError("artifact descriptor must be an open file descriptor")
if not isinstance(declaration, ArtifactSchema):
raise ValueError("artifact declaration must be an ArtifactSchema")
@@ -127,9 +134,10 @@ class LocalArtifactStore:
try:
if not stat.S_ISREG(os.fstat(source_fd).st_mode):
raise ValueError("artifact sink accepts only regular files")
with os.fdopen(source_fd, "rb", closefd=True) as source_file, os.fdopen(
temporary_fd, "wb", closefd=True
) as destination_file:
with (
os.fdopen(source_fd, "rb", closefd=True) as source_file,
os.fdopen(temporary_fd, "wb", closefd=True) as destination_file,
):
source_fd = -1
temporary_fd = -1
for block in iter(lambda: source_file.read(1024 * 1024), b""):
@@ -150,14 +158,22 @@ class LocalArtifactStore:
declaration,
)
if records is not None and records != measured_records:
raise ValueError("artifact record summary does not match inspected content")
raise ValueError(
"artifact record summary does not match inspected content"
)
if dimensions and dimensions != measured_dimensions:
raise ValueError("artifact dimension summary does not match inspected content")
raise ValueError(
"artifact dimension summary does not match inspected content"
)
if declaration.max_records is not None:
if measured_records is None:
raise ValueError("artifact validator did not produce a required record count")
raise ValueError(
"artifact validator did not produce a required record count"
)
if measured_records > declaration.max_records:
raise ValueError("sealed artifact exceeds its declared record limit")
raise ValueError(
"sealed artifact exceeds its declared record limit"
)
if declaration.max_dimensions:
if len(measured_dimensions) != len(declaration.max_dimensions) or any(
actual > maximum
@@ -166,7 +182,9 @@ class LocalArtifactStore:
declaration.max_dimensions,
)
):
raise ValueError("sealed artifact exceeds its declared dimension limits")
raise ValueError(
"sealed artifact exceeds its declared dimension limits"
)
reference = ArtifactRef(
artifact_id,
digest,
@@ -178,7 +196,9 @@ class LocalArtifactStore:
)
with self._lock:
if destination.is_symlink():
raise ValueError("local artifact destination must not be a symbolic link")
raise ValueError(
"local artifact destination must not be a symbolic link"
)
if destination.exists():
if not destination.is_file() or _sha256_file(destination) != digest:
raise ValueError("local artifact identity collision")
@@ -188,7 +208,9 @@ class LocalArtifactStore:
destination.chmod(0o444)
existing = self._references.get(artifact_id)
if existing is not None and existing != reference:
raise ValueError("local artifact identity was reused with different metadata")
raise ValueError(
"local artifact identity was reused with different metadata"
)
self._paths[artifact_id] = destination
self._references[artifact_id] = reference
self._refcounts[artifact_id] = self._refcounts.get(artifact_id, 0) + 1
@@ -254,7 +276,9 @@ class LocalArtifactStore:
configuration = dict(declaration.validator_configuration)
unknown = set(configuration) - {"columns", "required_columns"}
if unknown:
raise ValueError("delimited-table validator configuration has unknown fields")
raise ValueError(
"delimited-table validator configuration has unknown fields"
)
columns = configuration.get("columns")
required = configuration.get("required_columns", ())
if columns is not None and not isinstance(columns, (list, tuple)):
@@ -267,34 +291,54 @@ class LocalArtifactStore:
(expected_columns or (), "columns"),
(required_columns, "required_columns"),
):
if (
any(not isinstance(value, str) or not value for value in values)
or len(values) != len(set(values))
):
if any(not isinstance(value, str) or not value for value in values) or len(
values
) != len(set(values)):
raise ValueError(f"delimited-table {field_name} must be unique strings")
delimiter = "\t" if declaration.media_type == "text/tab-separated-values" else ","
delimiter = (
"\t" if declaration.media_type == "text/tab-separated-values" else ","
)
try:
with path.open("r", encoding="utf-8", newline="") as source_file:
reader = csv.reader(source_file, delimiter=delimiter)
try:
header = tuple(next(reader))
except StopIteration as error:
raise ValueError("delimited-table artifact must contain a header") from error
if not header or any(not value for value in header) or len(header) != len(set(header)):
raise ValueError(
"delimited-table artifact must contain a header"
) from error
if (
not header
or any(not value for value in header)
or len(header) != len(set(header))
):
raise ValueError("delimited-table artifact has an invalid header")
if expected_columns is not None and header != expected_columns:
raise ValueError("delimited-table artifact header does not match its schema")
raise ValueError(
"delimited-table artifact header does not match its schema"
)
if not set(required_columns).issubset(header):
raise ValueError("delimited-table artifact is missing required columns")
raise ValueError(
"delimited-table artifact is missing required columns"
)
count = 0
for row in reader:
if len(row) != len(header):
raise ValueError("delimited-table artifact has an inconsistent row width")
raise ValueError(
"delimited-table artifact has an inconsistent row width"
)
count += 1
if declaration.max_records is not None and count > declaration.max_records:
raise ValueError("sealed artifact exceeds its declared record limit")
if (
declaration.max_records is not None
and count > declaration.max_records
):
raise ValueError(
"sealed artifact exceeds its declared record limit"
)
except (UnicodeError, csv.Error) as error:
raise ValueError("sealed tabular artifact is not valid bounded text") from error
raise ValueError(
"sealed tabular artifact is not valid bounded text"
) from error
return count, ()
@staticmethod
@@ -325,8 +369,15 @@ class LocalArtifactStore:
ValueError("non-finite JSON number")
),
)
except (UnicodeError, json.JSONDecodeError, ValueError, RecursionError) as error:
raise ValueError("sealed JSON artifact is not a valid bounded document") from error
except (
UnicodeError,
json.JSONDecodeError,
ValueError,
RecursionError,
) as error:
raise ValueError(
"sealed JSON artifact is not a valid bounded document"
) from error
if top_level == "array" and not isinstance(value, list):
raise ValueError("JSON artifact must contain a top-level array")
if top_level == "object" and not isinstance(value, dict):
@@ -392,9 +443,13 @@ class LocalArtifactStore:
path = self._paths[artifact.artifact_id]
stored = self._references[artifact.artifact_id]
except KeyError as error:
raise ValueError("artifact is not present in the local store") from error
raise ValueError(
"artifact is not present in the local store"
) from error
if stored != artifact:
raise ValueError("artifact metadata does not match the sealed local reference")
raise ValueError(
"artifact metadata does not match the sealed local reference"
)
if (
path.is_symlink()
or not path.is_file()
@@ -479,13 +534,13 @@ class ScopedArtifactSink:
try:
lexical_relative = lexical.relative_to(self._workspace)
except ValueError as error:
raise ValueError("attempt artifact must remain inside its workspace") from error
raise ValueError(
"attempt artifact must remain inside its workspace"
) from error
if not lexical_relative.parts:
raise ValueError("attempt artifact must name a file inside its workspace")
directory_flags = (
os.O_RDONLY
| getattr(os, "O_DIRECTORY", 0)
| getattr(os, "O_NOFOLLOW", 0)
os.O_RDONLY | getattr(os, "O_DIRECTORY", 0) | getattr(os, "O_NOFOLLOW", 0)
)
opened_directories: list[int] = []
file_descriptor = -1
@@ -530,11 +585,15 @@ class ScopedArtifactSink:
existing = self._sealed.get(reference.artifact_id)
if existing is not None and existing != reference:
self._store.release(reference)
raise ValueError("attempt sealed conflicting metadata for one artifact")
raise ValueError(
"attempt sealed conflicting metadata for one artifact"
)
if existing is None:
if self._sealed_bytes + reference.size_bytes > self._max_bytes:
self._store.release(reference)
raise ValueError("attempt artifact bytes exceed their sink limit")
raise ValueError(
"attempt artifact bytes exceed their sink limit"
)
self._sealed[reference.artifact_id] = reference
self._sealed_bytes += reference.size_bytes
if self._transaction is not None:
@@ -567,7 +626,9 @@ class ScopedArtifactCatalog:
raise ValueError("catalog allowlist must contain ArtifactRef values")
existing = allowed.get(artifact.artifact_id)
if existing is not None and existing != artifact:
raise ValueError("catalog allowlist contains conflicting artifact metadata")
raise ValueError(
"catalog allowlist contains conflicting artifact metadata"
)
allowed[artifact.artifact_id] = artifact
self.__allowed = allowed
resolved_workspace = workspace.resolve()
@@ -620,6 +681,8 @@ class ScopedArtifactCatalog:
class CancellationFlag:
"""A thread-safe cancellation signal for local conformance runs."""
def __init__(self) -> None:
self._event = Event()
@@ -636,6 +699,12 @@ class CancellationFlag:
@dataclass(frozen=True, slots=True)
class LocalPlanningContext:
"""Scoped catalog/sink/workspace handed to a planner by the local executor.
Wraps a ``LocalArtifactStore`` into attempt-scoped catalog and sink so a
planner can materialize job inputs and seal shards.
"""
catalog: ArtifactCatalog
sink: ArtifactSink
workspace: Path
@@ -651,8 +720,14 @@ class LocalPlanningContext:
if any(not isinstance(value, ArtifactRef) for value in allowed_artifacts):
raise ValueError("allowed_artifacts must contain ArtifactRef values")
object.__setattr__(self, "allowed_artifacts", allowed_artifacts)
object.__setattr__(self, "max_artifacts", require_positive_int(self.max_artifacts, "max_artifacts"))
object.__setattr__(self, "max_bytes", require_positive_int(self.max_bytes, "max_bytes"))
object.__setattr__(
self,
"max_artifacts",
require_positive_int(self.max_artifacts, "max_artifacts"),
)
object.__setattr__(
self, "max_bytes", require_positive_int(self.max_bytes, "max_bytes")
)
if isinstance(self.catalog, LocalArtifactStore):
object.__setattr__(
self,
@@ -675,6 +750,12 @@ class LocalPlanningContext:
@dataclass(frozen=True, slots=True)
class LocalTaskContext:
"""Scoped context handed to a map/reduce handler by the local executor.
Carries the digest-pinned ``TaskSpec``, attempt workspace, cancellation
flag, provenance, and the accepted inputs for reducers.
"""
task: TaskSpec
catalog: ArtifactCatalog
sink: ArtifactSink
@@ -689,8 +770,14 @@ class LocalTaskContext:
def __post_init__(self) -> None:
workspace = self.workspace.resolve()
object.__setattr__(self, "workspace", workspace)
object.__setattr__(self, "max_artifacts", require_positive_int(self.max_artifacts, "max_artifacts"))
object.__setattr__(self, "max_bytes", require_positive_int(self.max_bytes, "max_bytes"))
object.__setattr__(
self,
"max_artifacts",
require_positive_int(self.max_artifacts, "max_artifacts"),
)
object.__setattr__(
self, "max_bytes", require_positive_int(self.max_bytes, "max_bytes")
)
if isinstance(self.catalog, LocalArtifactStore):
allowed_artifacts = tuple(
item.artifact
@@ -741,7 +828,9 @@ def _provenance(
job_id: str,
task_id: str,
) -> Provenance:
parameters_digest = hashlib.sha256(canonical_json(task.parameters).encode("utf-8")).hexdigest()
parameters_digest = hashlib.sha256(
canonical_json(task.parameters).encode("utf-8")
).hexdigest()
return Provenance(
workload=definition.manifest.workload,
sdk_api_version=task.sdk_api_version,
@@ -755,9 +844,7 @@ def _provenance(
item.artifact.schema
for collection in task.inputs.values()
for item in collection.items
}.union(
port.schema.ref for port in task.expected_outputs.values()
),
}.union(port.schema.ref for port in task.expected_outputs.values()),
key=lambda value: value.canonical,
)
),
@@ -768,7 +855,9 @@ def _provenance(
allocated_resource_ids=(allocation.allocation_id,) + allocation.accelerator_ids,
parameters_digest=parameters_digest,
input_collection_digest=_input_digest(task.inputs),
execution_contract_digest=hashlib.sha256(task.to_json().encode("utf-8")).hexdigest(),
execution_contract_digest=hashlib.sha256(
task.to_json().encode("utf-8")
).hexdigest(),
selected_features=task.selected_features,
optional_fallbacks=task.optional_fallbacks,
job_id=job_id,
@@ -800,7 +889,7 @@ def _verification_binding(manifest: OutputManifest) -> VerificationBinding:
manifest_schema_version=provenance.manifest_schema_version,
workflow_schema_version=provenance.workflow_schema_version,
artifact_schemas=provenance.artifact_schemas,
trust_mode=provenance.trust_mode,
trust_mode=TrustMode(provenance.trust_mode),
)
@@ -828,15 +917,23 @@ class LocalCoreBatchExecutor:
self.resources = ResourcePool(runtime.inventory, max_concurrency=1)
@staticmethod
def _assert_supported_profile(request: JobRequest, definition: WorkloadDefinition) -> None:
def _assert_supported_profile(
request: JobRequest, definition: WorkloadDefinition
) -> None:
if request.trust_mode is not TrustMode.TRUSTED:
raise ValueError("local conformance execution supports only trusted workloads")
raise ValueError(
"local conformance execution supports only trusted workloads"
)
workflow = definition.manifest.workflow
if workflow.failure_policy is not WorkflowFailurePolicy.FAIL_FAST:
raise ValueError("local conformance execution supports only fail-fast workflows")
raise ValueError(
"local conformance execution supports only fail-fast workflows"
)
for stage in workflow.stages:
if stage.kind not in {StageKind.MAP, StageKind.REDUCE}:
raise ValueError("local conformance execution does not implement advanced stages")
raise ValueError(
"local conformance execution does not implement advanced stages"
)
execution = stage.execution
if (
execution.process_model is not ProcessModel.SINGLE
@@ -845,7 +942,9 @@ class LocalCoreBatchExecutor:
or execution.native_threads != 1
or execution.nested_parallelism
):
raise ValueError("local conformance execution supports one non-nested host thread")
raise ValueError(
"local conformance execution supports one non-nested host thread"
)
if execution.network is not NetworkPolicy.TRUSTED:
raise ValueError(
"local conformance execution cannot enforce a restricted network policy"
@@ -857,19 +956,22 @@ class LocalCoreBatchExecutor:
or stage.gang is not None
or stage.resources.accelerator_count
):
raise ValueError("local conformance execution cannot enforce this stage profile")
raise ValueError(
"local conformance execution cannot enforce this stage profile"
)
if stage.retry.max_attempts != 1:
raise ValueError("local conformance execution does not implement retries")
reducers = tuple(stage for stage in workflow.stages if stage.kind is StageKind.REDUCE)
raise ValueError(
"local conformance execution does not implement retries"
)
reducers = tuple(
stage for stage in workflow.stages if stage.kind is StageKind.REDUCE
)
if len(reducers) == 1:
reducer = reducers[0]
if (
set(workflow.outputs) != set(reducer.outputs)
or any(
external_name != reference.port
or reference.stage_id != reducer.stage_id
for external_name, reference in workflow.outputs.items()
)
if set(workflow.outputs) != set(reducer.outputs) or any(
external_name != reference.port
or reference.stage_id != reducer.stage_id
for external_name, reference in workflow.outputs.items()
):
raise ValueError(
"local conformance execution requires identity-mapped reducer outputs"
@@ -930,11 +1032,14 @@ class LocalCoreBatchExecutor:
or task.protocol_version != self.runtime.protocol_version
or task.manifest_schema_version
!= definition.manifest.manifest_schema_version
or task.workflow_schema_version != definition.manifest.workflow.schema_version
or task.workflow_schema_version
!= definition.manifest.workflow.schema_version
or task.environment_digest != definition.manifest.environment.digest
or task.verifier != stage.verifier
):
raise ValueError("task resolved pins do not match the selected runtime and manifest")
raise ValueError(
"task resolved pins do not match the selected runtime and manifest"
)
provenance = _provenance(
definition,
self.runtime,
@@ -960,9 +1065,13 @@ class LocalCoreBatchExecutor:
if not isinstance(manifest, OutputManifest):
raise ValueError("workload handler must return an OutputManifest")
if manifest.task_key != task.task_key:
raise ValueError("handler output task_key does not match its trusted task")
raise ValueError(
"handler output task_key does not match its trusted task"
)
if manifest.provenance != provenance:
raise ValueError("handler output provenance does not match its trusted context")
raise ValueError(
"handler output provenance does not match its trusted context"
)
manifest.validate_against(
task.expected_outputs,
max_output_bytes=max_output_bytes,
@@ -1012,14 +1121,16 @@ class LocalCoreBatchExecutor:
expected_outputs, # type: ignore[arg-type]
max_output_bytes,
binding=_verification_binding(output),
trust_mode=output.provenance.trust_mode,
trust_mode=TrustMode(output.provenance.trust_mode),
),
CandidateOutputs((output,)),
)
if not isinstance(decision, VerificationDecision):
raise ValueError("declared verifier must return a VerificationDecision")
if decision.verifier != verifier_ref:
raise ValueError("verification decision identity does not match the declared verifier")
raise ValueError(
"verification decision identity does not match the declared verifier"
)
if decision.status is not VerificationStatus.ACCEPTED:
raise ValueError("task output did not pass its declared verifier")
@@ -1059,13 +1170,18 @@ class LocalCoreBatchExecutor:
self._assert_supported_profile(request, definition)
workflow = definition.manifest.workflow
map_stages = [stage for stage in workflow.stages if stage.kind is StageKind.MAP]
reduce_stages = [stage for stage in workflow.stages if stage.kind is StageKind.REDUCE]
reduce_stages = [
stage for stage in workflow.stages if stage.kind is StageKind.REDUCE
]
unsupported = [
stage for stage in workflow.stages
stage
for stage in workflow.stages
if stage.kind not in {StageKind.MAP, StageKind.REDUCE}
]
if len(map_stages) != 1 or len(reduce_stages) != 1 or unsupported:
raise ValueError("local core-batch executor supports one static map stage and one reducer")
raise ValueError(
"local core-batch executor supports one static map stage and one reducer"
)
limits = definition.manifest.limits
output_limit = min(limits.max_output_bytes, workflow.max_output_bytes)
job_id = str(uuid4())
@@ -1103,7 +1219,9 @@ class LocalCoreBatchExecutor:
}
for issued in planning.sink.sealed_references:
if planned_references.get(issued.artifact_id) != issued:
raise ValueError("planner sealed an artifact that is not referenced by its plan")
raise ValueError(
"planner sealed an artifact that is not referenced by its plan"
)
authorized_plan_inputs = {
item.artifact.artifact_id: item.artifact
for collection in request.inputs.values()
@@ -1141,7 +1259,9 @@ class LocalCoreBatchExecutor:
for reference in workflow.outputs.values()
)
):
raise ValueError("local core-batch executor requires a canonical map-to-reduce DAG")
raise ValueError(
"local core-batch executor requires a canonical map-to-reduce DAG"
)
runner = definition.runners[map_stage.entry_point]
map_results: list[OutputManifest] = []
for task_index, task in enumerate(plan.tasks):
@@ -1171,16 +1291,22 @@ class LocalCoreBatchExecutor:
raise ValueError("job exceeds the cumulative output byte limit")
map_results.append(manifest)
if len(map_results) != len(plan.tasks):
raise ValueError("map execution did not produce exactly one accepted result per task")
raise ValueError(
"map execution did not produce exactly one accepted result per task"
)
if len(map_stage.outputs) != 1 or len(reducer_stage.inputs) != 1:
raise ValueError("core map/reduce adapter requires one map output and one reducer input")
raise ValueError(
"core map/reduce adapter requires one map output and one reducer input"
)
map_port = next(iter(map_stage.outputs))
reducer_input_name = next(iter(reducer_stage.inputs))
partial_items: list[ArtifactItem] = []
for task, result in zip(plan.tasks, map_results):
collection = result.outputs[map_port]
if len(collection.items) != 1:
raise ValueError("core map stage must produce exactly one partial per planned task")
raise ValueError(
"core map stage must produce exactly one partial per planned task"
)
partial_items.append(
ArtifactItem(
collection.items[0].artifact,
@@ -1188,8 +1314,11 @@ class LocalCoreBatchExecutor:
)
)
if len(partial_items) != len(plan.tasks):
raise ValueError("core map stage must produce exactly one partial per planned task")
raise ValueError(
"core map stage must produce exactly one partial per planned task"
)
accepted = ArtifactCollection(CollectionKind.KEYED, tuple(partial_items))
assert reducer_stage.verifier is not None
reducer_task = TaskSpec(
workload=plan.workload,
package_digest=plan.package_digest,
@@ -1234,5 +1363,7 @@ class LocalCoreBatchExecutor:
)
if output_bytes > output_limit:
raise ValueError("job exceeds the cumulative output byte limit")
final.validate_against(definition.manifest.outputs, max_output_bytes=output_limit)
final.validate_against(
definition.manifest.outputs, max_output_bytes=output_limit
)
return final
+155 -31
View File
@@ -23,6 +23,12 @@ from .resources import ResourceAllocation, ResourceRequirements
class ProcessModel(str, Enum):
"""How a task's code is executed.
``SINGLE`` is the only profile the local conformance runtime executes;
the other models require runtime features that are fail-closed in v1.
"""
SINGLE = "single"
PROCESS_POOL = "process_pool"
THREAD_POOL = "thread_pool"
@@ -30,6 +36,12 @@ class ProcessModel(str, Enum):
class NetworkPolicy(str, Enum):
"""Network isolation declared for a stage.
``TRUSTED`` is the only policy the local conformance runtime accepts;
stricter policies require runtime enforcement.
"""
NONE = "none"
COORDINATOR_ARTIFACTS_ONLY = "coordinator_artifacts_only"
ALLOWLISTED_EGRESS = "allowlisted_egress"
@@ -37,6 +49,8 @@ class NetworkPolicy(str, Enum):
class FailureCategory(str, Enum):
"""Sanitized failure classification for retry and reporting policy."""
INPUT = "input"
SCIENTIFIC = "scientific"
RESOURCE = "resource"
@@ -48,13 +62,23 @@ class FailureCategory(str, Enum):
@dataclass(frozen=True, slots=True)
class RetryPolicy:
"""Retry budget and backoff for a stage.
``max_attempts`` greater than one requires the runtime to advertise
``retries``; the local conformance executor rejects retries.
"""
max_attempts: int = 1
retryable_categories: tuple[FailureCategory, ...] = ()
initial_backoff_seconds: int = 1
max_backoff_seconds: int = 60
def __post_init__(self) -> None:
object.__setattr__(self, "max_attempts", require_positive_int(self.max_attempts, "retry.max_attempts"))
object.__setattr__(
self,
"max_attempts",
require_positive_int(self.max_attempts, "retry.max_attempts"),
)
categories = tuple(
enum_value(FailureCategory, value, "retryable_category")
for value in self.retryable_categories
@@ -65,22 +89,30 @@ class RetryPolicy:
object.__setattr__(
self,
"initial_backoff_seconds",
require_nonnegative_int(self.initial_backoff_seconds, "retry.initial_backoff_seconds"),
require_nonnegative_int(
self.initial_backoff_seconds, "retry.initial_backoff_seconds"
),
)
object.__setattr__(
self,
"max_backoff_seconds",
require_nonnegative_int(self.max_backoff_seconds, "retry.max_backoff_seconds"),
require_nonnegative_int(
self.max_backoff_seconds, "retry.max_backoff_seconds"
),
)
if self.max_backoff_seconds < self.initial_backoff_seconds:
raise ValueError("retry max_backoff_seconds must not be less than initial_backoff_seconds")
raise ValueError(
"retry max_backoff_seconds must not be less than initial_backoff_seconds"
)
if self.max_attempts == 1 and categories:
raise ValueError("a non-retrying policy must not list retryable categories")
def to_dict(self) -> dict[str, object]:
return {
"max_attempts": self.max_attempts,
"retryable_categories": [category.value for category in self.retryable_categories],
"retryable_categories": [
category.value for category in self.retryable_categories
],
"initial_backoff_seconds": self.initial_backoff_seconds,
"max_backoff_seconds": self.max_backoff_seconds,
}
@@ -89,7 +121,12 @@ class RetryPolicy:
def from_dict(cls, value: object) -> "RetryPolicy":
if not isinstance(value, Mapping):
raise ValueError("retry policy must be an object")
fields = {"max_attempts", "retryable_categories", "initial_backoff_seconds", "max_backoff_seconds"}
fields = {
"max_attempts",
"retryable_categories",
"initial_backoff_seconds",
"max_backoff_seconds",
}
require_exact_keys(value, fields, "retry policy")
categories = value["retryable_categories"]
if not isinstance(categories, list):
@@ -104,6 +141,12 @@ class RetryPolicy:
@dataclass(frozen=True, slots=True)
class CheckpointPolicy:
"""Checkpoint declaration for a stage.
Declared but not executable until a runtime advertises ``checkpoints``;
enabled policies require a schema and a compatibility version.
"""
enabled: bool = False
schema: SchemaRef | None = None
compatibility_version: int | None = None
@@ -113,23 +156,38 @@ class CheckpointPolicy:
if not isinstance(self.enabled, bool):
raise ValueError("checkpoint.enabled must be a boolean")
if not self.enabled:
if any(value is not None for value in (self.schema, self.compatibility_version, self.interval_seconds)):
raise ValueError("disabled checkpoint policy must not declare checkpoint fields")
if any(
value is not None
for value in (
self.schema,
self.compatibility_version,
self.interval_seconds,
)
):
raise ValueError(
"disabled checkpoint policy must not declare checkpoint fields"
)
return
if not isinstance(self.schema, SchemaRef):
raise ValueError("enabled checkpoint policy requires a schema")
if self.compatibility_version is None:
raise ValueError("enabled checkpoint policy requires a compatibility_version")
raise ValueError(
"enabled checkpoint policy requires a compatibility_version"
)
object.__setattr__(
self,
"compatibility_version",
require_positive_int(self.compatibility_version, "checkpoint.compatibility_version"),
require_positive_int(
self.compatibility_version, "checkpoint.compatibility_version"
),
)
if self.interval_seconds is not None:
object.__setattr__(
self,
"interval_seconds",
require_positive_int(self.interval_seconds, "checkpoint.interval_seconds"),
require_positive_int(
self.interval_seconds, "checkpoint.interval_seconds"
),
)
def to_dict(self) -> dict[str, object]:
@@ -171,27 +229,52 @@ class ExecutionProfile:
secret_handles: tuple[str, ...] = ()
def __post_init__(self) -> None:
object.__setattr__(self, "profile", require_identifier(self.profile, "execution.profile"))
object.__setattr__(self, "process_model", enum_value(ProcessModel, self.process_model, "process_model"))
object.__setattr__(self, "max_processes", require_positive_int(self.max_processes, "max_processes"))
object.__setattr__(
self, "profile", require_identifier(self.profile, "execution.profile")
)
object.__setattr__(
self,
"process_model",
enum_value(ProcessModel, self.process_model, "process_model"),
)
object.__setattr__(
self,
"max_processes",
require_positive_int(self.max_processes, "max_processes"),
)
object.__setattr__(
self,
"threads_per_process",
require_positive_int(self.threads_per_process, "threads_per_process"),
)
object.__setattr__(self, "native_threads", require_positive_int(self.native_threads, "native_threads"))
object.__setattr__(
self,
"native_threads",
require_positive_int(self.native_threads, "native_threads"),
)
if not isinstance(self.nested_parallelism, bool):
raise ValueError("nested_parallelism must be a boolean")
object.__setattr__(self, "network", enum_value(NetworkPolicy, self.network, "network"))
object.__setattr__(self, "timeout_seconds", require_positive_int(self.timeout_seconds, "timeout_seconds"))
object.__setattr__(
self, "network", enum_value(NetworkPolicy, self.network, "network")
)
object.__setattr__(
self,
"timeout_seconds",
require_positive_int(self.timeout_seconds, "timeout_seconds"),
)
object.__setattr__(
self,
"cancellation_grace_seconds",
require_nonnegative_int(self.cancellation_grace_seconds, "cancellation_grace_seconds"),
require_nonnegative_int(
self.cancellation_grace_seconds, "cancellation_grace_seconds"
),
)
if not isinstance(self.checkpoint, CheckpointPolicy):
raise ValueError("checkpoint must be a CheckpointPolicy")
egress = tuple(require_string(value, "allowed_egress", max_length=253) for value in self.allowed_egress)
egress = tuple(
require_string(value, "allowed_egress", max_length=253)
for value in self.allowed_egress
)
if len(egress) != len(set(egress)):
raise ValueError("allowed_egress must be unique")
if self.network is NetworkPolicy.ALLOWLISTED_EGRESS and not egress:
@@ -199,7 +282,9 @@ class ExecutionProfile:
if self.network is not NetworkPolicy.ALLOWLISTED_EGRESS and egress:
raise ValueError("allowed_egress is valid only for allowlisted egress")
object.__setattr__(self, "allowed_egress", egress)
handles = tuple(require_identifier(value, "secret_handle") for value in self.secret_handles)
handles = tuple(
require_identifier(value, "secret_handle") for value in self.secret_handles
)
if len(handles) != len(set(handles)):
raise ValueError("secret_handles must be unique")
if handles and self.network is NetworkPolicy.NONE:
@@ -208,8 +293,14 @@ class ExecutionProfile:
if self.process_model is ProcessModel.SINGLE and (
self.max_processes != 1 or self.threads_per_process != 1
):
raise ValueError("single process model requires one process and one Python thread")
if not self.nested_parallelism and self.threads_per_process > 1 and self.native_threads > 1:
raise ValueError(
"single process model requires one process and one Python thread"
)
if (
not self.nested_parallelism
and self.threads_per_process > 1
and self.native_threads > 1
):
raise ValueError("nested thread pools require nested_parallelism=true")
@property
@@ -222,7 +313,9 @@ class ExecutionProfile:
if self.timeout_seconds > resources.max_duration_seconds:
raise ValueError("execution timeout exceeds the resource maximum duration")
def allocation_environment(self, allocation: ResourceAllocation) -> Mapping[str, str]:
def allocation_environment(
self, allocation: ResourceAllocation
) -> Mapping[str, str]:
"""Return only allocation-derived thread/device isolation variables."""
if not isinstance(allocation, ResourceAllocation):
raise ValueError("allocation must be a ResourceAllocation")
@@ -261,15 +354,26 @@ class ExecutionProfile:
if not isinstance(value, Mapping):
raise ValueError("execution profile must be an object")
fields = {
"profile", "process_model", "max_processes", "threads_per_process",
"native_threads", "nested_parallelism", "network", "timeout_seconds",
"cancellation_grace_seconds", "checkpoint", "allowed_egress", "secret_handles",
"profile",
"process_model",
"max_processes",
"threads_per_process",
"native_threads",
"nested_parallelism",
"network",
"timeout_seconds",
"cancellation_grace_seconds",
"checkpoint",
"allowed_egress",
"secret_handles",
}
require_exact_keys(value, fields, "execution profile")
allowed_egress = value["allowed_egress"]
secret_handles = value["secret_handles"]
if not isinstance(allowed_egress, list) or not isinstance(secret_handles, list):
raise ValueError("execution allowed_egress and secret_handles must be arrays")
raise ValueError(
"execution allowed_egress and secret_handles must be arrays"
)
return cls(
profile=value["profile"], # type: ignore[arg-type]
process_model=value["process_model"], # type: ignore[arg-type]
@@ -288,6 +392,12 @@ class ExecutionProfile:
@dataclass(frozen=True, slots=True)
class FailureReport:
"""A sanitized, location-free failure report for durable audit trails.
Message and evidence are bounded and reject local paths and transport
URLs; evidence is limited to 16 KiB.
"""
code: str
category: FailureCategory
retryable: bool
@@ -296,13 +406,27 @@ class FailureReport:
def __post_init__(self) -> None:
object.__setattr__(self, "code", require_identifier(self.code, "failure.code"))
object.__setattr__(self, "category", enum_value(FailureCategory, self.category, "failure.category"))
object.__setattr__(
self,
"category",
enum_value(FailureCategory, self.category, "failure.category"),
)
if not isinstance(self.retryable, bool):
raise ValueError("failure.retryable must be a boolean")
object.__setattr__(self, "message", require_safe_message(self.message, "failure.message", max_length=512))
evidence = freeze_json_mapping(self.evidence, "failure.evidence", forbid_locations=True)
object.__setattr__(
self,
"message",
require_safe_message(self.message, "failure.message", max_length=512),
)
evidence = freeze_json_mapping(
self.evidence, "failure.evidence", forbid_locations=True
)
import json
if len(json.dumps(thaw_json(evidence), allow_nan=False).encode("utf-8")) > 16_384:
if (
len(json.dumps(thaw_json(evidence), allow_nan=False).encode("utf-8"))
> 16_384
):
raise ValueError("failure evidence exceeds 16 KiB")
object.__setattr__(self, "evidence", evidence)
+47 -8
View File
@@ -1,4 +1,10 @@
"""Versioned identities used across the SciMesh workload SDK."""
"""Versioned identities used across the SciMesh workload SDK.
The module-level constants ``SDK_API_VERSION``, ``MANIFEST_SCHEMA_VERSION``,
``WORKFLOW_SCHEMA_VERSION``, ``TASK_SCHEMA_VERSION``, and
``OUTPUT_SCHEMA_VERSION`` pin the current wire/schema versions; manifests
declare explicit compatibility ranges against them.
"""
from __future__ import annotations
@@ -35,7 +41,9 @@ class VersionRange:
expression: str
def __post_init__(self) -> None:
object.__setattr__(self, "expression", validate_version_range(self.expression, "version range"))
object.__setattr__(
self, "expression", validate_version_range(self.expression, "version range")
)
def contains(self, version: str) -> bool:
return version_in_range(version, self.expression)
@@ -50,12 +58,16 @@ class VersionRange:
@dataclass(frozen=True, slots=True)
class WorkloadId:
"""The exact identity of a workload: canonical hyphenated name and semantic version."""
name: str
version: str
def __post_init__(self) -> None:
object.__setattr__(self, "name", require_workload_name(self.name))
object.__setattr__(self, "version", require_semver(self.version, "workload.version"))
object.__setattr__(
self, "version", require_semver(self.version, "workload.version")
)
def to_dict(self) -> dict[str, str]:
return {"name": self.name, "version": self.version}
@@ -70,12 +82,22 @@ class WorkloadId:
@dataclass(frozen=True, slots=True)
class SchemaRef:
"""A versioned artifact schema identity (``name@version``).
Schemas are content contracts: two artifacts share a schema only when
their ``SchemaRef`` values are equal.
"""
name: str
version: int
def __post_init__(self) -> None:
object.__setattr__(self, "name", require_identifier(self.name, "schema.name"))
if isinstance(self.version, bool) or not isinstance(self.version, int) or self.version < 1:
if (
isinstance(self.version, bool)
or not isinstance(self.version, int)
or self.version < 1
):
raise ValueError("schema.version must be a positive integer")
@property
@@ -111,8 +133,14 @@ class ComponentRef:
version: int
def __post_init__(self) -> None:
object.__setattr__(self, "name", require_identifier(self.name, "component.name"))
if isinstance(self.version, bool) or not isinstance(self.version, int) or self.version < 1:
object.__setattr__(
self, "name", require_identifier(self.name, "component.name")
)
if (
isinstance(self.version, bool)
or not isinstance(self.version, int)
or self.version < 1
):
raise ValueError("component.version must be a positive integer")
@property
@@ -135,6 +163,12 @@ class ComponentRef:
@dataclass(frozen=True, slots=True)
class FeatureRequirement:
"""A versioned feature a workload requires or optionally selects.
Optional features may declare a fallback that negotiation records when
the runtime does not provide the feature.
"""
name: str
versions: VersionRange
fallback: str | None = None
@@ -144,10 +178,15 @@ class FeatureRequirement:
if not isinstance(self.versions, VersionRange):
raise ValueError("feature.versions must be a VersionRange")
if self.fallback is not None:
object.__setattr__(self, "fallback", require_identifier(self.fallback, "feature.fallback"))
object.__setattr__(
self, "fallback", require_identifier(self.fallback, "feature.fallback")
)
def to_dict(self) -> dict[str, object]:
result: dict[str, object] = {"name": self.name, "versions": self.versions.expression}
result: dict[str, object] = {
"name": self.name,
"versions": self.versions.expression,
}
if self.fallback is not None:
result["fallback"] = self.fallback
return result
+44 -16
View File
@@ -4,6 +4,7 @@ from __future__ import annotations
import hashlib
import importlib.util
import os
from importlib import metadata
from pathlib import Path
@@ -22,11 +23,17 @@ def installed_distribution_digest(
files are excluded because they are neither stable wheel payloads nor used
by the registry's cache-isolated discovery import.
"""
installed = metadata.distribution(distribution) if isinstance(distribution, str) else distribution
installed = (
metadata.distribution(distribution)
if isinstance(distribution, str)
else distribution
)
raw_top_level = installed.read_text("top_level.txt")
if raw_top_level is None:
raise ValueError("installed distribution does not declare top-level packages")
declared_top_levels = [line.strip() for line in raw_top_level.splitlines() if line.strip()]
declared_top_levels = [
line.strip() for line in raw_top_level.splitlines() if line.strip()
]
if any(not value.isidentifier() for value in declared_top_levels):
raise ValueError("installed distribution declares an invalid top-level package")
top_levels = set(declared_top_levels)
@@ -34,55 +41,74 @@ def installed_distribution_digest(
raise ValueError("installed distribution has no measurable top-level package")
declared_files = tuple(installed.files or ())
editable_bootstrap = any(
Path(str(item)).name.startswith("__editable__") and Path(str(item)).suffix == ".pth"
Path(str(item)).name.startswith("__editable__")
and Path(str(item)).suffix == ".pth"
for item in declared_files
)
if editable_bootstrap and not allow_editable:
raise ValueError("editable workload installations are not accepted for secure discovery")
raise ValueError(
"editable workload installations are not accepted for secure discovery"
)
for item in declared_files:
relative = Path(str(item))
suffix = relative.suffix.lower()
if suffix == ".pth" and not allow_editable:
raise ValueError("installed workload distribution declares a .pth bootstrap")
raise ValueError(
"installed workload distribution declares a .pth bootstrap"
)
if suffix in {".pyc", ".pyo"} and "__pycache__" not in relative.parts:
raise ValueError("installed workload distribution declares sourceless bytecode")
raise ValueError(
"installed workload distribution declares sourceless bytecode"
)
selected: list[tuple[str, Path]] = []
for top_level in sorted(top_levels):
root = Path(installed.locate_file(top_level))
root = Path(str(installed.locate_file(top_level)))
if not root.exists():
# PEP 660 editable distributions may expose source packages through
# a meta-path finder rather than a physical site-packages path.
spec = importlib.util.find_spec(top_level)
locations = tuple(spec.submodule_search_locations or ()) if spec is not None else ()
locations = (
tuple(spec.submodule_search_locations or ()) if spec is not None else ()
)
if len(locations) > 1:
raise ValueError("shared namespace packages are not supported for workload integrity")
raise ValueError(
"shared namespace packages are not supported for workload integrity"
)
if locations:
root = Path(locations[0])
if root.is_symlink():
raise ValueError("installed workload package root must not be a symbolic link")
raise ValueError(
"installed workload package root must not be a symbolic link"
)
if root.is_dir():
candidates = root.rglob("*")
for path in candidates:
if path.is_symlink():
raise ValueError("installed workload package contains a symbolic-link payload")
raise ValueError(
"installed workload package contains a symbolic-link payload"
)
if not path.is_file():
continue
relative_parts = path.relative_to(root).parts
if "__pycache__" in relative_parts:
continue
if path.suffix.lower() in {".pyc", ".pyo"}:
raise ValueError("installed workload package contains sourceless bytecode")
raise ValueError(
"installed workload package contains sourceless bytecode"
)
relative = f"{top_level}/{path.relative_to(root).as_posix()}"
selected.append((relative, path))
continue
module = Path(installed.locate_file(top_level + ".py"))
module = Path(str(installed.locate_file(top_level + ".py")))
if not module.exists():
spec = importlib.util.find_spec(top_level)
if spec is not None and spec.origin is not None:
module = Path(spec.origin)
if module.is_symlink() or not module.is_file():
raise ValueError("installed workload package contains a missing top-level payload")
raise ValueError(
"installed workload package contains a missing top-level payload"
)
selected.append((top_level + ".py", module))
# Include declared package data outside top-level import trees. Generated
# console wrappers and installer metadata are excluded; executable .pth and
@@ -106,9 +132,11 @@ def installed_distribution_digest(
or "__pycache__" in relative.parts
):
continue
path = Path(installed.locate_file(item))
path = Path(str(installed.locate_file(item)))
if path.is_symlink():
raise ValueError("installed workload distribution contains a symbolic-link payload")
raise ValueError(
"installed workload distribution contains a symbolic-link payload"
)
if not path.is_file() or text in selected_names:
continue
selected.append((text, path))
+203 -44
View File
@@ -34,6 +34,12 @@ from .schema import validate_schema_definition
class DeterminismProfile(str, Enum):
"""How a workload's output is guaranteed to repeat.
``BYTE_EXACT`` is the only profile eligible for ``untrusted_quorum`` in
v1; the other profiles require future verifier or trust policies.
"""
BYTE_EXACT = "byte_exact"
CANONICAL_EXACT = "canonical_exact"
NUMERIC_TOLERANCE = "numeric_tolerance"
@@ -43,6 +49,13 @@ class DeterminismProfile(str, Enum):
class TrustMode(str, Enum):
"""Who may execute a workload and what acceptance requires.
``TRUSTED`` accepts a single execution; ``VERIFIED`` requires a
coordinator-owned binding; ``UNTRUSTED_QUORUM`` requires distinct owners
to produce identical whole-artifact SHA-256 digests.
"""
TRUSTED = "trusted"
VERIFIED = "verified"
UNTRUSTED_QUORUM = "untrusted_quorum"
@@ -50,27 +63,49 @@ class TrustMode(str, Enum):
@dataclass(frozen=True, slots=True)
class PackageSpec:
"""Identity pin of the installed distribution providing the workload.
``digest`` is the measured content pin (``sha256:`` prefix) that
discovery compares before importing an entry point.
"""
distribution: str
digest: str
signature: str | None = None
def __post_init__(self) -> None:
distribution = require_string(self.distribution, "package.distribution", max_length=128).lower()
distribution = require_string(
self.distribution, "package.distribution", max_length=128
).lower()
if not re.fullmatch(r"[a-z0-9]+(?:[-_.][a-z0-9]+)*", distribution):
raise ValueError("package.distribution must be a canonical Python distribution name")
raise ValueError(
"package.distribution must be a canonical Python distribution name"
)
object.__setattr__(self, "distribution", distribution.replace("_", "-"))
object.__setattr__(self, "digest", require_sha256(self.digest, "package.digest", prefixed=True))
object.__setattr__(
self, "digest", require_sha256(self.digest, "package.digest", prefixed=True)
)
if self.signature is not None:
object.__setattr__(self, "signature", require_string(self.signature, "package.signature", max_length=512))
object.__setattr__(
self,
"signature",
require_string(self.signature, "package.signature", max_length=512),
)
def to_dict(self) -> dict[str, object]:
return {"distribution": self.distribution, "digest": self.digest, "signature": self.signature}
return {
"distribution": self.distribution,
"digest": self.digest,
"signature": self.signature,
}
@classmethod
def from_dict(cls, value: object) -> "PackageSpec":
if not isinstance(value, Mapping):
raise ValueError("package specification must be an object")
require_exact_keys(value, {"distribution", "digest", "signature"}, "package specification")
require_exact_keys(
value, {"distribution", "digest", "signature"}, "package specification"
)
return cls(
distribution=value["distribution"], # type: ignore[arg-type]
digest=value["digest"], # type: ignore[arg-type]
@@ -80,23 +115,43 @@ class PackageSpec:
@dataclass(frozen=True, slots=True)
class EnvironmentSpec:
"""Pinned execution environment (kind, digest, metadata).
Negotiation fails unless the runtime inventory advertises this exact
environment digest.
"""
kind: str
digest: str
metadata: Mapping[str, Any]
def __post_init__(self) -> None:
object.__setattr__(self, "kind", require_identifier(self.kind, "environment.kind"))
object.__setattr__(self, "digest", require_sha256(self.digest, "environment.digest", prefixed=True))
object.__setattr__(self, "metadata", freeze_json_mapping(self.metadata, "environment.metadata"))
object.__setattr__(
self, "kind", require_identifier(self.kind, "environment.kind")
)
object.__setattr__(
self,
"digest",
require_sha256(self.digest, "environment.digest", prefixed=True),
)
object.__setattr__(
self, "metadata", freeze_json_mapping(self.metadata, "environment.metadata")
)
def to_dict(self) -> dict[str, object]:
return {"kind": self.kind, "digest": self.digest, "metadata": thaw_json(self.metadata)}
return {
"kind": self.kind,
"digest": self.digest,
"metadata": thaw_json(self.metadata),
}
@classmethod
def from_dict(cls, value: object) -> "EnvironmentSpec":
if not isinstance(value, Mapping):
raise ValueError("environment specification must be an object")
require_exact_keys(value, {"kind", "digest", "metadata"}, "environment specification")
require_exact_keys(
value, {"kind", "digest", "metadata"}, "environment specification"
)
return cls(
kind=value["kind"], # type: ignore[arg-type]
digest=value["digest"], # type: ignore[arg-type]
@@ -106,6 +161,12 @@ class EnvironmentSpec:
@dataclass(frozen=True, slots=True)
class VerifierSpec:
"""The manifest acceptance verifier and its bounded configuration.
The verifier must be installed in the definition and its handler
configuration must match this declaration exactly.
"""
verifier: ComponentRef
configuration: Mapping[str, Any]
@@ -128,7 +189,9 @@ class VerifierSpec:
def from_dict(cls, value: object) -> "VerifierSpec":
if not isinstance(value, Mapping):
raise ValueError("verifier specification must be an object")
require_exact_keys(value, {"verifier", "configuration"}, "verifier specification")
require_exact_keys(
value, {"verifier", "configuration"}, "verifier specification"
)
return cls(
verifier=ComponentRef.from_dict(value["verifier"]),
configuration=value["configuration"], # type: ignore[arg-type]
@@ -137,6 +200,12 @@ class VerifierSpec:
@dataclass(frozen=True, slots=True)
class WorkloadLimits:
"""Hard resource and size bounds enforced by planning and execution.
Covers input bytes, task count, output bytes, parameter bytes, and the
total artifact count for one job.
"""
max_input_bytes: int
max_tasks: int
max_output_bytes: int
@@ -145,9 +214,17 @@ class WorkloadLimits:
def __post_init__(self) -> None:
for field in (
"max_input_bytes", "max_tasks", "max_output_bytes", "max_parameter_bytes", "max_artifacts"
"max_input_bytes",
"max_tasks",
"max_output_bytes",
"max_parameter_bytes",
"max_artifacts",
):
object.__setattr__(self, field, require_positive_int(getattr(self, field), f"limits.{field}"))
object.__setattr__(
self,
field,
require_positive_int(getattr(self, field), f"limits.{field}"),
)
def to_dict(self) -> dict[str, int]:
return {
@@ -163,7 +240,11 @@ class WorkloadLimits:
if not isinstance(value, Mapping):
raise ValueError("workload limits must be an object")
fields = {
"max_input_bytes", "max_tasks", "max_output_bytes", "max_parameter_bytes", "max_artifacts",
"max_input_bytes",
"max_tasks",
"max_output_bytes",
"max_parameter_bytes",
"max_artifacts",
}
require_exact_keys(value, fields, "workload limits")
return cls(**value) # type: ignore[arg-type]
@@ -186,6 +267,15 @@ def _ports(
@dataclass(frozen=True, slots=True)
class WorkloadManifest:
"""The installed workload's complete, immutable declaration.
Pins identity, SDK/protocol compatibility ranges, package and environment
digests, the strict parameter schema, the workflow DAG, external ports,
determinism, trust modes, the acceptance verifier, limits, capabilities,
and conformance profiles. ``digest`` is the canonical JSON content pin
carried by every plan and task.
"""
sdk_api: VersionRange
protocol: VersionRange
workload: WorkloadId
@@ -212,16 +302,33 @@ class WorkloadManifest:
MANIFEST_SCHEMA_VERSION,
"manifest_schema_version",
)
if not isinstance(self.sdk_api, VersionRange) or not isinstance(self.protocol, VersionRange):
raise ValueError("sdk_api and protocol must be explicit VersionRange values")
if not isinstance(self.sdk_api, VersionRange) or not isinstance(
self.protocol, VersionRange
):
raise ValueError(
"sdk_api and protocol must be explicit VersionRange values"
)
if not isinstance(self.workload, WorkloadId):
raise ValueError("workload must be a WorkloadId")
object.__setattr__(self, "description", require_string(self.description, "description", max_length=512))
if not isinstance(self.package, PackageSpec) or not isinstance(self.environment, EnvironmentSpec):
raise ValueError("manifest package and environment declarations are required")
object.__setattr__(
self,
"description",
require_string(self.description, "description", max_length=512),
)
if not isinstance(self.package, PackageSpec) or not isinstance(
self.environment, EnvironmentSpec
):
raise ValueError(
"manifest package and environment declarations are required"
)
schema = freeze_json_mapping(self.parameters_schema, "parameters_schema")
if schema.get("type") != "object" or schema.get("additionalProperties") is not False:
raise ValueError("parameters_schema must be an object schema with additionalProperties=false")
if (
schema.get("type") != "object"
or schema.get("additionalProperties") is not False
):
raise ValueError(
"parameters_schema must be an object schema with additionalProperties=false"
)
properties = schema.get("properties")
if not isinstance(properties, Mapping):
raise ValueError("parameters_schema.properties must be an object")
@@ -231,14 +338,22 @@ class WorkloadManifest:
object.__setattr__(self, "parameters_schema", schema)
if not isinstance(self.workflow, WorkflowSpec):
raise ValueError("workflow must be a WorkflowSpec")
object.__setattr__(self, "inputs", _ports(self.inputs, "manifest.inputs", allow_empty=True))
object.__setattr__(
self, "inputs", _ports(self.inputs, "manifest.inputs", allow_empty=True)
)
object.__setattr__(self, "outputs", _ports(self.outputs, "manifest.outputs"))
if dict(self.inputs) != dict(self.workflow.inputs):
raise ValueError("manifest inputs must match workflow inputs")
if dict(self.outputs) != dict(self.workflow.output_ports()):
raise ValueError("manifest outputs must match workflow outputs")
object.__setattr__(self, "determinism", enum_value(DeterminismProfile, self.determinism, "determinism"))
modes = tuple(enum_value(TrustMode, mode, "trust_mode") for mode in self.trust_modes)
object.__setattr__(
self,
"determinism",
enum_value(DeterminismProfile, self.determinism, "determinism"),
)
modes = tuple(
enum_value(TrustMode, mode, "trust_mode") for mode in self.trust_modes
)
if not modes or len(modes) != len(set(modes)):
raise ValueError("trust_modes must be non-empty and unique")
object.__setattr__(self, "trust_modes", modes)
@@ -250,9 +365,13 @@ class WorkloadManifest:
}
for stage in self.workflow.stages:
if not set(stage.trust_modes).issubset(manifest_mode_values):
raise ValueError("stage trust modes must be a subset of manifest trust_modes")
raise ValueError(
"stage trust modes must be a subset of manifest trust_modes"
)
if stage.verifier is None:
raise ValueError("every output-producing stage requires an acceptance verifier")
raise ValueError(
"every output-producing stage requires an acceptance verifier"
)
resource_sets = (stage.resources,) + (
(stage.gang.per_replica_resources,) if stage.gang is not None else ()
)
@@ -279,19 +398,26 @@ class WorkloadManifest:
raise ValueError("workflow max_tasks exceeds the workload limit")
if self.workflow.max_output_bytes > self.limits.max_output_bytes:
raise ValueError("workflow max_output_bytes exceeds the workload limit")
capabilities = tuple(require_identifier(value, "capability") for value in self.capabilities)
capabilities = tuple(
require_identifier(value, "capability") for value in self.capabilities
)
if not capabilities or len(capabilities) != len(set(capabilities)):
raise ValueError("capabilities must be non-empty and unique")
if self.workload.name not in capabilities:
raise ValueError("capabilities must include the canonical workload name")
object.__setattr__(self, "capabilities", capabilities)
profiles = tuple(require_identifier(value, "conformance_profile") for value in self.conformance_profiles)
profiles = tuple(
require_identifier(value, "conformance_profile")
for value in self.conformance_profiles
)
if "core-batch-v1" not in profiles or len(profiles) != len(set(profiles)):
raise ValueError("conformance_profiles must uniquely include core-batch-v1")
object.__setattr__(self, "conformance_profiles", profiles)
required = tuple(self.required_features)
optional = tuple(self.optional_features)
if any(not isinstance(item, FeatureRequirement) for item in required + optional):
if any(
not isinstance(item, FeatureRequirement) for item in required + optional
):
raise ValueError("features must contain FeatureRequirement values")
names = [item.name for item in required + optional]
if len(names) != len(set(names)):
@@ -308,23 +434,34 @@ class WorkloadManifest:
if self.determinism is DeterminismProfile.BYTE_EXACT and not exact:
raise ValueError("byte_exact workloads require exact-artifact verifier")
if self.determinism is DeterminismProfile.CANONICAL_EXACT and not canonical:
raise ValueError("canonical_exact workloads require canonical-record verifier")
raise ValueError(
"canonical_exact workloads require canonical-record verifier"
)
if self.determinism is DeterminismProfile.NUMERIC_TOLERANCE and not numeric:
raise ValueError("numeric_tolerance workloads require numeric-tolerance verifier")
raise ValueError(
"numeric_tolerance workloads require numeric-tolerance verifier"
)
if TrustMode.UNTRUSTED_QUORUM in self.trust_modes:
if self.determinism is not DeterminismProfile.BYTE_EXACT or not exact:
raise ValueError("untrusted_quorum v1 requires byte_exact and exact-artifact")
if any(stage.kind is StageKind.SIDE_EFFECT for stage in self.workflow.stages):
raise ValueError(
"untrusted_quorum v1 requires byte_exact and exact-artifact"
)
if any(
stage.kind is StageKind.SIDE_EFFECT for stage in self.workflow.stages
):
raise ValueError("side-effect stages cannot use untrusted quorum")
if self.determinism is DeterminismProfile.SIDE_EFFECTING:
if self.trust_modes != (TrustMode.TRUSTED,):
raise ValueError("side_effecting workloads must be trusted-only")
if not any(stage.kind is StageKind.SIDE_EFFECT for stage in self.workflow.stages):
if not any(
stage.kind is StageKind.SIDE_EFFECT for stage in self.workflow.stages
):
raise ValueError("side_effecting workload requires a side-effect stage")
@property
def digest(self) -> str:
import hashlib
return hashlib.sha256(self.to_json().encode("utf-8")).hexdigest()
def to_dict(self) -> dict[str, object]:
@@ -358,21 +495,41 @@ class WorkloadManifest:
if not isinstance(value, Mapping):
raise ValueError("workload manifest must be an object")
fields = {
"manifest_schema_version", "sdk_api", "protocol", "workload", "description",
"package", "environment", "parameters_schema", "workflow", "inputs", "outputs",
"determinism", "trust_modes", "verifier", "limits", "capabilities",
"conformance_profiles", "required_features", "optional_features",
"manifest_schema_version",
"sdk_api",
"protocol",
"workload",
"description",
"package",
"environment",
"parameters_schema",
"workflow",
"inputs",
"outputs",
"determinism",
"trust_modes",
"verifier",
"limits",
"capabilities",
"conformance_profiles",
"required_features",
"optional_features",
}
require_exact_keys(value, fields, "workload manifest")
inputs, outputs = value["inputs"], value["outputs"]
arrays = (
value["trust_modes"], value["capabilities"], value["conformance_profiles"],
value["required_features"], value["optional_features"],
value["trust_modes"],
value["capabilities"],
value["conformance_profiles"],
value["required_features"],
value["optional_features"],
)
if not isinstance(inputs, Mapping) or not isinstance(outputs, Mapping):
raise ValueError("manifest inputs and outputs must be objects")
if any(not isinstance(item, list) for item in arrays):
raise ValueError("manifest trust, capability, profile, and feature fields must be arrays")
raise ValueError(
"manifest trust, capability, profile, and feature fields must be arrays"
)
return cls(
manifest_schema_version=value["manifest_schema_version"], # type: ignore[arg-type]
sdk_api=VersionRange.from_dict(value["sdk_api"]),
@@ -392,10 +549,12 @@ class WorkloadManifest:
capabilities=tuple(value["capabilities"]), # type: ignore[arg-type]
conformance_profiles=tuple(value["conformance_profiles"]), # type: ignore[arg-type]
required_features=tuple(
FeatureRequirement.from_dict(item) for item in value["required_features"] # type: ignore[union-attr]
FeatureRequirement.from_dict(item)
for item in value["required_features"] # type: ignore[union-attr]
),
optional_features=tuple(
FeatureRequirement.from_dict(item) for item in value["optional_features"] # type: ignore[union-attr]
FeatureRequirement.from_dict(item)
for item in value["optional_features"] # type: ignore[union-attr]
),
)
+202 -59
View File
@@ -85,6 +85,12 @@ def _fallbacks(value: object, field: str) -> Mapping[str, str]:
@dataclass(frozen=True, slots=True)
class JobRequest:
"""A user-requested job: workload identity, strict parameters, and inputs.
Parameters are frozen, JSON-safe, and location-free; required features
must be declared by the workload and available in the runtime.
"""
workload: WorkloadId
parameters: Mapping[str, Any]
inputs: Mapping[str, ArtifactCollection]
@@ -97,10 +103,15 @@ class JobRequest:
object.__setattr__(
self,
"parameters",
freeze_json_mapping(self.parameters, "job.parameters", forbid_locations=True),
freeze_json_mapping(
self.parameters, "job.parameters", forbid_locations=True
),
)
object.__setattr__(self, "inputs", _collections(self.inputs, "job.inputs"))
features = tuple(require_identifier(value, "required_feature") for value in self.required_features)
features = tuple(
require_identifier(value, "required_feature")
for value in self.required_features
)
if len(features) != len(set(features)):
raise ValueError("required_features must be unique")
object.__setattr__(self, "required_features", features)
@@ -112,7 +123,9 @@ class JobRequest:
@property
def parameters_digest(self) -> str:
return hashlib.sha256(canonical_json(self.parameters).encode("utf-8")).hexdigest()
return hashlib.sha256(
canonical_json(self.parameters).encode("utf-8")
).hexdigest()
def to_dict(self) -> dict[str, object]:
return {
@@ -135,11 +148,16 @@ class JobRequest:
inputs = value["inputs"]
features = value["required_features"]
if not isinstance(inputs, Mapping) or not isinstance(features, list):
raise ValueError("job inputs must be an object and required_features an array")
raise ValueError(
"job inputs must be an object and required_features an array"
)
return cls(
workload=WorkloadId.from_dict(value["workload"]),
parameters=value["parameters"], # type: ignore[arg-type]
inputs={name: ArtifactCollection.from_dict(item) for name, item in inputs.items()},
inputs={
name: ArtifactCollection.from_dict(item)
for name, item in inputs.items()
},
required_features=tuple(features),
trust_mode=value["trust_mode"], # type: ignore[arg-type]
)
@@ -155,6 +173,12 @@ class JobRequest:
@dataclass(frozen=True, slots=True)
class ValidatedJob:
"""A job after planner validation, carrying the resolved parameter set.
The planner may replace ambiguous parameters (for example a query id)
with their resolved values before tasks are planned.
"""
request: JobRequest
resolved_parameters: Mapping[str, Any]
@@ -173,7 +197,9 @@ class ValidatedJob:
@property
def parameters_digest(self) -> str:
return hashlib.sha256(canonical_json(self.resolved_parameters).encode("utf-8")).hexdigest()
return hashlib.sha256(
canonical_json(self.resolved_parameters).encode("utf-8")
).hexdigest()
@dataclass(frozen=True, slots=True)
@@ -201,7 +227,9 @@ class TaskSpec:
schema_version: int = TASK_SCHEMA_VERSION
def __post_init__(self) -> None:
require_schema_version(self.schema_version, TASK_SCHEMA_VERSION, "task schema_version")
require_schema_version(
self.schema_version, TASK_SCHEMA_VERSION, "task schema_version"
)
if not isinstance(self.workload, WorkloadId):
raise ValueError("task workload must be a WorkloadId")
object.__setattr__(
@@ -227,24 +255,32 @@ class TaskSpec:
object.__setattr__(
self,
"protocol_version",
require_string(self.protocol_version, "task protocol_version", max_length=32),
require_string(
self.protocol_version, "task protocol_version", max_length=32
),
)
parse_release(self.sdk_api_version, "task sdk_api_version")
parse_release(self.protocol_version, "task protocol_version")
object.__setattr__(
self,
"manifest_schema_version",
require_positive_int(self.manifest_schema_version, "task manifest_schema_version"),
require_positive_int(
self.manifest_schema_version, "task manifest_schema_version"
),
)
object.__setattr__(
self,
"workflow_schema_version",
require_positive_int(self.workflow_schema_version, "task workflow_schema_version"),
require_positive_int(
self.workflow_schema_version, "task workflow_schema_version"
),
)
object.__setattr__(
self,
"environment_digest",
require_sha256(self.environment_digest, "task environment_digest", prefixed=True),
require_sha256(
self.environment_digest, "task environment_digest", prefixed=True
),
)
if not isinstance(self.verifier, ComponentRef):
raise ValueError("task verifier must be a ComponentRef")
@@ -261,14 +297,22 @@ class TaskSpec:
if set(self.selected_features).intersection(self.optional_fallbacks):
raise ValueError("one task feature cannot be selected and fallbacked")
object.__setattr__(self, "task_key", require_task_key(self.task_key))
object.__setattr__(self, "stage_id", require_identifier(self.stage_id, "stage_id"))
object.__setattr__(
self, "stage_id", require_identifier(self.stage_id, "stage_id")
)
object.__setattr__(
self,
"parameters",
freeze_json_mapping(self.parameters, "task.parameters", forbid_locations=True),
freeze_json_mapping(
self.parameters, "task.parameters", forbid_locations=True
),
)
object.__setattr__(self, "inputs", _collections(self.inputs, "task.inputs"))
object.__setattr__(self, "expected_outputs", _ports(self.expected_outputs, "task.expected_outputs"))
object.__setattr__(
self,
"expected_outputs",
_ports(self.expected_outputs, "task.expected_outputs"),
)
if not self.expected_outputs:
raise ValueError("task expected_outputs must not be empty")
if not isinstance(self.resources, ResourceRequirements):
@@ -283,9 +327,9 @@ class TaskSpec:
canonical_port = require_identifier(port_name, "expected input key port")
if not isinstance(keys, (list, tuple)):
raise ValueError("expected input keys must be arrays")
canonical_keys = tuple(sorted(
require_identifier(key, "expected input key") for key in keys
))
canonical_keys = tuple(
sorted(require_identifier(key, "expected input key") for key in keys)
)
if not canonical_keys or len(canonical_keys) != len(set(canonical_keys)):
raise ValueError("expected input keys must be non-empty and unique")
expected_keys[canonical_port] = canonical_keys
@@ -310,7 +354,9 @@ class TaskSpec:
item.key for item in self.inputs[name].items if item.key is not None
)
if set(actual_keys) != set(expected_keys):
raise ValueError("task keyed input does not match its coordinator expected keys")
raise ValueError(
"task keyed input does not match its coordinator expected keys"
)
keyed_many_ports = {
name
for name, declaration in stage.inputs.items()
@@ -326,7 +372,9 @@ class TaskSpec:
if self.resources != stage.resources or self.execution != stage.execution:
raise ValueError("task execution requirements do not match the stage")
if self.verifier != stage.verifier:
raise ValueError("task verifier does not match the stage acceptance verifier")
raise ValueError(
"task verifier does not match the stage acceptance verifier"
)
if self.trust_mode.value not in stage.trust_modes:
raise ValueError("task trust mode is not allowed by the stage")
return self
@@ -350,7 +398,9 @@ class TaskSpec:
"stage_id": self.stage_id,
"parameters": thaw_json(self.parameters),
"inputs": {name: value.to_dict() for name, value in self.inputs.items()},
"expected_outputs": {name: value.to_dict() for name, value in self.expected_outputs.items()},
"expected_outputs": {
name: value.to_dict() for name, value in self.expected_outputs.items()
},
"resources": self.resources.to_dict(),
"execution": self.execution.to_dict(),
"expected_input_keys": {
@@ -371,12 +421,27 @@ class TaskSpec:
if not isinstance(value, Mapping):
raise ValueError("task specification must be an object")
fields = {
"schema_version", "workload", "package_digest", "manifest_digest", "trust_mode",
"sdk_api_version", "protocol_version", "manifest_schema_version",
"workflow_schema_version", "environment_digest", "verifier",
"selected_features", "optional_fallbacks",
"task_key", "stage_id", "parameters", "inputs",
"expected_outputs", "resources", "execution", "expected_input_keys",
"schema_version",
"workload",
"package_digest",
"manifest_digest",
"trust_mode",
"sdk_api_version",
"protocol_version",
"manifest_schema_version",
"workflow_schema_version",
"environment_digest",
"verifier",
"selected_features",
"optional_fallbacks",
"task_key",
"stage_id",
"parameters",
"inputs",
"expected_outputs",
"resources",
"execution",
"expected_input_keys",
}
require_exact_keys(value, fields, "task specification")
inputs, outputs = value["inputs"], value["expected_outputs"]
@@ -399,8 +464,13 @@ class TaskSpec:
task_key=value["task_key"], # type: ignore[arg-type]
stage_id=value["stage_id"], # type: ignore[arg-type]
parameters=value["parameters"], # type: ignore[arg-type]
inputs={name: ArtifactCollection.from_dict(item) for name, item in inputs.items()},
expected_outputs={name: PortSpec.from_dict(item) for name, item in outputs.items()},
inputs={
name: ArtifactCollection.from_dict(item)
for name, item in inputs.items()
},
expected_outputs={
name: PortSpec.from_dict(item) for name, item in outputs.items()
},
resources=ResourceRequirements.from_dict(value["resources"]),
execution=ExecutionProfile.from_dict(value["execution"]),
expected_input_keys=value["expected_input_keys"], # type: ignore[arg-type]
@@ -417,6 +487,12 @@ class TaskSpec:
@dataclass(frozen=True, slots=True)
class WorkflowPlan:
"""The immutable result of planning: tasks plus the exact workload pin.
Every task must carry the same package, manifest, environment, trust
mode, schema versions, and negotiated features as the plan itself.
"""
workload: WorkloadId
package_digest: str
manifest_digest: str
@@ -461,24 +537,32 @@ class WorkflowPlan:
object.__setattr__(
self,
"protocol_version",
require_string(self.protocol_version, "plan protocol_version", max_length=32),
require_string(
self.protocol_version, "plan protocol_version", max_length=32
),
)
parse_release(self.sdk_api_version, "plan sdk_api_version")
parse_release(self.protocol_version, "plan protocol_version")
object.__setattr__(
self,
"manifest_schema_version",
require_positive_int(self.manifest_schema_version, "plan manifest_schema_version"),
require_positive_int(
self.manifest_schema_version, "plan manifest_schema_version"
),
)
object.__setattr__(
self,
"workflow_schema_version",
require_positive_int(self.workflow_schema_version, "plan workflow_schema_version"),
require_positive_int(
self.workflow_schema_version, "plan workflow_schema_version"
),
)
object.__setattr__(
self,
"environment_digest",
require_sha256(self.environment_digest, "plan environment_digest", prefixed=True),
require_sha256(
self.environment_digest, "plan environment_digest", prefixed=True
),
)
if not isinstance(self.verifier, ComponentRef):
raise ValueError("plan verifier must be a ComponentRef")
@@ -494,7 +578,9 @@ class WorkflowPlan:
)
if set(self.selected_features).intersection(self.optional_fallbacks):
raise ValueError("one plan feature cannot be selected and fallbacked")
object.__setattr__(self, "workflow_id", require_identifier(self.workflow_id, "workflow_id"))
object.__setattr__(
self, "workflow_id", require_identifier(self.workflow_id, "workflow_id")
)
object.__setattr__(
self,
"resolved_parameters",
@@ -524,7 +610,9 @@ class WorkflowPlan:
or task.selected_features != self.selected_features
or task.optional_fallbacks != self.optional_fallbacks
):
raise ValueError("workflow plan tasks must carry the plan's exact workload pin")
raise ValueError(
"workflow plan tasks must carry the plan's exact workload pin"
)
object.__setattr__(self, "tasks", tasks)
def validate_workflow(self, workflow: WorkflowSpec) -> "WorkflowPlan":
@@ -538,10 +626,14 @@ class WorkflowPlan:
try:
task.validate_stage(stages[task.stage_id])
except KeyError as error:
raise ValueError(f"workflow plan references unknown stage: {task.stage_id}") from error
raise ValueError(
f"workflow plan references unknown stage: {task.stage_id}"
) from error
task_counts[task.stage_id] = task_counts.get(task.stage_id, 0) + 1
if task_counts[task.stage_id] > stages[task.stage_id].max_fan_out:
raise ValueError(f"workflow plan exceeds max_fan_out for stage {task.stage_id}")
raise ValueError(
f"workflow plan exceeds max_fan_out for stage {task.stage_id}"
)
return self
@property
@@ -576,11 +668,22 @@ class WorkflowPlan:
if not isinstance(value, Mapping):
raise ValueError("workflow plan must be an object")
fields = {
"schema_version", "workload", "package_digest", "manifest_digest", "trust_mode",
"sdk_api_version", "protocol_version", "manifest_schema_version",
"workflow_schema_version", "environment_digest", "verifier",
"selected_features", "optional_fallbacks",
"workflow_id", "resolved_parameters", "tasks",
"schema_version",
"workload",
"package_digest",
"manifest_digest",
"trust_mode",
"sdk_api_version",
"protocol_version",
"manifest_schema_version",
"workflow_schema_version",
"environment_digest",
"verifier",
"selected_features",
"optional_fallbacks",
"workflow_id",
"resolved_parameters",
"tasks",
}
require_exact_keys(value, fields, "workflow plan")
tasks = value["tasks"]
@@ -625,14 +728,20 @@ class ExpansionManifest:
schema_version: int = 1
def __post_init__(self) -> None:
require_schema_version(self.schema_version, 1, "expansion manifest schema_version")
object.__setattr__(self, "job_id", require_uuid(self.job_id, "expansion job_id"))
require_schema_version(
self.schema_version, 1, "expansion manifest schema_version"
)
object.__setattr__(
self, "job_id", require_uuid(self.job_id, "expansion job_id")
)
object.__setattr__(
self,
"parent_task_id",
require_uuid(self.parent_task_id, "expansion parent_task_id"),
)
object.__setattr__(self, "parent_task_key", require_task_key(self.parent_task_key))
object.__setattr__(
self, "parent_task_key", require_task_key(self.parent_task_key)
)
object.__setattr__(
self,
"parent_execution_contract_digest",
@@ -641,15 +750,23 @@ class ExpansionManifest:
"expansion parent_execution_contract_digest",
),
)
object.__setattr__(self, "max_children", require_positive_int(self.max_children, "max_children"))
object.__setattr__(
self,
"max_children",
require_positive_int(self.max_children, "max_children"),
)
tasks = tuple(self.tasks)
if not tasks or len(tasks) > self.max_children:
raise ValueError("expansion tasks must be non-empty and within max_children")
raise ValueError(
"expansion tasks must be non-empty and within max_children"
)
keys = [task.task_key for task in tasks]
if keys != sorted(keys) or len(keys) != len(set(keys)):
raise ValueError("expansion child task keys must be unique and ascending")
if any(not key.startswith(self.parent_task_key + "/") for key in keys):
raise ValueError("expansion child task keys must be namespaced by the parent")
raise ValueError(
"expansion child task keys must be namespaced by the parent"
)
first = tasks[0]
if any(
task.workload != first.workload
@@ -691,7 +808,9 @@ class ExpansionManifest:
raise ValueError("expansion workflow must be a WorkflowSpec")
if self.job_id != require_uuid(job_id, "coordinator job_id"):
raise ValueError("expansion belongs to another job")
if self.parent_task_id != require_uuid(parent_task_id, "coordinator parent_task_id"):
if self.parent_task_id != require_uuid(
parent_task_id, "coordinator parent_task_id"
):
raise ValueError("expansion belongs to another durable parent task")
if self.parent_task_key != parent.task_key:
raise ValueError("expansion parent task key does not match")
@@ -703,7 +822,9 @@ class ExpansionManifest:
try:
parent_stage = stages[parent.stage_id]
except KeyError as error:
raise ValueError("expansion parent references an unknown workflow stage") from error
raise ValueError(
"expansion parent references an unknown workflow stage"
) from error
parent.validate_stage(parent_stage)
if parent_stage.kind is not StageKind.PLAN:
raise ValueError("v1 expansion parent must be a plan stage")
@@ -726,8 +847,12 @@ class ExpansionManifest:
for port_name, collection in ports.items():
canonical_port = require_identifier(port_name, "authorized input port")
declaration = stages[canonical_stage].inputs.get(canonical_port)
if declaration is None or not isinstance(collection, ArtifactCollection):
raise ValueError("authorized_inputs references an unknown input port")
if declaration is None or not isinstance(
collection, ArtifactCollection
):
raise ValueError(
"authorized_inputs references an unknown input port"
)
declaration.validate_collection(
collection,
f"authorized input {canonical_stage}.{canonical_port}",
@@ -762,23 +887,34 @@ class ExpansionManifest:
or task.selected_features != parent.selected_features
or task.optional_fallbacks != parent.optional_fallbacks
):
raise ValueError("expansion child task does not share the parent workload pin")
raise ValueError(
"expansion child task does not share the parent workload pin"
)
try:
stage = stages[task.stage_id]
except KeyError as error:
raise ValueError("expansion child references an unknown workflow stage") from error
raise ValueError(
"expansion child references an unknown workflow stage"
) from error
if parent.stage_id not in stage.needs:
raise ValueError("v1 expansion child must be a direct successor of its parent stage")
raise ValueError(
"v1 expansion child must be a direct successor of its parent stage"
)
task.validate_stage(stage)
target_ports = allowed_by_target.get(task.stage_id, {})
for port_name, collection in task.inputs.items():
allowed = target_ports.get(port_name)
if allowed is None or collection.kind is not allowed.kind:
raise ValueError("expansion child input target is not coordinator-authorized")
raise ValueError(
"expansion child input target is not coordinator-authorized"
)
if collection.kind is CollectionKind.ORDERED:
cursor = 0
for item in collection.items:
while cursor < len(allowed.items) and allowed.items[cursor] != item:
while (
cursor < len(allowed.items)
and allowed.items[cursor] != item
):
cursor += 1
if cursor == len(allowed.items):
raise ValueError(
@@ -798,7 +934,9 @@ class ExpansionManifest:
@property
def digest(self) -> str:
return hashlib.sha256(canonical_json(self.to_dict()).encode("utf-8")).hexdigest()
return hashlib.sha256(
canonical_json(self.to_dict()).encode("utf-8")
).hexdigest()
def to_dict(self) -> dict[str, object]:
return {
@@ -819,8 +957,13 @@ class ExpansionManifest:
if not isinstance(value, Mapping):
raise ValueError("expansion manifest must be an object")
fields = {
"schema_version", "job_id", "parent_task_id", "parent_task_key",
"parent_execution_contract_digest", "max_children", "tasks",
"schema_version",
"job_id",
"parent_task_id",
"parent_task_key",
"parent_execution_contract_digest",
"max_children",
"tasks",
}
require_exact_keys(value, fields, "expansion manifest")
tasks = value["tasks"]
+34 -1
View File
@@ -5,7 +5,14 @@ from __future__ import annotations
from pathlib import Path
from typing import Mapping, Protocol, Sequence
from .artifacts import ArtifactCollection, ArtifactRef, ArtifactSchema, OutputManifest, Provenance
from .artifacts import (
ArtifactCollection,
ArtifactRef,
ArtifactSchema,
OutputManifest,
Provenance,
)
from .identity import ComponentRef
from .plans import JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
from .runtime import NegotiatedWorkload
from .verification import CandidateOutputs, VerificationDecision, VerifyContext
@@ -16,6 +23,7 @@ class ArtifactCatalog(Protocol):
def materialize(self, artifact: ArtifactRef) -> Path:
"""Return an attempt-scoped verified local copy without exposing credentials."""
...
class ArtifactSink(Protocol):
@@ -30,9 +38,12 @@ class ArtifactSink(Protocol):
dimensions: tuple[int, ...] = (),
) -> ArtifactRef:
"""Validate/upload bytes and return coordinator-owned immutable metadata."""
...
class CancellationToken(Protocol):
"""Cooperative cancellation observable by workload handlers."""
def cancelled(self) -> bool: ...
def raise_if_cancelled(self) -> None: ...
@@ -57,9 +68,13 @@ class PlanningContext(PlanningResources, Protocol):
@property
def negotiated(self) -> NegotiatedWorkload:
"""Resolved optional fallbacks and the exact negotiated manifest."""
...
class TaskContext(Protocol):
"""Everything a runner needs: the pinned task, scoped catalog/sink,
workspace, cancellation, and provenance to stamp on outputs."""
@property
def task(self) -> TaskSpec: ...
@@ -80,11 +95,18 @@ class TaskContext(Protocol):
class ReduceContext(TaskContext, Protocol):
"""TaskContext plus the keyed partial artifacts accepted for reduction."""
@property
def accepted_inputs(self) -> Mapping[str, ArtifactCollection]: ...
class Planner(Protocol):
"""Validates a job and produces a digest-pinned ``WorkflowPlan``.
``entry_point`` must match the workflow's PLAN stage when one exists.
"""
entry_point: str
def validate(self, request: JobRequest) -> ValidatedJob: ...
@@ -93,14 +115,25 @@ class Planner(Protocol):
class Runner(Protocol):
"""Executes one map (or verify) task and seals its partial output."""
def run(self, context: TaskContext) -> OutputManifest: ...
class Reducer(Protocol):
"""Merges accepted partial artifacts into the final result."""
def reduce(self, context: ReduceContext) -> OutputManifest: ...
class Verifier(Protocol):
"""Accepts or rejects candidate outputs with bounded sanitized evidence.
``identity`` must match the key under which the verifier is registered.
"""
identity: ComponentRef
def verify(
self,
context: VerifyContext,
+25 -13
View File
@@ -78,7 +78,7 @@ def _validate_entry_point_ownership(entry_point: metadata.EntryPoint) -> None:
if normalized_owners and normalized_owners != {expected_owner}:
raise ValueError("workload entry point top-level package is not uniquely owned")
package_root = Path(distribution.locate_file(root_name))
package_root = Path(str(distribution.locate_file(root_name)))
if not package_root.exists():
root_spec = util.find_spec(root_name)
locations = (
@@ -109,8 +109,8 @@ def _validate_entry_point_ownership(entry_point: metadata.EntryPoint) -> None:
else:
if len(parts) != 1:
raise ValueError("workload entry point module is outside its distribution")
ownership_root = Path(distribution.locate_file(".")).resolve()
module_base = Path(distribution.locate_file(root_name))
ownership_root = Path(str(distribution.locate_file("."))).resolve()
module_base = Path(str(distribution.locate_file(root_name)))
candidates = [
*(Path(str(module_base) + suffix) for suffix in machinery.SOURCE_SUFFIXES),
*(
@@ -125,6 +125,14 @@ def _validate_entry_point_ownership(entry_point: metadata.EntryPoint) -> None:
@dataclass(frozen=True, slots=True)
class WorkloadDefinition:
"""The immutable binding of a manifest to its installed handlers.
Validation at construction requires every stage entry point to have a
matching runner/reducer, the manifest verifier to be installed with
matching configuration, and verifier handler identities to match their
keys.
"""
manifest: WorkloadManifest
planner: Planner
runners: Mapping[str, Runner]
@@ -215,6 +223,12 @@ class WorkloadDefinition:
@dataclass(frozen=True, slots=True)
class AllowedPackage:
"""An administrator's approval to load one installed workload version.
Pins the distribution, the exact ``WorkloadId``, and the measured
``sha256:`` package digest that discovery must match.
"""
distribution: str
workload: WorkloadId
digest: str
@@ -272,6 +286,8 @@ def workload_allowlist_from_json(value: object) -> tuple[AllowedPackage, ...]:
@dataclass(frozen=True, slots=True)
class WorkloadDescription:
"""A read-only registry listing: identity, description, digest, and enablement."""
workload: WorkloadId
description: str
package_digest: str
@@ -581,11 +597,10 @@ class WorkloadRegistry:
discovered: set[tuple[str, str, str]] = set()
pending: list[WorkloadDefinition] = []
for entry_point in selected:
distribution = (
_normalized_distribution_name(entry_point.dist.name)
if entry_point.dist
else ""
)
entry_dist = entry_point.dist
if entry_dist is None:
continue
distribution = _normalized_distribution_name(entry_dist.name)
for key, approval in allowed.items():
if _normalized_distribution_name(key[0]) != distribution:
continue
@@ -606,7 +621,7 @@ class WorkloadRegistry:
prefix="scimesh-discovery-cache-"
) as cache_prefix,
):
measured_before = installed_distribution_digest(entry_point.dist)
measured_before = installed_distribution_digest(entry_dist)
if measured_before != approval.digest:
raise ValueError(
"installed package content does not match its allowlist digest"
@@ -626,10 +641,7 @@ class WorkloadRegistry:
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_dist) != measured_before:
raise ValueError(
"installed package content changed while loading its entry point"
)
+270 -66
View File
@@ -24,6 +24,8 @@ from ._validation import (
class AcceleratorMode(str, Enum):
"""How an accelerator is allocated: whole device or a managed partition."""
NONE = "none"
EXCLUSIVE_DEVICE = "exclusive_device"
FRACTIONAL = "fractional"
@@ -36,6 +38,12 @@ def _resource_id(value: object, field: str) -> str:
@dataclass(frozen=True, slots=True)
class AcceleratorDevice:
"""One physical accelerator advertised in a host inventory.
Declared but not schedulable until a runtime advertises the matching
accelerator features.
"""
kind: str
vendor: str
device_id: str
@@ -48,27 +56,55 @@ class AcceleratorDevice:
healthy: bool = True
def __post_init__(self) -> None:
object.__setattr__(self, "kind", require_identifier(self.kind, "accelerator.kind"))
object.__setattr__(self, "vendor", require_identifier(self.vendor, "accelerator.vendor"))
object.__setattr__(self, "device_id", _resource_id(self.device_id, "accelerator.device_id"))
object.__setattr__(self, "model", require_string(self.model, "accelerator.model", max_length=160))
object.__setattr__(self, "memory_mb", require_positive_int(self.memory_mb, "accelerator.memory_mb"))
modes = tuple(enum_value(AcceleratorMode, mode, "accelerator.mode") for mode in self.modes)
object.__setattr__(
self, "kind", require_identifier(self.kind, "accelerator.kind")
)
object.__setattr__(
self, "vendor", require_identifier(self.vendor, "accelerator.vendor")
)
object.__setattr__(
self, "device_id", _resource_id(self.device_id, "accelerator.device_id")
)
object.__setattr__(
self,
"model",
require_string(self.model, "accelerator.model", max_length=160),
)
object.__setattr__(
self,
"memory_mb",
require_positive_int(self.memory_mb, "accelerator.memory_mb"),
)
modes = tuple(
enum_value(AcceleratorMode, mode, "accelerator.mode") for mode in self.modes
)
if not modes or AcceleratorMode.NONE in modes or len(modes) != len(set(modes)):
raise ValueError("accelerator modes must contain unique allocation modes other than none")
raise ValueError(
"accelerator modes must contain unique allocation modes other than none"
)
object.__setattr__(self, "modes", modes)
capabilities = freeze_json_mapping(self.capabilities, "accelerator.capabilities")
capabilities = freeze_json_mapping(
self.capabilities, "accelerator.capabilities"
)
if any(not isinstance(value, str) for value in capabilities.values()):
raise ValueError("accelerator capabilities must use string values")
object.__setattr__(self, "capabilities", capabilities)
if self.topology_group is not None:
object.__setattr__(self, "topology_group", _resource_id(self.topology_group, "topology_group"))
object.__setattr__(
self,
"topology_group",
_resource_id(self.topology_group, "topology_group"),
)
if self.partition_id is not None:
object.__setattr__(self, "partition_id", _resource_id(self.partition_id, "partition_id"))
object.__setattr__(
self, "partition_id", _resource_id(self.partition_id, "partition_id")
)
if AcceleratorMode.PARTITION not in modes:
raise ValueError("a partition_id requires partition allocation support")
if AcceleratorMode.EXCLUSIVE_DEVICE in modes:
raise ValueError("an accelerator partition cannot be allocated as a whole device")
raise ValueError(
"an accelerator partition cannot be allocated as a whole device"
)
if not isinstance(self.healthy, bool):
raise ValueError("accelerator.healthy must be a boolean")
@@ -95,8 +131,16 @@ class AcceleratorDevice:
if not isinstance(value, Mapping):
raise ValueError("accelerator device must be an object")
fields = {
"kind", "vendor", "device_id", "model", "memory_mb", "modes",
"capabilities", "topology_group", "partition_id", "healthy",
"kind",
"vendor",
"device_id",
"model",
"memory_mb",
"modes",
"capabilities",
"topology_group",
"partition_id",
"healthy",
}
require_exact_keys(value, fields, "accelerator device")
modes = value["modes"]
@@ -118,6 +162,8 @@ class AcceleratorDevice:
@dataclass(frozen=True, slots=True)
class ResourceInventory:
"""What a host offers: CPU, memory, scratch, architecture, environments, accelerators."""
cpu_cores: int
memory_mb: int
scratch_mb: int
@@ -126,13 +172,31 @@ class ResourceInventory:
environment_digests: tuple[str, ...] = ()
def __post_init__(self) -> None:
object.__setattr__(self, "cpu_cores", require_positive_int(self.cpu_cores, "inventory.cpu_cores"))
object.__setattr__(self, "memory_mb", require_positive_int(self.memory_mb, "inventory.memory_mb"))
object.__setattr__(self, "scratch_mb", require_nonnegative_int(self.scratch_mb, "inventory.scratch_mb"))
object.__setattr__(self, "architecture", require_identifier(self.architecture, "inventory.architecture"))
object.__setattr__(
self,
"cpu_cores",
require_positive_int(self.cpu_cores, "inventory.cpu_cores"),
)
object.__setattr__(
self,
"memory_mb",
require_positive_int(self.memory_mb, "inventory.memory_mb"),
)
object.__setattr__(
self,
"scratch_mb",
require_nonnegative_int(self.scratch_mb, "inventory.scratch_mb"),
)
object.__setattr__(
self,
"architecture",
require_identifier(self.architecture, "inventory.architecture"),
)
devices = tuple(self.accelerators)
if any(not isinstance(device, AcceleratorDevice) for device in devices):
raise ValueError("inventory accelerators must contain AcceleratorDevice values")
raise ValueError(
"inventory accelerators must contain AcceleratorDevice values"
)
ids = [device.allocation_id for device in devices]
if len(ids) != len(set(ids)):
raise ValueError("inventory accelerator allocation IDs must be unique")
@@ -160,26 +224,40 @@ class ResourceInventory:
if not isinstance(value, Mapping):
raise ValueError("resource inventory must be an object")
fields = {
"cpu_cores", "memory_mb", "scratch_mb", "architecture",
"accelerators", "environment_digests",
"cpu_cores",
"memory_mb",
"scratch_mb",
"architecture",
"accelerators",
"environment_digests",
}
require_exact_keys(value, fields, "resource inventory")
accelerators = value["accelerators"]
digests = value["environment_digests"]
if not isinstance(accelerators, list) or not isinstance(digests, list):
raise ValueError("inventory accelerators and environment_digests must be arrays")
raise ValueError(
"inventory accelerators and environment_digests must be arrays"
)
return cls(
cpu_cores=value["cpu_cores"], # type: ignore[arg-type]
memory_mb=value["memory_mb"], # type: ignore[arg-type]
scratch_mb=value["scratch_mb"], # type: ignore[arg-type]
architecture=value["architecture"], # type: ignore[arg-type]
accelerators=tuple(AcceleratorDevice.from_dict(device) for device in accelerators),
accelerators=tuple(
AcceleratorDevice.from_dict(device) for device in accelerators
),
environment_digests=tuple(digests),
)
@dataclass(frozen=True, slots=True)
class ResourceRequirements:
"""What one task needs; eligibility is checked against the inventory.
``cpu_cores`` is a reservation, never a concurrency claim; accelerator
declarations remain fail-closed until runtime support exists.
"""
profile: str
cpu_cores: int
memory_mb: int
@@ -196,55 +274,107 @@ class ResourceRequirements:
max_duration_seconds: int = 3600
def __post_init__(self) -> None:
object.__setattr__(self, "profile", require_identifier(self.profile, "resources.profile"))
object.__setattr__(self, "cpu_cores", require_positive_int(self.cpu_cores, "resources.cpu_cores"))
object.__setattr__(self, "memory_mb", require_positive_int(self.memory_mb, "resources.memory_mb"))
object.__setattr__(self, "scratch_mb", require_nonnegative_int(self.scratch_mb, "resources.scratch_mb"))
object.__setattr__(
self, "profile", require_identifier(self.profile, "resources.profile")
)
object.__setattr__(
self,
"cpu_cores",
require_positive_int(self.cpu_cores, "resources.cpu_cores"),
)
object.__setattr__(
self,
"memory_mb",
require_positive_int(self.memory_mb, "resources.memory_mb"),
)
object.__setattr__(
self,
"scratch_mb",
require_nonnegative_int(self.scratch_mb, "resources.scratch_mb"),
)
object.__setattr__(
self,
"accelerator_count",
require_nonnegative_int(self.accelerator_count, "resources.accelerator_count"),
require_nonnegative_int(
self.accelerator_count, "resources.accelerator_count"
),
)
object.__setattr__(
self,
"accelerator_memory_mb",
require_nonnegative_int(self.accelerator_memory_mb, "resources.accelerator_memory_mb"),
require_nonnegative_int(
self.accelerator_memory_mb, "resources.accelerator_memory_mb"
),
)
object.__setattr__(
self,
"accelerator_mode",
enum_value(AcceleratorMode, self.accelerator_mode, "resources.accelerator_mode"),
enum_value(
AcceleratorMode, self.accelerator_mode, "resources.accelerator_mode"
),
)
if self.accelerator_count == 0:
if self.accelerator_kind is not None or self.accelerator_memory_mb or self.accelerator_mode is not AcceleratorMode.NONE:
raise ValueError("CPU-only resources must not declare accelerator constraints")
if (
self.accelerator_kind is not None
or self.accelerator_memory_mb
or self.accelerator_mode is not AcceleratorMode.NONE
):
raise ValueError(
"CPU-only resources must not declare accelerator constraints"
)
if self.topology_group is not None:
raise ValueError("CPU-only resources must not declare accelerator topology")
raise ValueError(
"CPU-only resources must not declare accelerator topology"
)
else:
if self.accelerator_kind is None:
raise ValueError("accelerator_kind is required when accelerator_count is non-zero")
object.__setattr__(self, "accelerator_kind", require_identifier(self.accelerator_kind, "accelerator_kind"))
raise ValueError(
"accelerator_kind is required when accelerator_count is non-zero"
)
object.__setattr__(
self,
"accelerator_kind",
require_identifier(self.accelerator_kind, "accelerator_kind"),
)
if self.accelerator_mode is AcceleratorMode.NONE:
raise ValueError("accelerator_mode is required when accelerator_count is non-zero")
raise ValueError(
"accelerator_mode is required when accelerator_count is non-zero"
)
if self.architecture is not None:
object.__setattr__(self, "architecture", require_identifier(self.architecture, "resources.architecture"))
object.__setattr__(
self,
"architecture",
require_identifier(self.architecture, "resources.architecture"),
)
if self.topology_group is not None:
object.__setattr__(self, "topology_group", _resource_id(self.topology_group, "resources.topology_group"))
object.__setattr__(
self,
"topology_group",
_resource_id(self.topology_group, "resources.topology_group"),
)
if self.environment_digest is not None:
object.__setattr__(
self,
"environment_digest",
require_sha256(self.environment_digest, "resources.environment_digest", prefixed=True),
require_sha256(
self.environment_digest,
"resources.environment_digest",
prefixed=True,
),
)
object.__setattr__(
self,
"estimated_input_bytes",
require_nonnegative_int(self.estimated_input_bytes, "estimated_input_bytes"),
require_nonnegative_int(
self.estimated_input_bytes, "estimated_input_bytes"
),
)
object.__setattr__(
self,
"estimated_output_bytes",
require_nonnegative_int(self.estimated_output_bytes, "estimated_output_bytes"),
require_nonnegative_int(
self.estimated_output_bytes, "estimated_output_bytes"
),
)
object.__setattr__(
self,
@@ -260,9 +390,15 @@ class ResourceRequirements:
errors.append("insufficient-memory")
if self.scratch_mb > inventory.scratch_mb:
errors.append("insufficient-scratch")
if self.architecture is not None and self.architecture != inventory.architecture:
if (
self.architecture is not None
and self.architecture != inventory.architecture
):
errors.append("architecture-mismatch")
if self.environment_digest is not None and self.environment_digest not in inventory.environment_digests:
if (
self.environment_digest is not None
and self.environment_digest not in inventory.environment_digests
):
errors.append("environment-unavailable")
matches = self._matching_devices(inventory.accelerators)
if len(matches) < self.accelerator_count:
@@ -286,22 +422,35 @@ class ResourceRequirements:
and device.memory_mb >= self.accelerator_memory_mb
and self.accelerator_mode in device.modes
and (
(self.accelerator_mode is AcceleratorMode.PARTITION and device.partition_id is not None)
(
self.accelerator_mode is AcceleratorMode.PARTITION
and device.partition_id is not None
)
or (
self.accelerator_mode is AcceleratorMode.EXCLUSIVE_DEVICE
and device.partition_id is None
)
or self.accelerator_mode is AcceleratorMode.FRACTIONAL
)
and (self.topology_group is None or device.topology_group == self.topology_group)
and (
self.topology_group is None
or device.topology_group == self.topology_group
)
]
if self.accelerator_count > 1 and self.topology_group is None:
groups: dict[str | None, list[AcceleratorDevice]] = {}
for device in matches:
groups.setdefault(device.topology_group, []).append(device)
sufficiently_large = [group for group in groups.values() if len(group) >= self.accelerator_count]
sufficiently_large = [
group
for group in groups.values()
if len(group) >= self.accelerator_count
]
if sufficiently_large:
matches = min(sufficiently_large, key=lambda group: tuple(item.allocation_id for item in group))
matches = min(
sufficiently_large,
key=lambda group: tuple(item.allocation_id for item in group),
)
return tuple(sorted(matches, key=lambda device: device.allocation_id))
def to_dict(self) -> dict[str, object]:
@@ -327,10 +476,20 @@ class ResourceRequirements:
if not isinstance(value, Mapping):
raise ValueError("resource requirements must be an object")
fields = {
"profile", "cpu_cores", "memory_mb", "scratch_mb", "accelerator_count",
"accelerator_kind", "accelerator_memory_mb", "accelerator_mode", "architecture",
"topology_group", "environment_digest", "estimated_input_bytes",
"estimated_output_bytes", "max_duration_seconds",
"profile",
"cpu_cores",
"memory_mb",
"scratch_mb",
"accelerator_count",
"accelerator_kind",
"accelerator_memory_mb",
"accelerator_mode",
"architecture",
"topology_group",
"environment_digest",
"estimated_input_bytes",
"estimated_output_bytes",
"max_duration_seconds",
}
require_exact_keys(value, fields, "resource requirements")
return cls(**value) # type: ignore[arg-type]
@@ -346,16 +505,32 @@ class ResourceAllocation:
accelerator_ids: tuple[str, ...]
def __post_init__(self) -> None:
object.__setattr__(self, "allocation_id", _resource_id(self.allocation_id, "allocation_id"))
object.__setattr__(
self, "allocation_id", _resource_id(self.allocation_id, "allocation_id")
)
object.__setattr__(
self,
"owner_id",
require_string(self.owner_id, "reservation owner_id", max_length=256),
)
object.__setattr__(self, "cpu_cores", require_positive_int(self.cpu_cores, "allocation.cpu_cores"))
object.__setattr__(self, "memory_mb", require_positive_int(self.memory_mb, "allocation.memory_mb"))
object.__setattr__(self, "scratch_mb", require_nonnegative_int(self.scratch_mb, "allocation.scratch_mb"))
ids = tuple(_resource_id(value, "accelerator_id") for value in self.accelerator_ids)
object.__setattr__(
self,
"cpu_cores",
require_positive_int(self.cpu_cores, "allocation.cpu_cores"),
)
object.__setattr__(
self,
"memory_mb",
require_positive_int(self.memory_mb, "allocation.memory_mb"),
)
object.__setattr__(
self,
"scratch_mb",
require_nonnegative_int(self.scratch_mb, "allocation.scratch_mb"),
)
ids = tuple(
_resource_id(value, "accelerator_id") for value in self.accelerator_ids
)
if len(ids) != len(set(ids)):
raise ValueError("accelerator_ids must be unique")
object.__setattr__(self, "accelerator_ids", ids)
@@ -379,7 +554,9 @@ class ResourcePool:
this local state.
"""
def __init__(self, inventory: ResourceInventory, *, max_concurrency: int = 1) -> None:
def __init__(
self, inventory: ResourceInventory, *, max_concurrency: int = 1
) -> None:
if not isinstance(inventory, ResourceInventory):
raise ValueError("inventory must be a ResourceInventory")
self.inventory = inventory
@@ -396,20 +573,33 @@ class ResourcePool:
return True
return left.partition_id == right.partition_id
def reserve(self, owner_id: str, requirements: ResourceRequirements) -> ResourceAllocation:
def reserve(
self, owner_id: str, requirements: ResourceRequirements
) -> ResourceAllocation:
if not isinstance(requirements, ResourceRequirements):
raise ValueError("requirements must be ResourceRequirements")
owner_id = require_string(owner_id, "reservation owner_id", max_length=256)
if requirements.accelerator_mode is AcceleratorMode.FRACTIONAL:
raise ResourceUnavailableError("fractional-accelerator-unsupported")
with self._lock:
if any(allocation.owner_id == owner_id for allocation in self._allocations.values()):
raise ValueError("reservation owner already has an active resource allocation")
if any(
allocation.owner_id == owner_id
for allocation in self._allocations.values()
):
raise ValueError(
"reservation owner already has an active resource allocation"
)
if len(self._allocations) >= self.max_concurrency:
raise ResourceUnavailableError("execution-slot-unavailable")
used_cpu = sum(allocation.cpu_cores for allocation in self._allocations.values())
used_memory = sum(allocation.memory_mb for allocation in self._allocations.values())
used_scratch = sum(allocation.scratch_mb for allocation in self._allocations.values())
used_cpu = sum(
allocation.cpu_cores for allocation in self._allocations.values()
)
used_memory = sum(
allocation.memory_mb for allocation in self._allocations.values()
)
used_scratch = sum(
allocation.scratch_mb for allocation in self._allocations.values()
)
if used_cpu + requirements.cpu_cores > self.inventory.cpu_cores:
raise ResourceUnavailableError("insufficient-cpu")
if used_memory + requirements.memory_mb > self.inventory.memory_mb:
@@ -419,7 +609,13 @@ class ResourcePool:
static_errors = tuple(
error
for error in requirements.eligibility_errors(self.inventory)
if error not in {"insufficient-cpu", "insufficient-memory", "insufficient-scratch", "accelerator-unavailable"}
if error
not in {
"insufficient-cpu",
"insufficient-memory",
"insufficient-scratch",
"accelerator-unavailable",
}
)
if static_errors:
raise ResourceUnavailableError(static_errors[0])
@@ -431,12 +627,18 @@ class ResourcePool:
available_devices = tuple(
device
for device in self.inventory.accelerators
if not any(self._devices_conflict(device, reserved) for reserved in reserved_devices)
if not any(
self._devices_conflict(device, reserved)
for reserved in reserved_devices
)
)
devices = requirements._matching_devices(available_devices)
if len(devices) < requirements.accelerator_count:
raise ResourceUnavailableError("accelerator-unavailable")
selected = tuple(device.allocation_id for device in devices[: requirements.accelerator_count])
selected = tuple(
device.allocation_id
for device in devices[: requirements.accelerator_count]
)
allocation = ResourceAllocation(
allocation_id=str(uuid4()),
owner_id=owner_id,
@@ -460,4 +662,6 @@ class ResourcePool:
def active_allocations(self) -> tuple[ResourceAllocation, ...]:
with self._lock:
return tuple(sorted(self._allocations.values(), key=lambda item: item.owner_id))
return tuple(
sorted(self._allocations.values(), key=lambda item: item.owner_id)
)
+75 -20
View File
@@ -6,7 +6,12 @@ from dataclasses import dataclass
from types import MappingProxyType
from typing import Mapping
from ._validation import require_identifier, require_string, validate_version_range, version_in_range
from ._validation import (
require_identifier,
require_string,
validate_version_range,
version_in_range,
)
from .identity import SDK_API_VERSION
from .execution import NetworkPolicy, ProcessModel
from .manifest import TrustMode, WorkloadManifest
@@ -15,6 +20,8 @@ from .workflow import StageKind
class CompatibilityError(ValueError):
"""A fail-closed negotiation failure with a stable machine-readable code."""
def __init__(self, code: str, message: str) -> None:
self.code = require_identifier(code, "compatibility error code")
super().__init__(message)
@@ -22,6 +29,9 @@ class CompatibilityError(ValueError):
@dataclass(frozen=True, slots=True)
class RuntimeCapabilities:
"""What a runtime advertises: SDK/protocol versions, profiles, features,
workload capabilities, inventory, and enforceable trust modes."""
sdk_api_version: str
protocol_version: str
profiles: tuple[str, ...]
@@ -31,13 +41,23 @@ class RuntimeCapabilities:
trust_modes: tuple[TrustMode, ...] = (TrustMode.TRUSTED,)
def __post_init__(self) -> None:
object.__setattr__(self, "sdk_api_version", require_string(self.sdk_api_version, "sdk_api_version"))
object.__setattr__(self, "protocol_version", require_string(self.protocol_version, "protocol_version"))
object.__setattr__(
self,
"sdk_api_version",
require_string(self.sdk_api_version, "sdk_api_version"),
)
object.__setattr__(
self,
"protocol_version",
require_string(self.protocol_version, "protocol_version"),
)
# Parsing as an equality range provides the same numeric release rules
# used by manifest ranges without accepting an implicit/latest value.
validate_version_range(f"=={self.sdk_api_version}", "sdk_api_version")
validate_version_range(f"=={self.protocol_version}", "protocol_version")
profiles = tuple(require_identifier(value, "runtime profile") for value in self.profiles)
profiles = tuple(
require_identifier(value, "runtime profile") for value in self.profiles
)
if len(profiles) != len(set(profiles)):
raise ValueError("runtime profiles must be unique")
object.__setattr__(self, "profiles", profiles)
@@ -50,7 +70,10 @@ class RuntimeCapabilities:
validate_version_range(f"=={text}", "runtime feature version")
features[canonical] = text
object.__setattr__(self, "features", MappingProxyType(features))
capabilities = tuple(require_identifier(value, "workload capability") for value in self.workload_capabilities)
capabilities = tuple(
require_identifier(value, "workload capability")
for value in self.workload_capabilities
)
if len(capabilities) != len(set(capabilities)):
raise ValueError("workload_capabilities must be unique")
object.__setattr__(self, "workload_capabilities", capabilities)
@@ -59,7 +82,9 @@ class RuntimeCapabilities:
try:
trust_modes = tuple(TrustMode(value) for value in self.trust_modes)
except (TypeError, ValueError) as error:
raise ValueError("runtime trust_modes contain an unsupported value") from error
raise ValueError(
"runtime trust_modes contain an unsupported value"
) from error
if not trust_modes or len(trust_modes) != len(set(trust_modes)):
raise ValueError("runtime trust_modes must be non-empty and unique")
object.__setattr__(self, "trust_modes", trust_modes)
@@ -67,6 +92,9 @@ class RuntimeCapabilities:
@dataclass(frozen=True, slots=True)
class NegotiatedWorkload:
"""The result of successful negotiation: selected features, fallbacks,
and the exact manifest a plan must pin."""
manifest: WorkloadManifest
optional_fallbacks: Mapping[str, str]
sdk_api_version: str
@@ -76,19 +104,29 @@ class NegotiatedWorkload:
def __post_init__(self) -> None:
if not isinstance(self.manifest, WorkloadManifest):
raise ValueError("negotiated manifest must be a WorkloadManifest")
object.__setattr__(self, "optional_fallbacks", MappingProxyType(dict(self.optional_fallbacks)))
object.__setattr__(
self, "optional_fallbacks", MappingProxyType(dict(self.optional_fallbacks))
)
object.__setattr__(
self,
"sdk_api_version",
require_string(self.sdk_api_version, "negotiated sdk_api_version", max_length=32),
require_string(
self.sdk_api_version, "negotiated sdk_api_version", max_length=32
),
)
object.__setattr__(
self,
"protocol_version",
require_string(self.protocol_version, "negotiated protocol_version", max_length=32),
require_string(
self.protocol_version, "negotiated protocol_version", max_length=32
),
)
validate_version_range(
f"=={self.sdk_api_version}", "negotiated sdk_api_version"
)
validate_version_range(
f"=={self.protocol_version}", "negotiated protocol_version"
)
validate_version_range(f"=={self.sdk_api_version}", "negotiated sdk_api_version")
validate_version_range(f"=={self.protocol_version}", "negotiated protocol_version")
selected: dict[str, str] = {}
for name, version in self.selected_features.items():
selected[require_identifier(name, "negotiated feature")] = require_string(
@@ -108,22 +146,33 @@ def negotiate_manifest(
runtime: RuntimeCapabilities,
) -> NegotiatedWorkload:
"""Resolve compatibility before any package handler or planner is invoked."""
if not isinstance(manifest, WorkloadManifest) or not isinstance(runtime, RuntimeCapabilities):
raise ValueError("negotiation requires WorkloadManifest and RuntimeCapabilities")
if not isinstance(manifest, WorkloadManifest) or not isinstance(
runtime, RuntimeCapabilities
):
raise ValueError(
"negotiation requires WorkloadManifest and RuntimeCapabilities"
)
if runtime.sdk_api_version != SDK_API_VERSION:
raise CompatibilityError(
"runtime-sdk-mismatch",
"runtime SDK declaration does not match this SDK implementation",
)
if not manifest.sdk_api.contains(runtime.sdk_api_version):
raise CompatibilityError("sdk-api-mismatch", "runtime SDK API is outside the manifest range")
raise CompatibilityError(
"sdk-api-mismatch", "runtime SDK API is outside the manifest range"
)
if not manifest.protocol.contains(runtime.protocol_version):
raise CompatibilityError("protocol-mismatch", "runtime protocol is outside the manifest range")
missing_profiles = sorted(set(manifest.conformance_profiles) - set(runtime.profiles))
raise CompatibilityError(
"protocol-mismatch", "runtime protocol is outside the manifest range"
)
missing_profiles = sorted(
set(manifest.conformance_profiles) - set(runtime.profiles)
)
if missing_profiles:
raise CompatibilityError(
"profile-unavailable",
"runtime does not support required profiles: " + ", ".join(missing_profiles),
"runtime does not support required profiles: "
+ ", ".join(missing_profiles),
)
if manifest.workload.name not in runtime.workload_capabilities:
raise CompatibilityError(
@@ -131,7 +180,9 @@ def negotiate_manifest(
"runtime does not advertise the canonical workload capability",
)
if manifest.environment.digest not in runtime.inventory.environment_digests:
raise CompatibilityError("environment-unavailable", "pinned workload environment is unavailable")
raise CompatibilityError(
"environment-unavailable", "pinned workload environment is unavailable"
)
for feature in manifest.required_features:
version = runtime.features.get(feature.name)
if version is None or not feature.versions.contains(version):
@@ -167,7 +218,9 @@ def negotiate_manifest(
def require_declared(condition: bool, feature: str, message: str) -> None:
if condition and feature not in declared_required:
raise CompatibilityError("feature-undeclared", message + f" requires {feature}")
raise CompatibilityError(
"feature-undeclared", message + f" requires {feature}"
)
for stage in manifest.workflow.stages:
shape_feature = required_by_shape.get(stage.kind)
@@ -177,7 +230,9 @@ def negotiate_manifest(
f"stage {stage.stage_id} requires declared feature {shape_feature}",
)
if stage.gang is not None and "gang-leases" not in declared_required:
raise CompatibilityError("feature-undeclared", "gang execution requires gang-leases")
raise CompatibilityError(
"feature-undeclared", "gang execution requires gang-leases"
)
execution = stage.execution
require_declared(
execution.process_model is ProcessModel.PROCESS_POOL,
+50 -11
View File
@@ -91,7 +91,10 @@ def _json_key(value: object, depth: int = 0) -> object:
if isinstance(value, bool):
return ("boolean", value)
if isinstance(value, (int, float)):
return ("number", Fraction(value) if isinstance(value, int) else Fraction.from_float(value))
return (
"number",
Fraction(value) if isinstance(value, int) else Fraction.from_float(value),
)
if isinstance(value, str):
return ("string", value)
if isinstance(value, Mapping):
@@ -148,7 +151,14 @@ def _is_json_multiple(value: int | float, divisor: int | float) -> bool:
return (value_fraction / divisor_fraction).denominator == 1
def validate_schema_definition(schema: Mapping[str, object], *, _depth: int = 0) -> None:
def validate_schema_definition(
schema: Mapping[str, object], *, _depth: int = 0
) -> None:
"""Validate a parameter schema against the bounded JSON Schema subset.
Raises ``ValueError`` on unknown keywords, unsupported types, unsafe
patterns, or malformed bounds.
"""
if _depth > 64:
raise _schema_error("nesting exceeds 64 levels")
if not isinstance(schema, Mapping):
@@ -167,7 +177,9 @@ def validate_schema_definition(schema: Mapping[str, object], *, _depth: int = 0)
raise _schema_error("type alternatives must be unique")
properties = schema.get("properties")
if properties is not None:
if not isinstance(properties, Mapping) or any(not isinstance(name, str) for name in properties):
if not isinstance(properties, Mapping) or any(
not isinstance(name, str) for name in properties
):
raise _schema_error("properties must be an object")
for child in properties.values():
validate_schema_definition(child, _depth=_depth + 1) # type: ignore[arg-type]
@@ -178,7 +190,9 @@ def validate_schema_definition(schema: Mapping[str, object], *, _depth: int = 0)
validate_schema_definition(additional, _depth=_depth + 1)
required = schema.get("required")
if required is not None:
if not isinstance(required, (list, tuple)) or any(not isinstance(name, str) for name in required):
if not isinstance(required, (list, tuple)) or any(
not isinstance(name, str) for name in required
):
raise _schema_error("required must be an array of strings")
if len(required) != len(set(required)):
raise _schema_error("required names must be unique")
@@ -205,20 +219,35 @@ def validate_schema_definition(schema: Mapping[str, object], *, _depth: int = 0)
raise _schema_error("enum values must be unique")
seen_enum.add(key)
for keyword in (
"minProperties", "maxProperties", "minItems", "maxItems", "minLength", "maxLength"
"minProperties",
"maxProperties",
"minItems",
"maxItems",
"minLength",
"maxLength",
):
value = schema.get(keyword)
if value is not None and (isinstance(value, bool) or not isinstance(value, int) or value < 0):
if value is not None and (
isinstance(value, bool) or not isinstance(value, int) or value < 0
):
raise _schema_error(f"{keyword} must be a non-negative integer")
for minimum, maximum in (
("minProperties", "maxProperties"),
("minItems", "maxItems"),
("minLength", "maxLength"),
):
if minimum in schema and maximum in schema and schema[minimum] > schema[maximum]: # type: ignore[operator]
if (
minimum in schema
and maximum in schema
and schema[minimum] > schema[maximum] # type: ignore[operator]
):
raise _schema_error(f"{minimum} must not exceed {maximum}")
for keyword in (
"minimum", "maximum", "exclusiveMinimum", "exclusiveMaximum", "multipleOf"
"minimum",
"maximum",
"exclusiveMinimum",
"exclusiveMaximum",
"multipleOf",
):
value = schema.get(keyword)
if value is not None and (
@@ -262,7 +291,9 @@ def _type_matches(value: object, expected: str) -> bool:
def _failure(path: str, reason: str) -> ParameterValidationError:
return ParameterValidationError(f"job parameters violate their schema at {path}: {reason}")
return ParameterValidationError(
f"job parameters violate their schema at {path}: {reason}"
)
def validate_parameter_instance(
@@ -272,6 +303,11 @@ def validate_parameter_instance(
path: str = "$",
_depth: int = 0,
) -> None:
"""Validate one parameter value against a schema subset node.
Raises ``ParameterValidationError`` (a ``ValueError``) with a sanitized
JSON path on the first violation.
"""
if _depth > 64:
raise _failure(path, "nesting exceeds 64 levels")
raw_type = schema.get("type")
@@ -280,7 +316,8 @@ def validate_parameter_instance(
if not any(_type_matches(value, item) for item in expected):
raise _failure(path, "type mismatch")
if "enum" in schema and not any(
_json_equal(value, candidate) for candidate in schema["enum"] # type: ignore[union-attr]
_json_equal(value, candidate)
for candidate in schema["enum"] # type: ignore[union-attr]
):
raise _failure(path, "value is outside enum")
if "const" in schema and not _json_equal(value, schema["const"]):
@@ -334,7 +371,9 @@ def validate_parameter_instance(
elif additional is False:
raise _failure(path, f"unknown field {name}")
elif isinstance(additional, Mapping):
validate_parameter_instance(child, additional, path=f"{path}.{name}", _depth=_depth + 1)
validate_parameter_instance(
child, additional, path=f"{path}.{name}", _depth=_depth + 1
)
if isinstance(value, (list, tuple)):
minimum = schema.get("minItems")
maximum = schema.get("maxItems")
+284 -92
View File
@@ -12,7 +12,7 @@ from decimal import Decimal
from dataclasses import dataclass, field
from enum import Enum
from types import MappingProxyType
from typing import Any, Callable, Iterable, Mapping, Sequence
from typing import Any, Callable, Iterable, Mapping, Sequence, cast
from ._validation import (
enum_value,
@@ -34,6 +34,12 @@ from .manifest import TrustMode
class VerificationStatus(str, Enum):
"""The outcome of verification.
Only ``ACCEPTED`` satisfies a stage; ``INCONCLUSIVE`` must never be
treated as success by a reducer default.
"""
ACCEPTED = "accepted"
REJECTED = "rejected"
INCONCLUSIVE = "inconclusive"
@@ -41,6 +47,12 @@ class VerificationStatus(str, Enum):
@dataclass(frozen=True, slots=True)
class VerificationDecision:
"""An immutable verifier outcome with bounded sanitized evidence.
``accepted_digest`` is set only for accepted decisions; evidence is
limited to 16 KiB and rejects local paths and transport URLs.
"""
status: VerificationStatus
verifier: ComponentRef
reason_code: str
@@ -48,12 +60,29 @@ class VerificationDecision:
accepted_digest: str | None = None
def __post_init__(self) -> None:
object.__setattr__(self, "status", enum_value(VerificationStatus, self.status, "verification.status"))
object.__setattr__(
self,
"status",
enum_value(VerificationStatus, self.status, "verification.status"),
)
if not isinstance(self.verifier, ComponentRef):
raise ValueError("verification verifier must be a ComponentRef")
object.__setattr__(self, "reason_code", require_identifier(self.reason_code, "verification.reason_code"))
evidence = freeze_json_mapping(self.evidence, "verification.evidence", forbid_locations=True)
if len(json.dumps(thaw_json(evidence), sort_keys=True, allow_nan=False).encode("utf-8")) > 16_384:
object.__setattr__(
self,
"reason_code",
require_identifier(self.reason_code, "verification.reason_code"),
)
evidence = freeze_json_mapping(
self.evidence, "verification.evidence", forbid_locations=True
)
if (
len(
json.dumps(thaw_json(evidence), sort_keys=True, allow_nan=False).encode(
"utf-8"
)
)
> 16_384
):
raise ValueError("verification evidence exceeds 16 KiB")
object.__setattr__(self, "evidence", evidence)
if self.accepted_digest is not None:
@@ -64,7 +93,10 @@ class VerificationDecision:
)
if self.status is VerificationStatus.ACCEPTED and self.accepted_digest is None:
raise ValueError("accepted verification requires an accepted_digest")
if self.status is not VerificationStatus.ACCEPTED and self.accepted_digest is not None:
if (
self.status is not VerificationStatus.ACCEPTED
and self.accepted_digest is not None
):
raise ValueError("only accepted verification may carry an accepted_digest")
def to_dict(self) -> dict[str, object]:
@@ -126,7 +158,9 @@ class VerificationBinding:
object.__setattr__(
self,
"package_digest",
require_sha256(self.package_digest, "verification package_digest", prefixed=True),
require_sha256(
self.package_digest, "verification package_digest", prefixed=True
),
)
object.__setattr__(
self,
@@ -173,7 +207,9 @@ class VerificationBinding:
)
for name, version in selected_features.items():
require_identifier(name, "verification selected feature")
require_string(version, "verification selected feature version", max_length=32)
require_string(
version, "verification selected feature version", max_length=32
)
parse_release(version, "verification selected feature version")
for name, fallback in optional_fallbacks.items():
require_identifier(name, "verification fallback feature")
@@ -183,19 +219,28 @@ class VerificationBinding:
object.__setattr__(self, "selected_features", selected_features)
object.__setattr__(self, "optional_fallbacks", optional_fallbacks)
from ._validation import require_uuid
object.__setattr__(self, "job_id", require_uuid(self.job_id, "verification job_id"))
object.__setattr__(self, "task_id", require_uuid(self.task_id, "verification task_id"))
object.__setattr__(
self, "job_id", require_uuid(self.job_id, "verification job_id")
)
object.__setattr__(
self, "task_id", require_uuid(self.task_id, "verification task_id")
)
if not isinstance(self.verifier, ComponentRef):
raise ValueError("verification binding verifier must be a ComponentRef")
object.__setattr__(
self,
"sdk_api_version",
require_string(self.sdk_api_version, "verification sdk_api_version", max_length=32),
require_string(
self.sdk_api_version, "verification sdk_api_version", max_length=32
),
)
object.__setattr__(
self,
"protocol_version",
require_string(self.protocol_version, "verification protocol_version", max_length=32),
require_string(
self.protocol_version, "verification protocol_version", max_length=32
),
)
parse_release(self.sdk_api_version, "verification sdk_api_version")
parse_release(self.protocol_version, "verification protocol_version")
@@ -217,11 +262,15 @@ class VerificationBinding:
)
schemas = tuple(self.artifact_schemas)
if not schemas or any(not isinstance(schema, SchemaRef) for schema in schemas):
raise ValueError("verification artifact_schemas must contain schema identities")
raise ValueError(
"verification artifact_schemas must contain schema identities"
)
if len(schemas) != len(set(schemas)) or schemas != tuple(
sorted(schemas, key=lambda schema: schema.canonical)
):
raise ValueError("verification artifact_schemas must be unique and canonical")
raise ValueError(
"verification artifact_schemas must be unique and canonical"
)
object.__setattr__(self, "artifact_schemas", schemas)
try:
trust_mode = TrustMode(self.trust_mode)
@@ -283,14 +332,25 @@ class VerificationBinding:
if not isinstance(value, Mapping):
raise ValueError("verification binding must be an object")
fields = {
"workload", "task_key", "package_digest", "manifest_digest",
"environment_digest", "parameters_digest", "input_collection_digest",
"workload",
"task_key",
"package_digest",
"manifest_digest",
"environment_digest",
"parameters_digest",
"input_collection_digest",
"execution_contract_digest",
"selected_features", "optional_fallbacks",
"job_id", "task_id",
"verifier", "sdk_api_version", "protocol_version",
"manifest_schema_version", "workflow_schema_version",
"artifact_schemas", "trust_mode",
"selected_features",
"optional_fallbacks",
"job_id",
"task_id",
"verifier",
"sdk_api_version",
"protocol_version",
"manifest_schema_version",
"workflow_schema_version",
"artifact_schemas",
"trust_mode",
}
require_exact_keys(value, fields, "verification binding")
schemas = value["artifact_schemas"]
@@ -321,6 +381,12 @@ class VerificationBinding:
@dataclass(frozen=True, slots=True)
class VerifyContext:
"""Coordinator-owned verification inputs: expected outputs, byte budget,
quorum size, an optional reference, and the binding for non-trusted modes.
Multi-vote contexts automatically require distinct authenticated owners.
"""
expected_outputs: Mapping[str, PortSpec]
max_output_bytes: int
minimum_matches: int = 1
@@ -339,8 +405,16 @@ class VerifyContext:
raise ValueError("expected_outputs values must be PortSpec values")
ports[canonical] = port
object.__setattr__(self, "expected_outputs", MappingProxyType(ports))
object.__setattr__(self, "max_output_bytes", require_positive_int(self.max_output_bytes, "max_output_bytes"))
object.__setattr__(self, "minimum_matches", require_positive_int(self.minimum_matches, "minimum_matches"))
object.__setattr__(
self,
"max_output_bytes",
require_positive_int(self.max_output_bytes, "max_output_bytes"),
)
object.__setattr__(
self,
"minimum_matches",
require_positive_int(self.minimum_matches, "minimum_matches"),
)
if not isinstance(self.require_distinct_owners, bool):
raise ValueError("require_distinct_owners must be a boolean")
# A multi-vote quorum is never allowed to fall back to anonymous
@@ -348,7 +422,9 @@ class VerifyContext:
# convenient, while every quorum must carry coordinator-owned owners.
if self.minimum_matches > 1:
object.__setattr__(self, "require_distinct_owners", True)
if self.binding is not None and not isinstance(self.binding, VerificationBinding):
if self.binding is not None and not isinstance(
self.binding, VerificationBinding
):
raise ValueError("binding must be a VerificationBinding")
try:
trust_mode = TrustMode(self.trust_mode)
@@ -356,21 +432,29 @@ class VerifyContext:
raise ValueError("verification trust_mode is unsupported") from error
object.__setattr__(self, "trust_mode", trust_mode)
if self.binding is not None and self.binding.trust_mode is not trust_mode:
raise ValueError("verification context trust mode does not match its binding")
raise ValueError(
"verification context trust mode does not match its binding"
)
if trust_mode is not TrustMode.TRUSTED and self.binding is None:
raise ValueError("non-trusted verification requires a coordinator binding")
if trust_mode is TrustMode.UNTRUSTED_QUORUM:
if self.minimum_matches < 2:
raise ValueError("untrusted quorum requires at least two matching owners")
raise ValueError(
"untrusted quorum requires at least two matching owners"
)
object.__setattr__(self, "require_distinct_owners", True)
if self.require_distinct_owners and self.binding is None:
raise ValueError("multi-owner verification requires a coordinator binding")
if self.reference is not None:
if not isinstance(self.reference, OutputManifest):
raise ValueError("reference must be an OutputManifest")
self.reference.validate_against(self.expected_outputs, max_output_bytes=self.max_output_bytes)
self.reference.validate_against(
self.expected_outputs, max_output_bytes=self.max_output_bytes
)
if self.binding is not None and not self.binding.matches(self.reference):
raise ValueError("reference output does not match the coordinator binding")
raise ValueError(
"reference output does not match the coordinator binding"
)
_CANDIDATE_ID_PATTERN = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_.:-]{0,127}$")
@@ -432,9 +516,15 @@ class CandidateOutput:
).encode("utf-8")
def authenticated_by(self, key: bytes) -> bool:
if self.authentication_tag is None or not isinstance(key, bytes) or len(key) < 32:
if (
self.authentication_tag is None
or not isinstance(key, bytes)
or len(key) < 32
):
return False
expected = hmac.new(key, self._authentication_payload(), hashlib.sha256).hexdigest()
expected = hmac.new(
key, self._authentication_payload(), hashlib.sha256
).hexdigest()
return hmac.compare_digest(self.authentication_tag, expected)
@classmethod
@@ -451,7 +541,9 @@ class CandidateOutput:
cannot confer this process-local authority marker.
"""
if not isinstance(authentication_key, bytes) or len(authentication_key) < 32:
raise ValueError("candidate authentication key must contain at least 32 bytes")
raise ValueError(
"candidate authentication key must contain at least 32 bytes"
)
unsigned = cls(candidate_id, owner_id, manifest)
tag = hmac.new(
authentication_key,
@@ -531,7 +623,7 @@ class CandidateOutputs:
else:
if any(not isinstance(value, CandidateOutput) for value in values):
raise ValueError("candidates must contain CandidateOutput values")
normalized = values
normalized = cast(tuple[CandidateOutput, ...], values)
candidate_ids = [value.candidate_id for value in normalized]
if len(candidate_ids) != len(set(candidate_ids)):
raise ValueError("candidate_id values must be unique")
@@ -553,7 +645,9 @@ class CandidateOutputs:
candidates = value["candidates"]
if not isinstance(candidates, list):
raise ValueError("candidate outputs candidates must be an array")
return cls(candidates=tuple(CandidateOutput.from_dict(item) for item in candidates))
return cls(
candidates=tuple(CandidateOutput.from_dict(item) for item in candidates)
)
@classmethod
def from_authenticated_dict(
@@ -568,7 +662,9 @@ class CandidateOutputs:
"""
decoded = cls.from_dict(value)
if not isinstance(authentication_key, bytes) or len(authentication_key) < 32:
raise ValueError("candidate authentication key must contain at least 32 bytes")
raise ValueError(
"candidate authentication key must contain at least 32 bytes"
)
for candidate in decoded.candidates:
if not candidate.authenticated_by(authentication_key):
raise ValueError("candidate envelope authentication failed")
@@ -588,8 +684,7 @@ def _authentication_failure(
):
return None
invalid = sum(
candidate.owner_id is None
or not candidate.coordinator_authenticated
candidate.owner_id is None or not candidate.coordinator_authenticated
for candidate in candidates.candidates
)
if invalid:
@@ -612,7 +707,9 @@ def _verify_loaded_candidates(
compare: Callable[[OutputManifest, OutputManifest], VerificationDecision],
) -> VerificationDecision:
"""Apply a package-owned loader/comparator without trusting vote replay."""
if not isinstance(context, VerifyContext) or not isinstance(candidates, CandidateOutputs):
if not isinstance(context, VerifyContext) or not isinstance(
candidates, CandidateOutputs
):
raise ValueError("verifier requires VerifyContext and CandidateOutputs")
authentication_failure = _authentication_failure(context, candidates, identity)
if authentication_failure is not None:
@@ -643,7 +740,9 @@ def _verify_loaded_candidates(
continue
seen_owners.add(candidate.owner_id)
try:
if context.binding is not None and not context.binding.matches(candidate.manifest):
if context.binding is not None and not context.binding.matches(
candidate.manifest
):
raise ValueError("candidate does not match the coordinator binding")
candidate.manifest.validate_against(
context.expected_outputs,
@@ -686,6 +785,14 @@ def _verify_loaded_candidates(
class ExactArtifactVerifier:
"""Whole-artifact SHA-256 acceptance for byte-exact workloads.
Compares logical port/collection/schema/content digests while ignoring
coordinator UUIDs, timestamps, and worker identity; counts at most one
vote per owner and accepts only a declared reference match or an
unambiguous quorum.
"""
identity = ComponentRef("exact-artifact", 1)
configuration: Mapping[str, object] = MappingProxyType({})
@@ -694,9 +801,15 @@ class ExactArtifactVerifier:
context: VerifyContext,
candidates: CandidateOutputs,
) -> VerificationDecision:
if not isinstance(context, VerifyContext) or not isinstance(candidates, CandidateOutputs):
raise ValueError("exact verifier requires VerifyContext and CandidateOutputs")
authentication_failure = _authentication_failure(context, candidates, self.identity)
if not isinstance(context, VerifyContext) or not isinstance(
candidates, CandidateOutputs
):
raise ValueError(
"exact verifier requires VerifyContext and CandidateOutputs"
)
authentication_failure = _authentication_failure(
context, candidates, self.identity
)
if authentication_failure is not None:
return authentication_failure
if context.require_distinct_owners:
@@ -718,7 +831,9 @@ class ExactArtifactVerifier:
invalid = 0
for candidate in candidates.candidates:
try:
if context.binding is not None and not context.binding.matches(candidate.manifest):
if context.binding is not None and not context.binding.matches(
candidate.manifest
):
raise ValueError("candidate does not match the coordinator binding")
candidate.manifest.validate_against(
context.expected_outputs,
@@ -735,7 +850,10 @@ class ExactArtifactVerifier:
else VerificationStatus.INCONCLUSIVE,
self.identity,
"no-valid-candidates" if candidates.candidates else "no-candidates",
{"candidate_count": len(candidates.candidates), "invalid_count": invalid},
{
"candidate_count": len(candidates.candidates),
"invalid_count": invalid,
},
)
owner_digests: dict[str, set[str]] = {}
for candidate in valid:
@@ -819,12 +937,14 @@ class ExactArtifactVerifier:
VerificationStatus.ACCEPTED,
self.identity,
"quorum-match",
with_duplicate_evidence({
"matched": matches,
"required": context.minimum_matches,
"distinct_digests": len(counts),
"invalid_count": invalid,
}),
with_duplicate_evidence(
{
"matched": matches,
"required": context.minimum_matches,
"distinct_digests": len(counts),
"invalid_count": invalid,
}
),
digest,
)
if tied and matches >= context.minimum_matches:
@@ -837,12 +957,14 @@ class ExactArtifactVerifier:
status,
self.identity,
reason,
with_duplicate_evidence({
"largest_group": matches,
"required": context.minimum_matches,
"distinct_digests": len(counts),
"invalid_count": invalid,
}),
with_duplicate_evidence(
{
"largest_group": matches,
"required": context.minimum_matches,
"distinct_digests": len(counts),
"invalid_count": invalid,
}
),
)
@@ -865,8 +987,7 @@ def _numeric_digest_value(value: object, depth: int = 0) -> object:
if any(not isinstance(key, str) for key in value):
raise ValueError("numeric objects must use JSON string keys")
return {
key: _numeric_digest_value(child, depth + 1)
for key, child in value.items()
key: _numeric_digest_value(child, depth + 1) for key, child in value.items()
}
if isinstance(value, (list, tuple)):
return [_numeric_digest_value(child, depth + 1) for child in value]
@@ -911,6 +1032,9 @@ def _decimal_evidence(value: Decimal) -> int | float | str:
@dataclass(frozen=True, slots=True)
class NumericTolerance:
"""Bounded numeric comparison policy: absolute/relative/ULP tolerances,
NaN policy, and a maximum element count for structured values."""
absolute: float = 0.0
relative: float = 0.0
max_ulps: int = 0
@@ -920,8 +1044,14 @@ class NumericTolerance:
def __post_init__(self) -> None:
for field in ("absolute", "relative"):
value = getattr(self, field)
if isinstance(value, bool) or not isinstance(value, (int, float)) or value < 0:
raise ValueError(f"numeric tolerance {field} must be a finite non-negative number")
if (
isinstance(value, bool)
or not isinstance(value, (int, float))
or value < 0
):
raise ValueError(
f"numeric tolerance {field} must be a finite non-negative number"
)
try:
converted = float(value)
except (OverflowError, ValueError) as error:
@@ -929,10 +1059,18 @@ class NumericTolerance:
f"numeric tolerance {field} must be a finite non-negative number"
) from error
if not math.isfinite(converted):
raise ValueError(f"numeric tolerance {field} must be a finite non-negative number")
raise ValueError(
f"numeric tolerance {field} must be a finite non-negative number"
)
object.__setattr__(self, field, converted)
if isinstance(self.max_ulps, bool) or not isinstance(self.max_ulps, int) or self.max_ulps < 0:
raise ValueError("numeric tolerance max_ulps must be a non-negative integer")
if (
isinstance(self.max_ulps, bool)
or not isinstance(self.max_ulps, int)
or self.max_ulps < 0
):
raise ValueError(
"numeric tolerance max_ulps must be a non-negative integer"
)
if self.nan_policy not in {"reject", "equal"}:
raise ValueError("numeric tolerance nan_policy must be reject or equal")
object.__setattr__(
@@ -943,6 +1081,13 @@ class NumericTolerance:
class NumericToleranceVerifier:
"""Reference-based structured numeric comparison verifier.
Requires a package-owned ``value_loader`` to turn artifacts into bounded
structured values; without one, verification returns ``inconclusive``
rather than accepting bytes it did not parse.
"""
identity = ComponentRef("numeric-tolerance", 1)
def __init__(
@@ -974,8 +1119,12 @@ class NumericToleranceVerifier:
context: VerifyContext,
candidates: CandidateOutputs,
) -> VerificationDecision:
if not isinstance(context, VerifyContext) or not isinstance(candidates, CandidateOutputs):
raise ValueError("numeric verifier requires VerifyContext and CandidateOutputs")
if not isinstance(context, VerifyContext) or not isinstance(
candidates, CandidateOutputs
):
raise ValueError(
"numeric verifier requires VerifyContext and CandidateOutputs"
)
if self._value_loader is None:
return VerificationDecision(
VerificationStatus.INCONCLUSIVE,
@@ -984,7 +1133,9 @@ class NumericToleranceVerifier:
{"candidate_count": len(candidates.candidates)},
)
def compare(reference: OutputManifest, candidate: OutputManifest) -> VerificationDecision:
def compare(
reference: OutputManifest, candidate: OutputManifest
) -> VerificationDecision:
assert self._value_loader is not None
return self.verify_values(
self._value_loader(reference),
@@ -1045,15 +1196,20 @@ class NumericToleranceVerifier:
difference_int = abs(actual - expected)
allowed_decimal = max(
Decimal(str(self.tolerance.absolute)),
Decimal(str(self.tolerance.relative)) * Decimal(max(abs(expected), abs(actual))),
Decimal(str(self.tolerance.relative))
* Decimal(max(abs(expected), abs(actual))),
)
if Decimal(difference_int) <= allowed_decimal:
return None
return "numeric-mismatch", location, {
"absolute_error": difference_int,
"allowed_error": _decimal_evidence(allowed_decimal),
"ulp_distance": 0,
}
return (
"numeric-mismatch",
location,
{
"absolute_error": difference_int,
"allowed_error": _decimal_evidence(allowed_decimal),
"ulp_distance": 0,
},
)
if (
isinstance(expected, int)
and abs(expected).bit_length() > 1024
@@ -1064,7 +1220,11 @@ class NumericToleranceVerifier:
left = float(expected) if isinstance(expected, int) else expected
right = float(actual) if isinstance(actual, int) else actual
if math.isnan(left) or math.isnan(right):
if self.tolerance.nan_policy == "equal" and math.isnan(left) and math.isnan(right):
if (
self.tolerance.nan_policy == "equal"
and math.isnan(left)
and math.isnan(right)
):
return None
return "nan-policy", location, {}
if not math.isfinite(left) or not math.isfinite(right):
@@ -1094,39 +1254,61 @@ class NumericToleranceVerifier:
ulp_distance is not None and ulp_distance <= self.tolerance.max_ulps
):
return None
return "numeric-mismatch", location, {
"absolute_error": _decimal_evidence(difference_decimal),
"allowed_error": _decimal_evidence(allowed_decimal),
"ulp_distance": ulp_distance if ulp_distance is not None else 0,
}
return (
"numeric-mismatch",
location,
{
"absolute_error": _decimal_evidence(difference_decimal),
"allowed_error": _decimal_evidence(allowed_decimal),
"ulp_distance": ulp_distance if ulp_distance is not None else 0,
},
)
if isinstance(expected, Mapping) and isinstance(actual, Mapping):
if any(not isinstance(key, str) for key in expected) or any(
not isinstance(key, str) for key in actual
):
return "type-mismatch", location, {}
if set(expected) != set(actual):
return "shape-mismatch", location, {
"missing_keys": sorted(str(key) for key in set(expected) - set(actual))[:32],
"extra_keys": sorted(str(key) for key in set(actual) - set(expected))[:32],
}
return (
"shape-mismatch",
location,
{
"missing_keys": sorted(
str(key) for key in set(expected) - set(actual)
)[:32],
"extra_keys": sorted(
str(key) for key in set(actual) - set(expected)
)[:32],
},
)
for key in sorted(expected, key=str):
mismatch = self._compare(expected[key], actual[key], f"{location}.{key}", depth + 1)
mismatch = self._compare(
expected[key], actual[key], f"{location}.{key}", depth + 1
)
if mismatch is not None:
return mismatch
return None
if isinstance(expected, (list, tuple)) and isinstance(actual, (list, tuple)):
if len(expected) != len(actual):
return "shape-mismatch", location, {"expected_length": len(expected), "actual_length": len(actual)}
return (
"shape-mismatch",
location,
{"expected_length": len(expected), "actual_length": len(actual)},
)
for index, (left, right) in enumerate(zip(expected, actual)):
mismatch = self._compare(left, right, f"{location}[{index}]", depth + 1)
if mismatch is not None:
return mismatch
return None
if type(expected) is not type(actual):
return "type-mismatch", location, {
"expected_type": type(expected).__name__,
"actual_type": type(actual).__name__,
}
return (
"type-mismatch",
location,
{
"expected_type": type(expected).__name__,
"actual_type": type(actual).__name__,
},
)
return None if expected == actual else ("value-mismatch", location, {})
@@ -1159,8 +1341,12 @@ class CanonicalRecordVerifier:
context: VerifyContext,
candidates: CandidateOutputs,
) -> VerificationDecision:
if not isinstance(context, VerifyContext) or not isinstance(candidates, CandidateOutputs):
raise ValueError("canonical verifier requires VerifyContext and CandidateOutputs")
if not isinstance(context, VerifyContext) or not isinstance(
candidates, CandidateOutputs
):
raise ValueError(
"canonical verifier requires VerifyContext and CandidateOutputs"
)
if self._record_loader is None:
return VerificationDecision(
VerificationStatus.INCONCLUSIVE,
@@ -1169,7 +1355,9 @@ class CanonicalRecordVerifier:
{"candidate_count": len(candidates.candidates)},
)
def compare(reference: OutputManifest, candidate: OutputManifest) -> VerificationDecision:
def compare(
reference: OutputManifest, candidate: OutputManifest
) -> VerificationDecision:
assert self._record_loader is not None
return self.verify_records(
self._record_loader(reference),
@@ -1178,12 +1366,16 @@ class CanonicalRecordVerifier:
return _verify_loaded_candidates(context, candidates, self.identity, compare)
def verify_records(self, expected: Iterable[object], actual: Iterable[object]) -> VerificationDecision:
def verify_records(
self, expected: Iterable[object], actual: Iterable[object]
) -> VerificationDecision:
expected_digest = hashlib.sha256()
actual_digest = hashlib.sha256()
counts = [0, 0]
try:
for index, (stream, digest) in enumerate(((expected, expected_digest), (actual, actual_digest))):
for index, (stream, digest) in enumerate(
((expected, expected_digest), (actual, actual_digest))
):
for record in stream:
counts[index] += 1
if counts[index] > self._max_records:
+290 -64
View File
@@ -24,6 +24,14 @@ from .resources import ResourceRequirements
class StageKind(str, Enum):
"""The kind of a workflow stage.
``PLAN`` stages expand dynamically, ``MAP`` stages fan out, ``REDUCE``
stages fan in, and the advanced kinds (loops, streams, services, side
effects) are declared but fail negotiation unless the runtime advertises
the corresponding features.
"""
PLAN = "plan"
MAP = "map"
REDUCE = "reduce"
@@ -35,6 +43,12 @@ class StageKind(str, Enum):
class WorkflowFailurePolicy(str, Enum):
"""How a workflow behaves when a stage fails.
``FAIL_FAST`` aborts on the first failure; the remaining policies require
coordinator/runtime support and are fail-closed in v1.
"""
FAIL_FAST = "fail_fast"
CONTINUE_INDEPENDENT = "continue_independent"
ALLOW_PARTIAL = "allow_partial"
@@ -43,6 +57,11 @@ class WorkflowFailurePolicy(str, Enum):
@dataclass(frozen=True, slots=True)
class LoopSpec:
"""Bounded loop declaration for a ``LOOP_CONTROLLER`` stage.
Declared but not executable until a runtime advertises ``bounded-loops``.
"""
state_schema: SchemaRef
max_iterations: int
max_wall_seconds: int
@@ -54,16 +73,34 @@ class LoopSpec:
def __post_init__(self) -> None:
if not isinstance(self.state_schema, SchemaRef):
raise ValueError("loop state_schema must be a SchemaRef")
object.__setattr__(self, "max_iterations", require_positive_int(self.max_iterations, "loop.max_iterations"))
object.__setattr__(self, "max_wall_seconds", require_positive_int(self.max_wall_seconds, "loop.max_wall_seconds"))
object.__setattr__(self, "body_workflow", require_identifier(self.body_workflow, "loop.body_workflow"))
object.__setattr__(
self,
"max_iterations",
require_positive_int(self.max_iterations, "loop.max_iterations"),
)
object.__setattr__(
self,
"max_wall_seconds",
require_positive_int(self.max_wall_seconds, "loop.max_wall_seconds"),
)
object.__setattr__(
self,
"body_workflow",
require_identifier(self.body_workflow, "loop.body_workflow"),
)
if not isinstance(self.continue_when, ComponentRef):
raise ValueError("loop continue_when must be a ComponentRef")
object.__setattr__(self, "checkpoint_every", require_positive_int(self.checkpoint_every, "loop.checkpoint_every"))
object.__setattr__(
self,
"checkpoint_every",
require_positive_int(self.checkpoint_every, "loop.checkpoint_every"),
)
if self.checkpoint_every > self.max_iterations:
raise ValueError("loop checkpoint_every must not exceed max_iterations")
if self.on_limit not in {"fail", "accept-best", "return-inconclusive"}:
raise ValueError("loop on_limit must be fail, accept-best, or return-inconclusive")
raise ValueError(
"loop on_limit must be fail, accept-best, or return-inconclusive"
)
def to_dict(self) -> dict[str, object]:
return {
@@ -81,8 +118,13 @@ class LoopSpec:
if not isinstance(value, Mapping):
raise ValueError("loop specification must be an object")
fields = {
"state_schema", "max_iterations", "max_wall_seconds", "body_workflow",
"continue_when", "checkpoint_every", "on_limit",
"state_schema",
"max_iterations",
"max_wall_seconds",
"body_workflow",
"continue_when",
"checkpoint_every",
"on_limit",
}
require_exact_keys(value, fields, "loop specification")
return cls(
@@ -98,6 +140,12 @@ class LoopSpec:
@dataclass(frozen=True, slots=True)
class StreamSpec:
"""Bounded stream declaration for a ``STREAM`` stage.
Declared but not executable until a runtime advertises
``stream-checkpoints``.
"""
source: str
partitioning: str
checkpoint_schema: SchemaRef
@@ -108,20 +156,40 @@ class StreamSpec:
max_windows: int
def __post_init__(self) -> None:
object.__setattr__(self, "source", require_identifier(self.source, "stream.source"))
object.__setattr__(self, "partitioning", require_identifier(self.partitioning, "stream.partitioning"))
object.__setattr__(
self, "source", require_identifier(self.source, "stream.source")
)
object.__setattr__(
self,
"partitioning",
require_identifier(self.partitioning, "stream.partitioning"),
)
if not isinstance(self.checkpoint_schema, SchemaRef):
raise ValueError("stream checkpoint_schema must be a SchemaRef")
object.__setattr__(self, "window_seconds", require_positive_int(self.window_seconds, "stream.window_seconds"))
object.__setattr__(self, "watermark_seconds", require_nonnegative_int(self.watermark_seconds, "stream.watermark_seconds"))
object.__setattr__(
self,
"window_seconds",
require_positive_int(self.window_seconds, "stream.window_seconds"),
)
object.__setattr__(
self,
"watermark_seconds",
require_nonnegative_int(self.watermark_seconds, "stream.watermark_seconds"),
)
object.__setattr__(
self,
"backpressure_limit",
require_positive_int(self.backpressure_limit, "stream.backpressure_limit"),
)
if self.delivery_guarantee not in {"at_least_once", "exactly_once"}:
raise ValueError("stream delivery_guarantee must be at_least_once or exactly_once")
object.__setattr__(self, "max_windows", require_positive_int(self.max_windows, "stream.max_windows"))
raise ValueError(
"stream delivery_guarantee must be at_least_once or exactly_once"
)
object.__setattr__(
self,
"max_windows",
require_positive_int(self.max_windows, "stream.max_windows"),
)
def to_dict(self) -> dict[str, object]:
return {
@@ -140,8 +208,14 @@ class StreamSpec:
if not isinstance(value, Mapping):
raise ValueError("stream specification must be an object")
fields = {
"source", "partitioning", "checkpoint_schema", "window_seconds",
"watermark_seconds", "backpressure_limit", "delivery_guarantee", "max_windows",
"source",
"partitioning",
"checkpoint_schema",
"window_seconds",
"watermark_seconds",
"backpressure_limit",
"delivery_guarantee",
"max_windows",
}
require_exact_keys(value, fields, "stream specification")
return cls(
@@ -158,6 +232,11 @@ class StreamSpec:
@dataclass(frozen=True, slots=True)
class GangSpec:
"""Co-scheduled replica group for one stage.
Declared but not executable until a runtime advertises ``gang-leases``.
"""
replicas: int
per_replica_resources: ResourceRequirements
same_topology_group: bool = False
@@ -165,7 +244,9 @@ class GangSpec:
failure_mode: str = "fail_all"
def __post_init__(self) -> None:
object.__setattr__(self, "replicas", require_positive_int(self.replicas, "gang.replicas"))
object.__setattr__(
self, "replicas", require_positive_int(self.replicas, "gang.replicas")
)
if self.replicas < 2:
raise ValueError("gang execution requires at least two replicas")
if not isinstance(self.per_replica_resources, ResourceRequirements):
@@ -173,7 +254,11 @@ class GangSpec:
if not isinstance(self.same_topology_group, bool):
raise ValueError("gang same_topology_group must be a boolean")
if self.bandwidth_class is not None:
object.__setattr__(self, "bandwidth_class", require_identifier(self.bandwidth_class, "gang.bandwidth_class"))
object.__setattr__(
self,
"bandwidth_class",
require_identifier(self.bandwidth_class, "gang.bandwidth_class"),
)
if self.failure_mode != "fail_all":
raise ValueError("SDK v1 gang failure_mode must be fail_all")
@@ -191,13 +276,18 @@ class GangSpec:
if not isinstance(value, Mapping):
raise ValueError("gang specification must be an object")
fields = {
"replicas", "per_replica_resources", "same_topology_group",
"bandwidth_class", "failure_mode",
"replicas",
"per_replica_resources",
"same_topology_group",
"bandwidth_class",
"failure_mode",
}
require_exact_keys(value, fields, "gang specification")
return cls(
replicas=value["replicas"], # type: ignore[arg-type]
per_replica_resources=ResourceRequirements.from_dict(value["per_replica_resources"]),
per_replica_resources=ResourceRequirements.from_dict(
value["per_replica_resources"]
),
same_topology_group=value["same_topology_group"], # type: ignore[arg-type]
bandwidth_class=value["bandwidth_class"], # type: ignore[arg-type]
failure_mode=value["failure_mode"], # type: ignore[arg-type]
@@ -206,6 +296,12 @@ class GangSpec:
@dataclass(frozen=True, slots=True)
class SideEffectSpec:
"""External side-effect declaration for a ``SIDE_EFFECT`` stage.
Declared but trusted-only and not executable until a runtime advertises
``side-effect``; requires an idempotency key projected into the stage.
"""
target: str
idempotency_key_parameter: str
credential_scope: str
@@ -213,14 +309,26 @@ class SideEffectSpec:
manual_approval: bool = True
def __post_init__(self) -> None:
object.__setattr__(self, "target", require_identifier(self.target, "side_effect.target"))
object.__setattr__(
self, "target", require_identifier(self.target, "side_effect.target")
)
object.__setattr__(
self,
"idempotency_key_parameter",
require_identifier(self.idempotency_key_parameter, "side_effect.idempotency_key_parameter"),
require_identifier(
self.idempotency_key_parameter, "side_effect.idempotency_key_parameter"
),
)
object.__setattr__(
self,
"credential_scope",
require_identifier(self.credential_scope, "side_effect.credential_scope"),
)
object.__setattr__(
self,
"compensation",
require_identifier(self.compensation, "side_effect.compensation"),
)
object.__setattr__(self, "credential_scope", require_identifier(self.credential_scope, "side_effect.credential_scope"))
object.__setattr__(self, "compensation", require_identifier(self.compensation, "side_effect.compensation"))
if not isinstance(self.manual_approval, bool):
raise ValueError("side_effect.manual_approval must be a boolean")
@@ -238,7 +346,11 @@ class SideEffectSpec:
if not isinstance(value, Mapping):
raise ValueError("side-effect specification must be an object")
fields = {
"target", "idempotency_key_parameter", "credential_scope", "compensation", "manual_approval",
"target",
"idempotency_key_parameter",
"credential_scope",
"compensation",
"manual_approval",
}
require_exact_keys(value, fields, "side-effect specification")
return cls(**value) # type: ignore[arg-type]
@@ -252,9 +364,13 @@ class PortRef:
stage_id: str | None = None
def __post_init__(self) -> None:
object.__setattr__(self, "port", require_identifier(self.port, "port reference"))
object.__setattr__(
self, "port", require_identifier(self.port, "port reference")
)
if self.stage_id is not None:
object.__setattr__(self, "stage_id", require_identifier(self.stage_id, "stage reference"))
object.__setattr__(
self, "stage_id", require_identifier(self.stage_id, "stage reference")
)
def to_dict(self) -> dict[str, object]:
return {"stage_id": self.stage_id, "port": self.port}
@@ -269,6 +385,12 @@ class PortRef:
@dataclass(frozen=True, slots=True)
class ArtifactEdge:
"""A typed data flow from a source port to a stage input port.
Endpoints must declare compatible schemas; every stage input receives
exactly one edge.
"""
source: PortRef
target: PortRef
@@ -286,7 +408,10 @@ class ArtifactEdge:
if not isinstance(value, Mapping):
raise ValueError("artifact edge must be an object")
require_exact_keys(value, {"source", "target"}, "artifact edge")
return cls(source=PortRef.from_dict(value["source"]), target=PortRef.from_dict(value["target"]))
return cls(
source=PortRef.from_dict(value["source"]),
target=PortRef.from_dict(value["target"]),
)
def _port_mapping(value: Mapping[str, PortSpec], field: str) -> Mapping[str, PortSpec]:
@@ -303,6 +428,13 @@ def _port_mapping(value: Mapping[str, PortSpec], field: str) -> Mapping[str, Por
@dataclass(frozen=True, slots=True)
class StageSpec:
"""One typed stage: kind, handler entry point, ports, and policy.
``entry_point`` must be an installed handler key (runner for map/verify
stages, reducer for reduce stages); resources, execution profile, retry
policy, verifier, and trust modes are validated at construction.
"""
stage_id: str
kind: StageKind
entry_point: str
@@ -323,18 +455,29 @@ class StageSpec:
side_effect: SideEffectSpec | None = None
def __post_init__(self) -> None:
object.__setattr__(self, "stage_id", require_identifier(self.stage_id, "stage_id"))
object.__setattr__(
self, "stage_id", require_identifier(self.stage_id, "stage_id")
)
object.__setattr__(self, "kind", enum_value(StageKind, self.kind, "stage.kind"))
object.__setattr__(self, "entry_point", require_entry_point(self.entry_point, "stage.entry_point"))
object.__setattr__(
self,
"entry_point",
require_entry_point(self.entry_point, "stage.entry_point"),
)
needs = tuple(require_identifier(value, "stage.needs") for value in self.needs)
if self.stage_id in needs or len(needs) != len(set(needs)):
raise ValueError("stage.needs must contain unique other stage IDs")
object.__setattr__(self, "needs", needs)
object.__setattr__(self, "inputs", _port_mapping(self.inputs, "stage.inputs"))
object.__setattr__(self, "outputs", _port_mapping(self.outputs, "stage.outputs"))
object.__setattr__(
self, "outputs", _port_mapping(self.outputs, "stage.outputs")
)
if not self.outputs:
raise ValueError("a stage must declare at least one output port")
names = tuple(require_identifier(value, "parameter_name") for value in self.parameter_names)
names = tuple(
require_identifier(value, "parameter_name")
for value in self.parameter_names
)
if len(names) != len(set(names)):
raise ValueError("parameter_names must be unique")
object.__setattr__(self, "parameter_names", names)
@@ -347,13 +490,19 @@ class StageSpec:
raise ValueError("stage.retry must be RetryPolicy")
if self.verifier is not None and not isinstance(self.verifier, ComponentRef):
raise ValueError("stage.verifier must be a ComponentRef")
modes = tuple(require_identifier(value, "trust_mode") for value in self.trust_modes)
modes = tuple(
require_identifier(value, "trust_mode") for value in self.trust_modes
)
if not modes or len(modes) != len(set(modes)):
raise ValueError("stage.trust_modes must be non-empty and unique")
if not set(modes).issubset({"trusted", "verified", "untrusted_quorum"}):
raise ValueError("stage.trust_modes contains an unsupported trust mode")
object.__setattr__(self, "trust_modes", modes)
object.__setattr__(self, "max_fan_out", require_positive_int(self.max_fan_out, "stage.max_fan_out"))
object.__setattr__(
self,
"max_fan_out",
require_positive_int(self.max_fan_out, "stage.max_fan_out"),
)
if not isinstance(self.cacheable, bool):
raise ValueError("stage.cacheable must be a boolean")
advanced = {
@@ -370,8 +519,12 @@ class StageSpec:
if self.kind is kind and declaration is None:
raise ValueError(f"{kind.value} stage requires its bounded declaration")
if self.kind is not kind and declaration is not None:
raise ValueError(f"{kind.value} declaration is valid only for a {kind.value} stage")
if declaration is not None and not isinstance(declaration, expected_types[kind]):
raise ValueError(
f"{kind.value} declaration is valid only for a {kind.value} stage"
)
if declaration is not None and not isinstance(
declaration, expected_types[kind]
):
raise ValueError(f"{kind.value} declaration has the wrong type")
if self.gang is not None and not isinstance(self.gang, GangSpec):
raise ValueError("stage.gang must be a GangSpec")
@@ -382,11 +535,16 @@ class StageSpec:
if self.kind is StageKind.SIDE_EFFECT:
if self.cacheable:
raise ValueError("side-effect stages cannot be cached")
if self.execution.network not in {NetworkPolicy.ALLOWLISTED_EGRESS, NetworkPolicy.TRUSTED}:
if self.execution.network not in {
NetworkPolicy.ALLOWLISTED_EGRESS,
NetworkPolicy.TRUSTED,
}:
raise ValueError("side-effect stages require explicit egress")
assert self.side_effect is not None
if self.side_effect.idempotency_key_parameter not in self.parameter_names:
raise ValueError("side-effect idempotency key must be projected into the stage")
raise ValueError(
"side-effect idempotency key must be projected into the stage"
)
def to_dict(self) -> dict[str, object]:
return {
@@ -407,7 +565,9 @@ class StageSpec:
"loop": self.loop.to_dict() if self.loop is not None else None,
"stream": self.stream.to_dict() if self.stream is not None else None,
"gang": self.gang.to_dict() if self.gang is not None else None,
"side_effect": self.side_effect.to_dict() if self.side_effect is not None else None,
"side_effect": self.side_effect.to_dict()
if self.side_effect is not None
else None,
}
@classmethod
@@ -415,14 +575,31 @@ class StageSpec:
if not isinstance(value, Mapping):
raise ValueError("stage specification must be an object")
fields = {
"stage_id", "kind", "entry_point", "needs", "inputs", "outputs",
"parameter_names", "resources", "execution", "retry", "verifier",
"trust_modes", "max_fan_out", "cacheable", "loop", "stream", "gang", "side_effect",
"stage_id",
"kind",
"entry_point",
"needs",
"inputs",
"outputs",
"parameter_names",
"resources",
"execution",
"retry",
"verifier",
"trust_modes",
"max_fan_out",
"cacheable",
"loop",
"stream",
"gang",
"side_effect",
}
require_exact_keys(value, fields, "stage specification")
arrays = (value["needs"], value["parameter_names"], value["trust_modes"])
if any(not isinstance(item, list) for item in arrays):
raise ValueError("stage needs, parameter_names, and trust_modes must be arrays")
raise ValueError(
"stage needs, parameter_names, and trust_modes must be arrays"
)
inputs, outputs = value["inputs"], value["outputs"]
if not isinstance(inputs, Mapping) or not isinstance(outputs, Mapping):
raise ValueError("stage inputs and outputs must be objects")
@@ -437,19 +614,32 @@ class StageSpec:
resources=ResourceRequirements.from_dict(value["resources"]),
execution=ExecutionProfile.from_dict(value["execution"]),
retry=RetryPolicy.from_dict(value["retry"]),
verifier=None if value["verifier"] is None else ComponentRef.from_dict(value["verifier"]),
verifier=None
if value["verifier"] is None
else ComponentRef.from_dict(value["verifier"]),
trust_modes=tuple(value["trust_modes"]), # type: ignore[arg-type]
max_fan_out=value["max_fan_out"], # type: ignore[arg-type]
cacheable=value["cacheable"], # type: ignore[arg-type]
loop=None if value["loop"] is None else LoopSpec.from_dict(value["loop"]),
stream=None if value["stream"] is None else StreamSpec.from_dict(value["stream"]),
stream=None
if value["stream"] is None
else StreamSpec.from_dict(value["stream"]),
gang=None if value["gang"] is None else GangSpec.from_dict(value["gang"]),
side_effect=None if value["side_effect"] is None else SideEffectSpec.from_dict(value["side_effect"]),
side_effect=None
if value["side_effect"] is None
else SideEffectSpec.from_dict(value["side_effect"]),
)
@dataclass(frozen=True, slots=True)
class WorkflowSpec:
"""A versioned acyclic workflow: inputs, stages, edges, and outputs.
Construction validates complete input bindings, matching ``needs``
declarations, edge schema compatibility, acyclicity, and output port
resolution.
"""
workflow_id: str
inputs: Mapping[str, PortSpec]
stages: tuple[StageSpec, ...]
@@ -461,9 +651,15 @@ class WorkflowSpec:
schema_version: int = WORKFLOW_SCHEMA_VERSION
def __post_init__(self) -> None:
require_schema_version(self.schema_version, WORKFLOW_SCHEMA_VERSION, "workflow schema_version")
object.__setattr__(self, "workflow_id", require_identifier(self.workflow_id, "workflow_id"))
object.__setattr__(self, "inputs", _port_mapping(self.inputs, "workflow.inputs"))
require_schema_version(
self.schema_version, WORKFLOW_SCHEMA_VERSION, "workflow schema_version"
)
object.__setattr__(
self, "workflow_id", require_identifier(self.workflow_id, "workflow_id")
)
object.__setattr__(
self, "inputs", _port_mapping(self.inputs, "workflow.inputs")
)
stages = tuple(self.stages)
if not stages or any(not isinstance(stage, StageSpec) for stage in stages):
raise ValueError("workflow stages must contain at least one StageSpec")
@@ -489,9 +685,15 @@ class WorkflowSpec:
object.__setattr__(
self,
"failure_policy",
enum_value(WorkflowFailurePolicy, self.failure_policy, "workflow.failure_policy"),
enum_value(
WorkflowFailurePolicy, self.failure_policy, "workflow.failure_policy"
),
)
object.__setattr__(
self,
"max_tasks",
require_positive_int(self.max_tasks, "workflow.max_tasks"),
)
object.__setattr__(self, "max_tasks", require_positive_int(self.max_tasks, "workflow.max_tasks"))
object.__setattr__(
self,
"max_output_bytes",
@@ -499,12 +701,16 @@ class WorkflowSpec:
)
self._validate_graph(stage_by_id)
def _source_port(self, reference: PortRef, stages: Mapping[str, StageSpec]) -> PortSpec:
def _source_port(
self, reference: PortRef, stages: Mapping[str, StageSpec]
) -> PortSpec:
if reference.stage_id is None:
try:
return self.inputs[reference.port]
except KeyError as error:
raise ValueError(f"unknown workflow input port: {reference.port}") from error
raise ValueError(
f"unknown workflow input port: {reference.port}"
) from error
try:
stage = stages[reference.stage_id]
return stage.outputs[reference.port]
@@ -538,17 +744,23 @@ class WorkflowSpec:
and target_port.collection.value in {"ordered", "keyed", "set"}
)
if not same_schema or not (direct_match or map_fan_in):
raise ValueError("artifact edge source and target port declarations are incompatible")
raise ValueError(
"artifact edge source and target port declarations are incompatible"
)
if edge.source.stage_id is not None:
dependencies[edge.target.stage_id].add(edge.source.stage_id)
for stage in stages.values():
missing = [name for name in stage.inputs if (stage.stage_id, name) not in incoming]
missing = [
name for name in stage.inputs if (stage.stage_id, name) not in incoming
]
if missing:
raise ValueError(
f"stage {stage.stage_id} has unbound inputs: {', '.join(sorted(missing))}"
)
if dependencies[stage.stage_id] != set(stage.needs):
raise ValueError(f"stage {stage.stage_id} needs do not match its artifact edges")
raise ValueError(
f"stage {stage.stage_id} needs do not match its artifact edges"
)
remaining = {name: set(values) for name, values in dependencies.items()}
ready = sorted(name for name, values in remaining.items() if not values)
visited: list[str] = []
@@ -568,10 +780,12 @@ class WorkflowSpec:
def output_ports(self) -> Mapping[str, PortSpec]:
stages = {stage.stage_id: stage for stage in self.stages}
return MappingProxyType({
name: self._source_port(reference, stages)
for name, reference in self.outputs.items()
})
return MappingProxyType(
{
name: self._source_port(reference, stages)
for name, reference in self.outputs.items()
}
)
def to_dict(self) -> dict[str, object]:
return {
@@ -580,7 +794,9 @@ class WorkflowSpec:
"inputs": {name: port.to_dict() for name, port in self.inputs.items()},
"stages": [stage.to_dict() for stage in self.stages],
"edges": [edge.to_dict() for edge in self.edges],
"outputs": {name: reference.to_dict() for name, reference in self.outputs.items()},
"outputs": {
name: reference.to_dict() for name, reference in self.outputs.items()
},
"failure_policy": self.failure_policy.value,
"max_tasks": self.max_tasks,
"max_output_bytes": self.max_output_bytes,
@@ -591,8 +807,15 @@ class WorkflowSpec:
if not isinstance(value, Mapping):
raise ValueError("workflow specification must be an object")
fields = {
"schema_version", "workflow_id", "inputs", "stages", "edges",
"outputs", "failure_policy", "max_tasks", "max_output_bytes",
"schema_version",
"workflow_id",
"inputs",
"stages",
"edges",
"outputs",
"failure_policy",
"max_tasks",
"max_output_bytes",
}
require_exact_keys(value, fields, "workflow specification")
inputs, outputs = value["inputs"], value["outputs"]
@@ -607,7 +830,10 @@ class WorkflowSpec:
inputs={name: PortSpec.from_dict(port) for name, port in inputs.items()},
stages=tuple(StageSpec.from_dict(stage) for stage in stages),
edges=tuple(ArtifactEdge.from_dict(edge) for edge in edges),
outputs={name: PortRef.from_dict(reference) for name, reference in outputs.items()},
outputs={
name: PortRef.from_dict(reference)
for name, reference in outputs.items()
},
failure_policy=value["failure_policy"], # type: ignore[arg-type]
max_tasks=value["max_tasks"], # type: ignore[arg-type]
max_output_bytes=value["max_output_bytes"], # type: ignore[arg-type]
+23 -7
View File
@@ -22,6 +22,22 @@ def _clean_url(value: object | None) -> str | None:
return text.rstrip("/") or None
def _int_value(value: object | None, name: str) -> int | None:
if value is None:
return None
if isinstance(value, bool) or not isinstance(value, (int, float, str)):
raise ValueError(f"{name} must be a number")
return int(value)
def _float_value(value: object | None, name: str) -> float | None:
if value is None:
return None
if isinstance(value, bool) or not isinstance(value, (int, float, str)):
raise ValueError(f"{name} must be a number")
return float(value)
def _positive_number(value: object, name: str, *, allow_zero: bool = False) -> None:
if (
isinstance(value, bool)
@@ -190,16 +206,16 @@ class WorkerConfig:
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(poll_interval),
request_timeout=float(request_timeout),
heartbeat_interval=float(heartbeat_interval),
cpu_count=_int_value(cpu_count, "cpu_count") or 1,
memory_mb=_int_value(memory_mb, "memory_mb"),
poll_interval=_float_value(poll_interval, "poll_interval") or 2.0,
request_timeout=_float_value(request_timeout, "request_timeout") or 30.0,
heartbeat_interval=_float_value(heartbeat_interval, "heartbeat_interval") or 15.0,
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,
cleanup_after_seconds=_float_value(cleanup, "cleanup_after_seconds"),
max_tasks=_int_value(max_tasks, "max_tasks"),
exit_when_idle=bool(values.get("exit_when_idle", False)),
capabilities=_capabilities(capabilities),
workload_allowlist=_workload_allowlist(allowlist),
+71 -25
View File
@@ -16,7 +16,11 @@ from datetime import datetime, timezone
from .artifacts import ArtifactClient, sha256_file
from .config import WorkerConfig
from .coordinator import CoordinatorClient, CoordinatorConflictError, CoordinatorTransientError
from .coordinator import (
CoordinatorClient,
CoordinatorConflictError,
CoordinatorTransientError,
)
from .models import ClaimedTask, UploadedArtifact
from .runners import Runner
@@ -24,8 +28,11 @@ from .runners import Runner
class LeaseHeartbeat:
"""Renews a claimed task lease while local work is in progress."""
def __init__(self, task: ClaimedTask, coordinator: CoordinatorClient, config: WorkerConfig) -> None:
def __init__(
self, task: ClaimedTask, coordinator: CoordinatorClient, config: WorkerConfig
) -> None:
self.task, self.coordinator, self.config = task, coordinator, config
self._worker_id = config.worker_id or ""
self._stop = threading.Event()
self._error: Exception | None = None
self._thread: threading.Thread | None = None
@@ -33,11 +40,11 @@ class LeaseHeartbeat:
def start(self) -> None:
# Verify ownership before expensive download or calculation begins.
self._lease_expires_at = self.coordinator.heartbeat(
self.task, self.config.worker_id
)
self._lease_expires_at = self.coordinator.heartbeat(self.task, self._worker_id)
self._next_delay()
self._thread = threading.Thread(target=self._run, name=f"lease-{self.task.task_id}", daemon=True)
self._thread = threading.Thread(
target=self._run, name=f"lease-{self.task.task_id}", daemon=True
)
self._thread.start()
def stop(self) -> None:
@@ -54,10 +61,12 @@ class LeaseHeartbeat:
while not self._stop.wait(max(delay, 0.01)):
try:
self._lease_expires_at = self.coordinator.heartbeat(
self.task, self.config.worker_id
self.task, self._worker_id
)
delay = self._next_delay()
except Exception as error: # Surface the lease loss in the main state machine.
except (
Exception
) as error: # Surface the lease loss in the main state machine.
self._error = error
return
@@ -66,7 +75,9 @@ class LeaseHeartbeat:
def _seconds_until_expiry(self) -> float:
try:
expiry = datetime.fromisoformat(self._lease_expires_at.replace("Z", "+00:00"))
expiry = datetime.fromisoformat(
self._lease_expires_at.replace("Z", "+00:00")
)
except ValueError as error:
raise ValueError("invalid lease_expires_at") from error
seconds = (expiry - datetime.now(timezone.utc)).total_seconds()
@@ -84,8 +95,19 @@ class RunOnceOutcome:
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
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")
@@ -129,14 +151,20 @@ class WorkerDaemon:
)
return True
elif self.config.exit_when_idle:
self._log("stopped", reason="queue_empty", completed_tasks=completed_tasks)
self._log(
"stopped",
reason="queue_empty",
completed_tasks=completed_tasks,
)
return True
else:
self._sleep(self.config.poll_interval)
except CoordinatorTransientError as error:
failures += 1
self._log("failed", error_type=type(error).__name__)
self._sleep(min(self.config.poll_interval * 2 ** min(failures, 6), 60.0))
self._sleep(
min(self.config.poll_interval * 2 ** min(failures, 6), 60.0)
)
except KeyboardInterrupt:
self._log("stopped", reason="interrupted", completed_tasks=completed_tasks)
return False
@@ -183,11 +211,15 @@ class WorkerDaemon:
},
)
completed = True
self._log("completed", task, elapsed_seconds=round(time.monotonic() - started, 3))
self._log(
"completed", task, elapsed_seconds=round(time.monotonic() - started, 3)
)
except KeyboardInterrupt:
self._log("interrupted", task)
try:
self._report_failure(task, InterruptedError("worker interrupted by operator"))
self._report_failure(
task, InterruptedError("worker interrupted by operator")
)
except CoordinatorTransientError:
self._log("failed", task, error_type="FailureReportError")
raise
@@ -203,13 +235,16 @@ class WorkerDaemon:
def _report_failure(self, task: ClaimedTask, error: Exception) -> None:
message = self._sanitize_error_message(error)
try:
self.coordinator.fail(task, {
"worker_id": self._worker_id(),
"attempt": task.attempt,
"error_code": type(error).__name__,
"error_message": message,
"retryable": self._is_retryable(error),
})
self.coordinator.fail(
task,
{
"worker_id": self._worker_id(),
"attempt": task.attempt,
"error_code": type(error).__name__,
"error_message": message,
"retryable": self._is_retryable(error),
},
)
except CoordinatorTransientError:
raise
except Exception:
@@ -218,7 +253,9 @@ class WorkerDaemon:
@staticmethod
def _is_retryable(error: Exception) -> bool:
"""Retry transient worker/transport failures, never invalid scientific input."""
return not isinstance(error, (ValueError, FileNotFoundError, subprocess.CalledProcessError))
return not isinstance(
error, (ValueError, FileNotFoundError, subprocess.CalledProcessError)
)
def _sanitize_error_message(self, error: Exception) -> str:
"""Keep coordinator-visible failures useful without exposing local paths."""
@@ -264,12 +301,21 @@ class WorkerDaemon:
log_level: int = logging.INFO,
**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}
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.log(log_level, "worker_event %s", fields)
def _cleanup_expired_directories(self) -> None:
"""Remove only old task attempt directories when retention was configured."""
if self.config.cleanup_after_seconds is None or not self.config.work_dir.exists():
if (
self.config.cleanup_after_seconds is None
or not self.config.work_dir.exists()
):
return
cutoff = time.time() - self.config.cleanup_after_seconds
for task_dir in self.config.work_dir.iterdir():
+4 -2
View File
@@ -2,6 +2,8 @@
from __future__ import annotations
from http.client import HTTPMessage
from typing import IO
from urllib.request import HTTPRedirectHandler, Request
from urllib.parse import urlsplit
@@ -24,10 +26,10 @@ class SameOriginAuthRedirectHandler(HTTPRedirectHandler):
def redirect_request(
self,
req: Request,
fp: object,
fp: IO[bytes],
code: int,
msg: str,
headers: object,
headers: HTTPMessage,
newurl: str,
) -> Request | None:
redirected = super().redirect_request(req, fp, code, msg, headers, newurl)
@@ -9,6 +9,7 @@ partition, compute, and merge.
from __future__ import annotations
from pathlib import Path
from typing import Any, Mapping
from scimesh.sdk.artifacts import ArtifactSchema, ComponentRef, PortSpec
@@ -132,6 +133,7 @@ class DescriptorBatchWorkload(MapReduceWorkload):
raise ValueError("skip_invalid must be a boolean")
return value
def descriptor_batch_sdk_definition(
*,
shard_rows: int = 10_000,
+7 -2
View File
@@ -12,7 +12,7 @@ import os
import platform
from scimesh.sdk.identity import SDK_API_VERSION
from scimesh.sdk.registry import WorkloadRegistry
from scimesh.sdk.registry import AllowedPackage, WorkloadRegistry
from scimesh.sdk.resources import ResourceInventory
from scimesh.sdk.runtime import RuntimeCapabilities
@@ -80,7 +80,12 @@ def default_sdk_runtime(
features={"artifact-collections": "1.0.0", "exact-verifier": "1.0.0"},
workload_capabilities=(
workload_capabilities
or ("similarity-search", "similarity-graph", "descriptor-batch", "molwt-filter")
or (
"similarity-search",
"similarity-graph",
"descriptor-batch",
"molwt-filter",
)
),
inventory=ResourceInventory(
cpu_cores=max(os.cpu_count() or 1, 1),
+1 -1
View File
@@ -68,7 +68,7 @@ def filter_molecules_by_molwt(
if not skip_invalid:
raise ValueError(f"row {scanned} has an invalid canonical_smiles")
continue
molwt = Descriptors.MolWt(molecule)
molwt = Descriptors.MolWt(molecule) # type: ignore[attr-defined]
if minimum is not None and molwt < minimum:
continue
if maximum is not None and molwt > maximum: