Add workload SDK foundation

This commit is contained in:
Emil
2026-08-01 23:22:20 +03:00
parent 11e9333033
commit c43af32495
32 changed files with 12256 additions and 20 deletions
+225
View File
@@ -0,0 +1,225 @@
"""SciMesh Workload SDK v1.
The implemented profile is ``core-batch-v1``: strict manifests, typed artifact
ports, static map/reduce DAGs, resource eligibility/local reservation, exact
verification, and an adapter for existing distributed workloads. Advanced
dynamic, stream, accelerator, gang, and side-effect declarations are modeled
but fail compatibility negotiation unless an enforcing runtime advertises the
corresponding versioned features.
"""
from .artifacts import (
ArtifactCollection,
ArtifactItem,
ArtifactRef,
ArtifactSchema,
Cardinality,
CollectionKind,
OutputManifest,
PortSpec,
Provenance,
)
from .builtins import (
current_environment_digest,
current_scimesh_package_digest,
default_sdk_registry,
default_sdk_runtime,
similarity_search_sdk_adapter,
)
from .conformance import (
CancellationFlag,
LocalArtifactStore,
LocalCoreBatchExecutor,
LocalPlanningContext,
LocalTaskContext,
assert_manifest_round_trip,
)
from .execution import (
CheckpointPolicy,
ExecutionProfile,
FailureCategory,
FailureReport,
NetworkPolicy,
ProcessModel,
RetryPolicy,
)
from .identity import (
MANIFEST_SCHEMA_VERSION,
OUTPUT_SCHEMA_VERSION,
SDK_API_VERSION,
TASK_SCHEMA_VERSION,
WORKFLOW_SCHEMA_VERSION,
ComponentRef,
FeatureRequirement,
SchemaRef,
VersionRange,
WorkloadId,
)
from .integrity import installed_distribution_digest
from .manifest import (
DeterminismProfile,
EnvironmentSpec,
PackageSpec,
TrustMode,
VerifierSpec,
WorkloadLimits,
WorkloadManifest,
)
from .plans import ExpansionManifest, JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
from .protocols import (
ArtifactCatalog,
ArtifactSink,
CancellationToken,
Planner,
PlanningContext,
PlanningResources,
ReduceContext,
Reducer,
Runner,
TaskContext,
Verifier,
)
from .registry import (
AllowedPackage,
WorkloadDefinition,
WorkloadDescription,
WorkloadRegistry,
)
from .resources import (
AcceleratorDevice,
AcceleratorMode,
ResourceAllocation,
ResourceInventory,
ResourcePool,
ResourceRequirements,
ResourceUnavailableError,
)
from .runtime import (
CompatibilityError,
NegotiatedWorkload,
RuntimeCapabilities,
negotiate_manifest,
)
from .verification import (
CandidateOutput,
CandidateOutputs,
CanonicalRecordVerifier,
ExactArtifactVerifier,
NumericTolerance,
NumericToleranceVerifier,
VerificationDecision,
VerificationBinding,
VerificationStatus,
VerifyContext,
)
from .workflow import (
ArtifactEdge,
GangSpec,
LoopSpec,
PortRef,
SideEffectSpec,
StageKind,
StageSpec,
StreamSpec,
WorkflowFailurePolicy,
WorkflowSpec,
)
__all__ = [
"AcceleratorDevice",
"AcceleratorMode",
"AllowedPackage",
"ArtifactCatalog",
"ArtifactCollection",
"ArtifactEdge",
"ArtifactItem",
"ArtifactRef",
"ArtifactSchema",
"ArtifactSink",
"CancellationFlag",
"CancellationToken",
"CandidateOutput",
"CandidateOutputs",
"CanonicalRecordVerifier",
"Cardinality",
"CheckpointPolicy",
"CollectionKind",
"CompatibilityError",
"ComponentRef",
"DeterminismProfile",
"EnvironmentSpec",
"ExactArtifactVerifier",
"ExecutionProfile",
"ExpansionManifest",
"FailureCategory",
"FailureReport",
"FeatureRequirement",
"GangSpec",
"JobRequest",
"LocalArtifactStore",
"LocalCoreBatchExecutor",
"LocalPlanningContext",
"LocalTaskContext",
"LoopSpec",
"MANIFEST_SCHEMA_VERSION",
"NegotiatedWorkload",
"NetworkPolicy",
"NumericTolerance",
"NumericToleranceVerifier",
"OUTPUT_SCHEMA_VERSION",
"OutputManifest",
"PackageSpec",
"Planner",
"PlanningContext",
"PlanningResources",
"PortRef",
"PortSpec",
"ProcessModel",
"Provenance",
"ReduceContext",
"Reducer",
"ResourceAllocation",
"ResourceInventory",
"ResourcePool",
"ResourceRequirements",
"ResourceUnavailableError",
"RetryPolicy",
"Runner",
"RuntimeCapabilities",
"SDK_API_VERSION",
"SchemaRef",
"SideEffectSpec",
"StageKind",
"StageSpec",
"StreamSpec",
"TASK_SCHEMA_VERSION",
"TaskContext",
"TaskSpec",
"TrustMode",
"ValidatedJob",
"VerificationDecision",
"VerificationBinding",
"VerificationStatus",
"Verifier",
"VerifierSpec",
"VersionRange",
"VerifyContext",
"WORKFLOW_SCHEMA_VERSION",
"WorkflowFailurePolicy",
"WorkflowPlan",
"WorkflowSpec",
"WorkloadDefinition",
"WorkloadDescription",
"WorkloadId",
"WorkloadLimits",
"WorkloadManifest",
"WorkloadRegistry",
"assert_manifest_round_trip",
"current_environment_digest",
"current_scimesh_package_digest",
"default_sdk_registry",
"default_sdk_runtime",
"installed_distribution_digest",
"negotiate_manifest",
"similarity_search_sdk_adapter",
]
+380
View File
@@ -0,0 +1,380 @@
"""Internal validation helpers for strict, JSON-safe SDK value objects."""
from __future__ import annotations
import json
import math
import re
from types import MappingProxyType
from typing import Any, Mapping
from urllib.parse import unquote
from uuid import UUID
WORKLOAD_NAME_PATTERN = re.compile(r"^[a-z][a-z0-9]*(?:-[a-z0-9]+)*$")
IDENTIFIER_PATTERN = re.compile(r"^[a-z][a-z0-9]*(?:[-_.][a-z0-9]+)*$")
ENTRY_POINT_PATTERN = re.compile(
r"^[A-Za-z_][A-Za-z0-9_.]*:[A-Za-z_][A-Za-z0-9_.]*(?:@v[1-9][0-9]*)?$"
)
SEMVER_PATTERN = re.compile(
r"^(0|[1-9][0-9]*)\.(0|[1-9][0-9]*)\.(0|[1-9][0-9]*)"
r"(?:-([0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*))?"
r"(?:\+([0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*))?$"
)
_VERSION_PATTERN = re.compile(r"^(0|[1-9][0-9]*)(?:\.(0|[1-9][0-9]*))?(?:\.(0|[1-9][0-9]*))?$")
_VERSION_CLAUSE_PATTERN = re.compile(r"^(==|>=|<=|>|<)\s*(.+)$")
_FORBIDDEN_LOCATOR_PREFIXES = (
"file://",
"worker://",
"http://",
"https://",
"s3://",
"/",
)
_URI_SCHEME_PATTERN = re.compile(r"^[A-Za-z][A-Za-z0-9+.-]*:")
_WINDOWS_PATH_PATTERN = re.compile(r"^[A-Za-z]:(?:[\\/]|[^\s]*[\\/])")
_SECRET_ASSIGNMENT_PATTERN = re.compile(
r"(?i)(?:^|[^A-Za-z0-9_])"
r"(?:authorization|bearer|token|secret|password|api[-_]?key)\s*[:=]"
)
_PATH_ASSIGNMENT_PATTERN = re.compile(
r"(?i)(?:^|[^A-Za-z0-9_])"
r"(?:path|file|directory|dir|workspace|cwd|upload|download)\s*[:=]"
)
_PATH_SEGMENT_PATTERN = re.compile(r"^[A-Za-z0-9_.-]+$")
_FILE_NAME_PATTERN = re.compile(r"^[A-Za-z0-9_.-]+\.[A-Za-z0-9]{1,16}$")
_TASK_KEY_COMPONENT_PATTERN = re.compile(r"^[a-z0-9][a-z0-9_.-]*$")
def require_exact_keys(
value: Mapping[str, object],
expected: set[str],
label: str,
*,
optional: set[str] | None = None,
) -> None:
"""Reject unknown fields and report missing required fields."""
if any(not isinstance(key, str) for key in value):
raise ValueError(f"{label} must use string field names")
optional = optional or set()
actual = set(value)
missing = expected - actual
unknown = actual - expected - optional
if not missing and not unknown:
return
details: list[str] = []
if missing:
details.append("missing " + ", ".join(sorted(missing)))
if unknown:
details.append("unknown " + ", ".join(sorted(unknown)))
raise ValueError(f"{label} has invalid fields: {'; '.join(details)}")
def require_mapping(value: object, field: str) -> Mapping[str, object]:
if not isinstance(value, Mapping) or any(not isinstance(key, str) for key in value):
raise ValueError(f"{field} must be an object with string keys")
return value
def require_string(value: object, field: str, *, max_length: int = 256) -> str:
if not isinstance(value, str) or not value.strip() or len(value) > max_length:
raise ValueError(f"{field} must be a non-empty string of at most {max_length} characters")
if any(ord(character) < 32 for character in value):
raise ValueError(f"{field} must not contain control characters")
return value
def require_identifier(value: object, field: str) -> str:
text = require_string(value, field, max_length=128)
if not IDENTIFIER_PATTERN.fullmatch(text):
raise ValueError(f"{field} must be a canonical identifier")
return text
def require_workload_name(value: object, field: str = "workload.name") -> str:
text = require_string(value, field, max_length=128)
if not WORKLOAD_NAME_PATTERN.fullmatch(text):
raise ValueError(f"{field} must be a canonical hyphenated workload name")
return text
def require_entry_point(value: object, field: str) -> str:
text = require_string(value, field, max_length=256)
if not ENTRY_POINT_PATTERN.fullmatch(text):
raise ValueError(f"{field} must be a package-owned module:object entry point")
return text
def require_semver(value: object, field: str) -> str:
text = require_string(value, field, max_length=64)
match = SEMVER_PATTERN.fullmatch(text)
if match is None:
raise ValueError(f"{field} must be a semantic version such as 1.0.0")
prerelease = match.group(4)
if prerelease is not None and any(
identifier.isdigit() and len(identifier) > 1 and identifier.startswith("0")
for identifier in prerelease.split(".")
):
raise ValueError(f"{field} has a non-canonical numeric prerelease identifier")
return text
def require_uuid(value: object, field: str) -> str:
if not isinstance(value, str):
raise ValueError(f"{field} must be a UUID string")
try:
return str(UUID(value))
except ValueError as error:
raise ValueError(f"{field} must be a UUID string") from error
def require_sha256(value: object, field: str, *, prefixed: bool = False) -> str:
if not isinstance(value, str):
raise ValueError(f"{field} must be a SHA-256 digest")
digest = value[7:] if prefixed and value.startswith("sha256:") else value
if prefixed and not value.startswith("sha256:"):
raise ValueError(f"{field} must use the sha256:<hex> form")
if not re.fullmatch(r"[0-9a-f]{64}", digest):
raise ValueError(f"{field} must be a lowercase SHA-256 digest")
return f"sha256:{digest}" if prefixed else digest
def require_nonnegative_int(value: object, field: str) -> int:
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
raise ValueError(f"{field} must be a non-negative integer")
return value
def require_positive_int(value: object, field: str) -> int:
result = require_nonnegative_int(value, field)
if result == 0:
raise ValueError(f"{field} must be a positive integer")
return result
def require_schema_version(value: object, expected: int, field: str) -> int:
if isinstance(value, bool) or not isinstance(value, int) or value != expected:
raise ValueError(f"{field} must be the integer {expected}")
return value
def require_task_key(value: object, field: str = "task_key") -> str:
text = require_string(value, field, max_length=256)
parts = text.split("/")
if any(
not part or part in {".", ".."} or not _TASK_KEY_COMPONENT_PATTERN.fullmatch(part)
for part in parts
):
raise ValueError(f"{field} must be a canonical workflow-relative key")
return text
def contains_unsafe_location(value: str) -> bool:
stripped = value.strip()
variants = [stripped]
for _ in range(2):
decoded = unquote(variants[-1])
if decoded == variants[-1]:
break
variants.append(decoded)
for candidate in variants:
if _SECRET_ASSIGNMENT_PATTERN.search(candidate):
return True
fragments = (candidate,) + tuple(
fragment
for fragment in re.split(r"[\s=\"'()\[\]{}<>;,]+", candidate)
if fragment
)
for fragment in fragments:
lower = fragment.lower()
normalized = fragment.replace("\\", "/")
segments = normalized.split("/")
looks_relative = (
len(segments) >= 3
and all(_PATH_SEGMENT_PATTERN.fullmatch(segment) for segment in segments)
) or (
len(segments) >= 2
and all(_PATH_SEGMENT_PATTERN.fullmatch(segment) for segment in segments)
and bool(_FILE_NAME_PATTERN.fullmatch(segments[-1]))
)
if (
bool(_URI_SCHEME_PATTERN.match(fragment))
or lower.startswith(tuple(prefix.lower() for prefix in _FORBIDDEN_LOCATOR_PREFIXES))
or fragment.startswith(("./", "../", "~/", "\\\\"))
or bool(_WINDOWS_PATH_PATTERN.match(fragment))
or any(segment == ".." for segment in segments)
or looks_relative
or (
_PATH_ASSIGNMENT_PATTERN.search(candidate) is not None
and ("/" in fragment or "\\" in fragment)
)
):
return True
return False
def require_safe_message(value: object, field: str, *, max_length: int = 512) -> str:
text = require_string(value, field, max_length=max_length)
tokens = (text,) + tuple(text.split())
if any(contains_unsafe_location(token.strip("'\"()[]{}<>,;")) for token in tokens):
raise ValueError(f"{field} must not contain a URI or local path")
return text
def require_opaque_resource_id(value: object, field: str) -> str:
"""Validate a non-secret resource handle without treating it as a locator."""
text = require_string(value, field, max_length=160)
if (
contains_unsafe_location(text)
or "/" in text
or "\\" in text
or "," in text
or any(character.isspace() for character in text)
):
raise ValueError(f"{field} must be an opaque single resource identifier")
return text
def require_finite_number(value: object, field: str) -> int | float:
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ValueError(f"{field} must be a finite number")
if isinstance(value, int):
if abs(value).bit_length() > 4096:
raise ValueError(f"{field} exceeds the 4096-bit integer bound")
return value
if not math.isfinite(value):
raise ValueError(f"{field} must be a finite number")
return value
def freeze_json(
value: object,
field: str,
*,
forbid_locations: bool = False,
_depth: int = 0,
) -> Any:
"""Return an immutable deep copy of a JSON value.
Scientific task parameters use ``forbid_locations`` so durable payloads
cannot smuggle worker-local paths or transport URLs. Manifests and verifier
evidence use ordinary JSON validation because JSON Schema keywords and
sanitized references may legitimately contain URI-shaped strings.
"""
if _depth > 64:
raise ValueError(f"{field} nesting exceeds 64 levels")
if value is None or isinstance(value, bool):
return value
if isinstance(value, int):
if abs(value).bit_length() > 4096:
raise ValueError(f"{field} contains an integer above the 4096-bit JSON bound")
return value
if isinstance(value, str):
if any(ord(character) < 32 for character in value):
raise ValueError(f"{field} must not contain control characters")
if forbid_locations and contains_unsafe_location(value):
raise ValueError(f"{field} must not contain a URI or local path")
return value
if isinstance(value, float):
if not math.isfinite(value):
raise ValueError(f"{field} must not contain NaN or infinity")
return value
if isinstance(value, Mapping):
frozen: dict[str, Any] = {}
for key, child in value.items():
if not isinstance(key, str):
raise ValueError(f"{field} must use string object keys")
frozen[key] = freeze_json(
child,
f"{field}.{key}",
forbid_locations=forbid_locations,
_depth=_depth + 1,
)
return MappingProxyType(frozen)
if isinstance(value, (list, tuple)):
return tuple(
freeze_json(
child,
f"{field}[]",
forbid_locations=forbid_locations,
_depth=_depth + 1,
)
for child in value
)
raise ValueError(f"{field} must contain only JSON-compatible values")
def freeze_json_mapping(
value: object,
field: str,
*,
forbid_locations: bool = False,
) -> Mapping[str, Any]:
mapping = require_mapping(value, field)
frozen = freeze_json(mapping, field, forbid_locations=forbid_locations)
assert isinstance(frozen, Mapping)
return frozen
def thaw_json(value: object) -> Any:
if isinstance(value, Mapping):
return {key: thaw_json(child) for key, child in value.items()}
if isinstance(value, tuple):
return [thaw_json(child) for child in value]
return value
def canonical_json(value: object) -> str:
return json.dumps(thaw_json(value), sort_keys=True, separators=(",", ":"), allow_nan=False)
def parse_release(value: object, field: str = "version") -> tuple[int, int, int]:
text = require_string(value, field, max_length=32)
match = _VERSION_PATTERN.fullmatch(text)
if match is None:
raise ValueError(f"{field} must contain one to three numeric release components")
return tuple(int(part or 0) for part in match.groups()) # type: ignore[return-value]
def validate_version_range(expression: object, field: str) -> str:
text = require_string(expression, field, max_length=128)
clauses = [clause.strip() for clause in text.split(",")]
if not clauses or any(not clause for clause in clauses):
raise ValueError(f"{field} must be an explicit version range")
canonical_clauses: list[str] = []
for clause in clauses:
match = _VERSION_CLAUSE_PATTERN.fullmatch(clause)
if match is None:
raise ValueError(f"{field} must use ==, >=, <=, >, or < clauses")
bound = match.group(2).strip()
parse_release(bound, field)
canonical_clauses.append(match.group(1) + bound)
return ",".join(canonical_clauses)
def version_in_range(version: object, expression: str) -> bool:
candidate = parse_release(version)
for clause in expression.split(","):
match = _VERSION_CLAUSE_PATTERN.fullmatch(clause)
assert match is not None
operator, raw_bound = match.groups()
bound = parse_release(raw_bound)
if operator == "==" and candidate != bound:
return False
if operator == ">=" and candidate < bound:
return False
if operator == "<=" and candidate > bound:
return False
if operator == ">" and candidate <= bound:
return False
if operator == "<" and candidate >= bound:
return False
return True
def enum_value(enum_type: type[Any], value: object, field: str) -> Any:
try:
return enum_type(value)
except (TypeError, ValueError) as error:
allowed = ", ".join(member.value for member in enum_type)
raise ValueError(f"{field} must be one of: {allowed}") from error
+705
View File
@@ -0,0 +1,705 @@
"""Typed artifact ports, immutable collections, and output provenance."""
from __future__ import annotations
import hashlib
import json
from dataclasses import dataclass, field
from datetime import datetime, timezone
from enum import Enum
from types import MappingProxyType
from typing import Any, Mapping, Sequence
from ._validation import (
canonical_json,
enum_value,
freeze_json_mapping,
require_exact_keys,
require_finite_number,
require_identifier,
require_nonnegative_int,
require_opaque_resource_id,
require_positive_int,
require_sha256,
require_schema_version,
require_string,
require_task_key,
require_uuid,
thaw_json,
parse_release,
)
from .identity import ComponentRef, OUTPUT_SCHEMA_VERSION, SchemaRef, WorkloadId
class CollectionKind(str, Enum):
SINGLE = "single"
ORDERED = "ordered"
KEYED = "keyed"
SET = "set"
class Cardinality(str, Enum):
ONE = "one"
OPTIONAL = "optional"
MANY = "many"
@dataclass(frozen=True, slots=True)
class ArtifactSchema:
"""Logical artifact shape and hard parsing bounds."""
ref: SchemaRef
media_type: str
encoding: str | None
max_bytes: int
validator: ComponentRef
validator_configuration: Mapping[str, Any] = field(default_factory=dict)
max_records: int | None = None
max_dimensions: tuple[int, ...] = ()
streaming: bool = False
canonicalizer: str | None = None
privacy_class: str = "project"
retention_class: str = "durable"
allow_nested_collections: bool = False
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):
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"))
if not isinstance(self.validator, ComponentRef):
raise ValueError("artifact schema validator must be a ComponentRef")
object.__setattr__(
self,
"validator_configuration",
freeze_json_mapping(
self.validator_configuration,
"artifact validator_configuration",
forbid_locations=True,
),
)
if self.max_records is not None:
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
for value in dimensions
):
raise ValueError("max_dimensions must contain positive integers")
if len(dimensions) > 8:
raise ValueError("max_dimensions must contain at most 8 axes")
object.__setattr__(self, "max_dimensions", dimensions)
if self.canonicalizer is not None:
object.__setattr__(
self,
"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):
raise ValueError("streaming and allow_nested_collections must be booleans")
def to_dict(self) -> dict[str, object]:
return {
"ref": self.ref.canonical,
"media_type": self.media_type,
"encoding": self.encoding,
"max_bytes": self.max_bytes,
"validator": self.validator.canonical,
"validator_configuration": thaw_json(self.validator_configuration),
"max_records": self.max_records,
"max_dimensions": list(self.max_dimensions),
"streaming": self.streaming,
"canonicalizer": self.canonicalizer,
"privacy_class": self.privacy_class,
"retention_class": self.retention_class,
"allow_nested_collections": self.allow_nested_collections,
}
@classmethod
def from_dict(cls, value: object) -> "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",
}
require_exact_keys(value, fields, "artifact schema")
dimensions = value["max_dimensions"]
if not isinstance(dimensions, list):
raise ValueError("max_dimensions must be an array")
return cls(
ref=SchemaRef.from_dict(value["ref"]),
media_type=value["media_type"], # type: ignore[arg-type]
encoding=value["encoding"], # type: ignore[arg-type]
max_bytes=value["max_bytes"], # type: ignore[arg-type]
validator=ComponentRef.from_dict(value["validator"]),
validator_configuration=value["validator_configuration"], # type: ignore[arg-type]
max_records=value["max_records"], # type: ignore[arg-type]
max_dimensions=tuple(dimensions),
streaming=value["streaming"], # type: ignore[arg-type]
canonicalizer=value["canonicalizer"], # type: ignore[arg-type]
privacy_class=value["privacy_class"], # type: ignore[arg-type]
retention_class=value["retention_class"], # type: ignore[arg-type]
allow_nested_collections=value["allow_nested_collections"], # type: ignore[arg-type]
)
@dataclass(frozen=True, slots=True)
class PortSpec:
schema: ArtifactSchema
cardinality: Cardinality = Cardinality.ONE
collection: CollectionKind = CollectionKind.SINGLE
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:
raise ValueError("one and optional cardinality require a single collection")
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)
if self.cardinality is Cardinality.ONE and count != 1:
raise ValueError(f"{field} must contain exactly one artifact")
if self.cardinality is Cardinality.OPTIONAL and count > 1:
raise ValueError(f"{field} must contain at most one artifact")
if self.cardinality is Cardinality.MANY and count < 1:
raise ValueError(f"{field} must contain at least one artifact")
for item in value.items:
artifact = item.artifact
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")
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:
if artifact.records is None:
raise ValueError(f"{field} is missing its required record summary")
if artifact.records > self.schema.max_records:
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")
if len(artifact.dimensions) != len(self.schema.max_dimensions) or any(
actual > maximum
for actual, maximum in zip(artifact.dimensions, self.schema.max_dimensions)
):
raise ValueError(f"{field} exceeds its dimension limits")
def to_dict(self) -> dict[str, object]:
return {
"schema": self.schema.to_dict(),
"cardinality": self.cardinality.value,
"collection": self.collection.value,
}
@classmethod
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")
return cls(
schema=ArtifactSchema.from_dict(value["schema"]),
cardinality=value["cardinality"], # type: ignore[arg-type]
collection=value["collection"], # type: ignore[arg-type]
)
@dataclass(frozen=True, slots=True)
class ArtifactRef:
"""Coordinator-owned artifact identity; transport URIs are intentionally absent."""
artifact_id: str
sha256: str
schema: SchemaRef
media_type: str
size_bytes: int
records: int | None = None
dimensions: tuple[int, ...] = ()
def __post_init__(self) -> None:
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):
raise ValueError("media_type must be a valid type/subtype token")
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"))
dimensions = tuple(self.dimensions)
if any(
isinstance(value, bool) or not isinstance(value, int) or value < 0
for value in dimensions
):
raise ValueError("dimensions must contain non-negative integers")
if len(dimensions) > 8:
raise ValueError("dimensions must contain at most 8 axes")
object.__setattr__(self, "dimensions", dimensions)
def to_dict(self) -> dict[str, object]:
return {
"artifact_id": self.artifact_id,
"sha256": self.sha256,
"schema": self.schema.canonical,
"media_type": self.media_type,
"size_bytes": self.size_bytes,
"records": self.records,
"dimensions": list(self.dimensions),
}
@classmethod
def from_dict(cls, value: object) -> "ArtifactRef":
if not isinstance(value, Mapping):
raise ValueError("artifact reference must be an object")
require_exact_keys(
value,
{"artifact_id", "sha256", "schema", "media_type", "size_bytes", "records", "dimensions"},
"artifact reference",
)
dimensions = value["dimensions"]
if not isinstance(dimensions, list):
raise ValueError("artifact dimensions must be an array")
return cls(
artifact_id=value["artifact_id"], # type: ignore[arg-type]
sha256=value["sha256"], # type: ignore[arg-type]
schema=SchemaRef.from_dict(value["schema"]),
media_type=value["media_type"], # type: ignore[arg-type]
size_bytes=value["size_bytes"], # type: ignore[arg-type]
records=value["records"], # type: ignore[arg-type]
dimensions=tuple(dimensions),
)
@dataclass(frozen=True, slots=True)
class ArtifactItem:
artifact: ArtifactRef
key: str | None = None
def __post_init__(self) -> None:
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"))
def to_dict(self) -> dict[str, object]:
return {"key": self.key, "artifact": self.artifact.to_dict()}
@classmethod
def from_dict(cls, value: object) -> "ArtifactItem":
if not isinstance(value, Mapping):
raise ValueError("artifact item must be an object")
require_exact_keys(value, {"key", "artifact"}, "artifact item")
return cls(artifact=ArtifactRef.from_dict(value["artifact"]), key=value["key"]) # type: ignore[arg-type]
@dataclass(frozen=True, slots=True)
class ArtifactCollection:
kind: CollectionKind
items: tuple[ArtifactItem, ...]
def __post_init__(self) -> None:
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")
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")
keys = [item.key for item in items]
if len(keys) != len(set(keys)):
raise ValueError("keyed collection keys must be unique")
items = tuple(sorted(items, key=lambda item: item.key or ""))
else:
if any(item.key is not None for item in items):
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)
for item in items
]
if len(identities) != len(set(identities)):
raise ValueError("set collection must not contain duplicate artifacts")
items = tuple(
sorted(
items,
key=lambda item: (
item.artifact.schema.canonical,
item.artifact.sha256,
item.artifact.size_bytes,
),
)
)
object.__setattr__(self, "items", items)
@classmethod
def single(cls, artifact: ArtifactRef | None) -> "ArtifactCollection":
return cls(CollectionKind.SINGLE, () if artifact is None else (ArtifactItem(artifact),))
@property
def size_bytes(self) -> int:
return sum(item.artifact.size_bytes for item in self.items)
@property
def digest(self) -> str:
payload = {
"kind": self.kind.value,
"items": [
{
"key": item.key,
"sha256": item.artifact.sha256,
"schema": item.artifact.schema.canonical,
"media_type": item.artifact.media_type,
"size_bytes": item.artifact.size_bytes,
}
for item in self.items
],
}
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]}
@classmethod
def from_dict(cls, value: object) -> "ArtifactCollection":
if not isinstance(value, Mapping):
raise ValueError("artifact collection must be an object")
require_exact_keys(value, {"kind", "items"}, "artifact collection")
items = value["items"]
if not isinstance(items, list):
raise ValueError("artifact collection items must be an array")
return cls(
kind=value["kind"], # type: ignore[arg-type]
items=tuple(ArtifactItem.from_dict(item) for item in items),
)
def _timestamp(value: object, field: str) -> str:
text = require_string(value, field, max_length=64)
try:
parsed = datetime.fromisoformat(text.replace("Z", "+00:00"))
except ValueError as error:
raise ValueError(f"{field} must be an RFC 3339 timestamp") from error
if parsed.tzinfo is None:
raise ValueError(f"{field} must include a timezone")
return parsed.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
@dataclass(frozen=True, slots=True)
class Provenance:
workload: WorkloadId
sdk_api_version: str
protocol_version: str
manifest_schema_version: int
workflow_schema_version: int
verifier: ComponentRef
artifact_schemas: tuple[SchemaRef, ...]
package_digest: str
manifest_digest: str
environment_digest: str
worker_runtime: Mapping[str, Any]
allocated_resource_ids: tuple[str, ...]
parameters_digest: str
input_collection_digest: str
execution_contract_digest: str
selected_features: Mapping[str, str]
optional_fallbacks: Mapping[str, str]
job_id: str
task_id: str
started_at: str
finished_at: str
trust_mode: str = "trusted"
random_seed: int | None = None
checkpoint_lineage: tuple[str, ...] = ()
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"))
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"),
)
object.__setattr__(
self,
"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")
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)
if len(canonical_json(runtime).encode("utf-8")) > 65_536:
raise ValueError("worker_runtime exceeds 64 KiB")
object.__setattr__(self, "worker_runtime", runtime)
resource_ids = tuple(
require_opaque_resource_id(value, "allocated_resource_id")
for value in self.allocated_resource_ids
)
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,
"execution_contract_digest",
require_sha256(self.execution_contract_digest, "execution_contract_digest"),
)
selected_features = freeze_json_mapping(
self.selected_features,
"provenance.selected_features",
)
optional_fallbacks = freeze_json_mapping(
self.optional_fallbacks,
"provenance.optional_fallbacks",
)
for name, version in selected_features.items():
require_identifier(name, "provenance selected feature")
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")
require_identifier(fallback, "provenance fallback")
if set(selected_features).intersection(optional_fallbacks):
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")
):
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)):
raise ValueError("random_seed must be an integer")
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)
def to_dict(self) -> dict[str, object]:
return {
"workload": self.workload.to_dict(),
"sdk_api_version": self.sdk_api_version,
"protocol_version": self.protocol_version,
"manifest_schema_version": self.manifest_schema_version,
"workflow_schema_version": self.workflow_schema_version,
"verifier": self.verifier.canonical,
"artifact_schemas": [schema.canonical for schema in self.artifact_schemas],
"package_digest": self.package_digest,
"manifest_digest": self.manifest_digest,
"environment_digest": self.environment_digest,
"worker_runtime": thaw_json(self.worker_runtime),
"allocated_resource_ids": list(self.allocated_resource_ids),
"parameters_digest": self.parameters_digest,
"input_collection_digest": self.input_collection_digest,
"execution_contract_digest": self.execution_contract_digest,
"selected_features": thaw_json(self.selected_features),
"optional_fallbacks": thaw_json(self.optional_fallbacks),
"job_id": self.job_id,
"task_id": self.task_id,
"started_at": self.started_at,
"finished_at": self.finished_at,
"trust_mode": self.trust_mode,
"random_seed": self.random_seed,
"checkpoint_lineage": list(self.checkpoint_lineage),
}
@classmethod
def from_dict(cls, value: object) -> "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",
}
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")
return cls(
workload=WorkloadId.from_dict(value["workload"]),
sdk_api_version=value["sdk_api_version"], # type: ignore[arg-type]
protocol_version=value["protocol_version"], # type: ignore[arg-type]
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),
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]
worker_runtime=value["worker_runtime"], # type: ignore[arg-type]
allocated_resource_ids=tuple(resource_ids),
parameters_digest=value["parameters_digest"], # type: ignore[arg-type]
input_collection_digest=value["input_collection_digest"], # type: ignore[arg-type]
execution_contract_digest=value["execution_contract_digest"], # type: ignore[arg-type]
selected_features=value["selected_features"], # type: ignore[arg-type]
optional_fallbacks=value["optional_fallbacks"], # type: ignore[arg-type]
job_id=value["job_id"], # type: ignore[arg-type]
task_id=value["task_id"], # type: ignore[arg-type]
started_at=value["started_at"], # type: ignore[arg-type]
finished_at=value["finished_at"], # type: ignore[arg-type]
trust_mode=value["trust_mode"], # type: ignore[arg-type]
random_seed=value["random_seed"], # type: ignore[arg-type]
checkpoint_lineage=tuple(lineage),
)
@dataclass(frozen=True, slots=True)
class OutputManifest:
task_key: str
outputs: Mapping[str, ArtifactCollection]
metrics: Mapping[str, int | float]
provenance: Provenance
schema_version: int = OUTPUT_SCHEMA_VERSION
def __post_init__(self) -> None:
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")
outputs: dict[str, ArtifactCollection] = {}
for name, collection in self.outputs.items():
canonical = require_identifier(name, "output port")
if not isinstance(collection, ArtifactCollection):
raise ValueError("output values must be ArtifactCollection values")
outputs[canonical] = collection
object.__setattr__(self, "outputs", MappingProxyType(outputs))
if not isinstance(self.metrics, Mapping):
raise ValueError("metrics must be an object")
metrics: dict[str, int | float] = {}
for name, value in self.metrics.items():
canonical = require_identifier(name, "metric name")
metrics[canonical] = require_finite_number(value, "metric value")
if len(canonical_json(metrics).encode("utf-8")) > 16_384:
raise ValueError("output metrics exceed 16 KiB")
object.__setattr__(self, "metrics", MappingProxyType(metrics))
if not isinstance(self.provenance, Provenance):
raise ValueError("provenance must be a Provenance value")
def validate_against(
self,
expected: Mapping[str, PortSpec],
*,
max_output_bytes: int,
) -> "OutputManifest":
if set(self.outputs) != set(expected):
missing = sorted(set(expected) - set(self.outputs))
unexpected = sorted(set(self.outputs) - set(expected))
details = []
if missing:
details.append("missing " + ", ".join(missing))
if unexpected:
details.append("unexpected " + ", ".join(unexpected))
raise ValueError("output ports do not match the declaration: " + "; ".join(details))
total = 0
for name, port in expected.items():
if not isinstance(port, PortSpec):
raise ValueError("expected outputs must contain PortSpec values")
port.validate_collection(self.outputs[name], f"output {name}")
total += self.outputs[name].size_bytes
if total > require_positive_int(max_output_bytes, "max_output_bytes"):
raise ValueError("output manifest exceeds the total byte limit")
return self
@property
def digest(self) -> str:
payload = {
"outputs": {
name: {"kind": collection.kind.value, "digest": collection.digest}
for name, collection in sorted(self.outputs.items())
}
}
return hashlib.sha256(canonical_json(payload).encode("utf-8")).hexdigest()
@property
def manifest_digest(self) -> str:
"""Digest the complete audit manifest, including provenance and metrics."""
return hashlib.sha256(self.to_json().encode("utf-8")).hexdigest()
def to_dict(self) -> dict[str, object]:
return {
"schema_version": self.schema_version,
"task_key": self.task_key,
"outputs": {name: value.to_dict() for name, value in self.outputs.items()},
"metrics": dict(self.metrics),
"provenance": self.provenance.to_dict(),
}
def to_json(self) -> str:
return canonical_json(self.to_dict())
@classmethod
def from_dict(cls, value: object) -> "OutputManifest":
if not isinstance(value, Mapping):
raise ValueError("output manifest must be an object")
require_exact_keys(
value,
{"schema_version", "task_key", "outputs", "metrics", "provenance"},
"output manifest",
)
outputs = value["outputs"]
if not isinstance(outputs, Mapping):
raise ValueError("outputs must be an object")
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()},
metrics=value["metrics"], # type: ignore[arg-type]
provenance=Provenance.from_dict(value["provenance"]),
)
@classmethod
def from_json(cls, value: str) -> "OutputManifest":
try:
decoded = json.loads(value)
except (TypeError, json.JSONDecodeError, RecursionError) as error:
raise ValueError("output manifest must be valid JSON") from error
return cls.from_dict(decoded)
+146
View File
@@ -0,0 +1,146 @@
"""SDK definitions for existing SciMesh workloads and local core runtime."""
from __future__ import annotations
import hashlib
import os
import platform
import sys
from rdkit import rdBase
from scimesh.distributed.similarity_search import (
SimilaritySearchDistributedWorkload,
run_similarity_search_shard,
)
from .artifacts import ArtifactSchema, PortSpec
from .compat import LegacyDistributedWorkloadAdapter
from .identity import ComponentRef, SDK_API_VERSION, SchemaRef
from .integrity import installed_distribution_digest
from .registry import WorkloadRegistry
from .resources import ResourceInventory
from .runtime import RuntimeCapabilities
def current_scimesh_package_digest() -> str:
"""Hash installed SciMesh Python sources for the built-in trusted adapter.
This is a local immutable-code pin, not a package signature or container
attestation. Consequently the built-in compatibility manifest is trusted
only; an administrator must supply signed image metadata before enabling an
untrusted quorum policy.
"""
# Source/editable installs are allowed only for this explicit local
# development helper. Registry discovery keeps the secure default.
return installed_distribution_digest("scimesh", allow_editable=True)
def current_environment_digest() -> str:
payload = "\n".join(
(
current_scimesh_package_digest(),
f"python={sys.implementation.name}-{platform.python_version()}",
f"rdkit={rdBase.rdkitVersion}",
f"platform={sys.platform}-{platform.machine().lower()}",
)
)
return "sha256:" + hashlib.sha256(payload.encode("utf-8")).hexdigest()
def similarity_search_sdk_adapter(*, shard_rows: int = 10_000) -> LegacyDistributedWorkloadAdapter:
dataset_schema = ArtifactSchema(
SchemaRef("molecule-table", 1),
"text/tab-separated-values",
"utf-8",
max_bytes=10 * 1024 * 1024 * 1024,
validator=ComponentRef("delimited-table", 1),
validator_configuration={
"required_columns": ["canonical_smiles", "chembl_id"],
},
max_records=100_000_000,
canonicalizer="scimesh-tsv-v1",
)
partial_schema = ArtifactSchema(
SchemaRef("similarity-search-partial", 1),
"text/csv",
"utf-8",
max_bytes=1024 * 1024 * 1024,
validator=ComponentRef("delimited-table", 1),
validator_configuration={
"columns": ["rank", "chembl_id", "canonical_smiles", "similarity"],
},
max_records=100_000,
canonicalizer="scimesh-search-partial-v1",
)
result_schema = ArtifactSchema(
SchemaRef("similarity-search-result", 1),
"text/csv",
"utf-8",
max_bytes=1024 * 1024 * 1024,
validator=ComponentRef("delimited-table", 1),
validator_configuration={
"columns": ["rank", "chembl_id", "canonical_smiles", "similarity"],
},
max_records=100_000,
canonicalizer="scimesh-search-result-v1",
)
parameters_schema = {
"type": "object",
"additionalProperties": False,
"properties": {
"query_id": {"type": "string", "minLength": 1, "maxLength": 200},
"query_smiles": {"type": "string", "minLength": 1, "maxLength": 200},
"top_k": {"type": "integer", "minimum": 1},
"threshold": {"type": "number", "minimum": 0, "maximum": 1},
"threshold_direction": {"enum": ["greater", "less"]},
"max_rows": {"type": "integer", "minimum": 1},
"progress_every": {"type": "integer", "minimum": 0},
},
"oneOf": [
{"required": ["query_id"], "not": {"required": ["query_smiles"]}},
{"required": ["query_smiles"], "not": {"required": ["query_id"]}},
],
}
return LegacyDistributedWorkloadAdapter(
SimilaritySearchDistributedWorkload(),
run_similarity_search_shard,
version="1.0.0",
package_digest=current_scimesh_package_digest(),
environment_digest=current_environment_digest(),
parameters_schema=parameters_schema,
input_port=PortSpec(dataset_schema),
partial_port=PortSpec(partial_schema),
output_port=PortSpec(result_schema),
resolved_parameter_names=("query_source", "fingerprint"),
shard_rows=shard_rows,
)
def default_sdk_registry(*, shard_rows: int = 10_000) -> WorkloadRegistry:
registry = WorkloadRegistry()
registry.register(similarity_search_sdk_adapter(shard_rows=shard_rows).definition(), enabled=True)
return registry
def similarity_search_workload_definition():
"""Installed entry-point factory for the default shard-size definition."""
return similarity_search_sdk_adapter().definition()
def default_sdk_runtime() -> RuntimeCapabilities:
architecture = platform.machine().lower() or "unknown"
return RuntimeCapabilities(
sdk_api_version=SDK_API_VERSION,
protocol_version="1.0.0",
profiles=("core-batch-v1",),
features={"artifact-collections": "1.0.0", "exact-verifier": "1.0.0"},
workload_capabilities=("similarity-search",),
inventory=ResourceInventory(
cpu_cores=max(os.cpu_count() or 1, 1),
memory_mb=4096,
scratch_mb=4096,
architecture=architecture,
environment_digests=(current_environment_digest(),),
),
)
+5
View File
@@ -0,0 +1,5 @@
"""Adapters for versioned pre-SDK SciMesh workload contracts."""
from .distributed_v1 import LegacyDistributedWorkloadAdapter
__all__ = ["LegacyDistributedWorkloadAdapter"]
+369
View File
@@ -0,0 +1,369 @@
"""Compatibility adapter for the CTX-07 ``DistributedWorkload`` protocol."""
from __future__ import annotations
import hashlib
import shutil
from pathlib import Path
from typing import Any, Callable, Mapping, Sequence
from scimesh.distributed.models import (
ArtifactReference as LegacyArtifactReference,
CompletedPartial,
FinalResult,
)
from scimesh.distributed.workload import DistributedWorkload
from ..artifacts import (
ArtifactCollection,
ArtifactItem,
ArtifactRef,
Cardinality,
CollectionKind,
OutputManifest,
PortSpec,
)
from ..execution import CheckpointPolicy, ExecutionProfile, NetworkPolicy, RetryPolicy
from ..identity import ComponentRef, SchemaRef, VersionRange, WorkloadId
from ..manifest import (
DeterminismProfile,
EnvironmentSpec,
PackageSpec,
TrustMode,
VerifierSpec,
WorkloadLimits,
WorkloadManifest,
)
from ..plans import JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
from ..protocols import PlanningContext, ReduceContext, TaskContext
from ..registry import WorkloadDefinition
from ..resources import ResourceRequirements
from ..verification import ExactArtifactVerifier
from ..workflow import ArtifactEdge, PortRef, StageKind, StageSpec, WorkflowSpec
ShardRunner = Callable[[Path, Mapping[str, object], Path], Mapping[str, int | float]]
def _sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as source:
for block in iter(lambda: source.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
class LegacyDistributedWorkloadAdapter:
"""Expose a legacy map/reduce workload through the SDK core-batch profile.
The adapter preserves the old wire schema. Local files are materialized and
sealed only through bridge-owned contexts, and no path is included in a
``TaskSpec`` or ``WorkflowPlan``.
"""
MAP_ENTRY_POINT = "scimesh.sdk.compat.distributed_v1:run_legacy@v1"
REDUCE_ENTRY_POINT = "scimesh.sdk.compat.distributed_v1:reduce_legacy@v1"
def __init__(
self,
workload: DistributedWorkload,
shard_runner: ShardRunner,
*,
version: str,
package_digest: str,
environment_digest: str,
parameters_schema: Mapping[str, Any],
input_port: PortSpec,
partial_port: PortSpec,
output_port: PortSpec,
resolved_parameter_names: Sequence[str] = (),
shard_rows: int = 10_000,
resources: ResourceRequirements | None = None,
execution: ExecutionProfile | None = None,
limits: WorkloadLimits | None = None,
) -> None:
if not isinstance(workload.name, str) or not isinstance(workload.description, str):
raise ValueError("legacy workload must expose name and description")
if not callable(shard_runner):
raise ValueError("shard_runner must be callable")
if isinstance(shard_rows, bool) or not isinstance(shard_rows, int) or shard_rows < 1:
raise ValueError("shard_rows must be a positive integer")
self.workload = workload
self.shard_runner = shard_runner
self.shard_rows = shard_rows
self.input_port = input_port
self.partial_port = partial_port
self.output_port = output_port
resources = resources or ResourceRequirements(
profile="legacy-cpu-v1",
cpu_cores=1,
memory_mb=1024,
scratch_mb=1024,
max_duration_seconds=3600,
)
execution = execution or ExecutionProfile(
profile="legacy-python-process-v1",
network=NetworkPolicy.TRUSTED,
timeout_seconds=3600,
checkpoint=CheckpointPolicy(),
)
limits = limits or WorkloadLimits(
max_input_bytes=input_port.schema.max_bytes,
max_tasks=10_000,
max_output_bytes=output_port.schema.max_bytes,
)
parameter_names = tuple(sorted(parameters_schema.get("properties", {})))
reduce_parameter_names = tuple(sorted(set(parameter_names).union(resolved_parameter_names)))
map_stage = StageSpec(
stage_id="map",
kind=StageKind.MAP,
entry_point=self.MAP_ENTRY_POINT,
needs=(),
inputs={"input": input_port},
outputs={"partial": partial_port},
parameter_names=parameter_names,
resources=resources,
execution=execution,
retry=RetryPolicy(),
verifier=ComponentRef("exact-artifact", 1),
trust_modes=("trusted",),
max_fan_out=limits.max_tasks,
cacheable=True,
)
reduce_input = PortSpec(
schema=partial_port.schema,
cardinality=Cardinality.MANY,
collection=CollectionKind.KEYED,
)
reduce_stage = StageSpec(
stage_id="reduce",
kind=StageKind.REDUCE,
entry_point=self.REDUCE_ENTRY_POINT,
needs=("map",),
inputs={"partials": reduce_input},
outputs={"result": output_port},
parameter_names=reduce_parameter_names,
resources=resources,
execution=execution,
retry=RetryPolicy(),
verifier=ComponentRef("exact-artifact", 1),
trust_modes=("trusted",),
cacheable=True,
)
workflow = WorkflowSpec(
workflow_id="map-reduce-v1",
inputs={"input": input_port},
stages=(map_stage, reduce_stage),
edges=(
ArtifactEdge(PortRef("input"), PortRef("input", "map")),
ArtifactEdge(PortRef("partial", "map"), PortRef("partials", "reduce")),
),
outputs={"result": PortRef("result", "reduce")},
max_tasks=limits.max_tasks,
max_output_bytes=limits.max_output_bytes,
)
self.manifest = WorkloadManifest(
sdk_api=VersionRange(">=1.0,<2.0"),
protocol=VersionRange(">=1,<2"),
workload=WorkloadId(workload.name, version),
description=workload.description,
package=PackageSpec("scimesh", package_digest),
environment=EnvironmentSpec("python-process", environment_digest, {"adapter": "distributed-v1"}),
parameters_schema=parameters_schema,
workflow=workflow,
inputs={"input": input_port},
outputs={"result": output_port},
determinism=DeterminismProfile.BYTE_EXACT,
trust_modes=(TrustMode.TRUSTED,),
verifier=VerifierSpec(ComponentRef("exact-artifact", 1), {}),
limits=limits,
capabilities=(workload.name,),
conformance_profiles=("core-batch-v1",),
)
self._exact_verifier = ExactArtifactVerifier()
def definition(self) -> WorkloadDefinition:
return WorkloadDefinition(
manifest=self.manifest,
planner=self,
runners={self.MAP_ENTRY_POINT: self},
reducers={self.REDUCE_ENTRY_POINT: self},
verifiers={self._exact_verifier.identity.canonical: self._exact_verifier},
)
def validate(self, request: JobRequest) -> ValidatedJob:
if request.workload != self.manifest.workload:
raise ValueError("legacy adapter received a request for another workload")
self.workload.validate_job(request.parameters)
return ValidatedJob(request, request.parameters)
def plan(self, job: ValidatedJob, context: PlanningContext) -> WorkflowPlan:
if not isinstance(job, ValidatedJob):
raise ValueError("job must be a ValidatedJob")
collection = job.request.inputs.get("input")
if collection is None:
raise ValueError("legacy adapter requires the input port")
self.input_port.validate_collection(collection, "job input")
input_artifact = collection.items[0].artifact
input_path = context.catalog.materialize(input_artifact)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
legacy = self.workload.plan(
input_path,
input_artifact.artifact_id,
job.request.parameters,
self.shard_rows,
workspace,
)
if legacy.workload != self.workload.name:
raise ValueError("legacy planner returned a plan for another workload")
tasks: list[TaskSpec] = []
negotiated = context.negotiated
map_stage = self.manifest.workflow.stages[0]
used_paths: set[Path] = set()
for planned in legacy.tasks:
path = self._find_planned_file(workspace, planned.input_artifact.sha256, used_paths)
sealed = context.sink.seal(
path,
declaration=self.input_port.schema,
)
if sealed.sha256 != planned.input_artifact.sha256:
raise ValueError("artifact sink returned a checksum that differs from the legacy plan")
tasks.append(
TaskSpec(
workload=self.manifest.workload,
package_digest=self.manifest.package.digest,
manifest_digest=self.manifest.digest,
trust_mode=job.request.trust_mode,
sdk_api_version=negotiated.sdk_api_version,
protocol_version=negotiated.protocol_version,
manifest_schema_version=self.manifest.manifest_schema_version,
workflow_schema_version=self.manifest.workflow.schema_version,
environment_digest=self.manifest.environment.digest,
verifier=map_stage.verifier,
selected_features=negotiated.selected_features,
optional_fallbacks=negotiated.optional_fallbacks,
task_key=f"map/{planned.chunk_index:08d}",
stage_id="map",
parameters=planned.parameters,
inputs={"input": ArtifactCollection.single(sealed)},
expected_outputs={"partial": self.partial_port},
resources=map_stage.resources,
execution=map_stage.execution,
)
)
return WorkflowPlan(
workload=self.manifest.workload,
package_digest=self.manifest.package.digest,
manifest_digest=self.manifest.digest,
trust_mode=job.request.trust_mode,
sdk_api_version=negotiated.sdk_api_version,
protocol_version=negotiated.protocol_version,
manifest_schema_version=self.manifest.manifest_schema_version,
workflow_schema_version=self.manifest.workflow.schema_version,
environment_digest=self.manifest.environment.digest,
verifier=self.manifest.verifier.verifier,
selected_features=negotiated.selected_features,
optional_fallbacks=negotiated.optional_fallbacks,
workflow_id=self.manifest.workflow.workflow_id,
resolved_parameters=legacy.resolved_parameters,
tasks=tuple(tasks),
)
@staticmethod
def _find_planned_file(workspace: Path, expected_sha256: str, used: set[Path]) -> Path:
for candidate in sorted(workspace.rglob("*")):
if candidate in used or not candidate.is_file() or candidate.is_symlink():
continue
if _sha256_file(candidate) == expected_sha256:
used.add(candidate)
return candidate
raise ValueError("legacy planner did not materialize its planned artifact")
def run(self, context: TaskContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
collection = context.task.inputs.get("input")
if collection is None:
raise ValueError("legacy map task requires one input collection")
self.input_port.validate_collection(collection, "legacy map input")
source = context.catalog.materialize(collection.items[0].artifact)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
input_path = workspace / "input"
output_path = workspace / "result"
if source.resolve() != input_path.resolve():
shutil.copyfile(source, input_path)
metrics = self.shard_runner(input_path, context.task.parameters, output_path)
context.cancellation.raise_if_cancelled()
sealed = context.sink.seal(
output_path,
declaration=self.partial_port.schema,
)
return OutputManifest(
context.task.task_key,
{"partial": ArtifactCollection.single(sealed)},
metrics,
context.provenance,
).validate_against(context.task.expected_outputs, max_output_bytes=self.manifest.limits.max_output_bytes)
def reduce(self, context: ReduceContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
collection = context.accepted_inputs.get("partials")
if collection is None or collection.kind is not CollectionKind.KEYED or not collection.items:
raise ValueError("legacy reducer requires a non-empty keyed partial collection")
self.manifest.workflow.stages[1].inputs["partials"].validate_collection(
collection,
"legacy reducer partials",
)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
partials: list[CompletedPartial] = []
indexed_items: list[tuple[int, ArtifactItem]] = []
for item in collection.items:
key = item.key or ""
prefix = "map."
raw_index = key[len(prefix):] if key.startswith(prefix) else ""
if len(raw_index) != 8 or not raw_index.isdigit():
raise ValueError("legacy partial key must use map.<eight-digit-index>")
indexed_items.append((int(raw_index), item))
indices = [index for index, _ in indexed_items]
expected_keys = context.task.expected_input_keys.get("partials")
if expected_keys is None or {item.key for item in collection.items} != set(expected_keys):
raise ValueError("legacy partial keys do not match the coordinator expected set")
if sorted(indices) != list(range(len(indexed_items))):
raise ValueError("legacy partial keys must be complete and contiguous")
for index, item in sorted(indexed_items):
artifact = item.artifact
source = context.catalog.materialize(artifact)
target = workspace / artifact.artifact_id
if source.resolve() != target.resolve():
shutil.copyfile(source, target)
if _sha256_file(target) != artifact.sha256:
raise ValueError("materialized partial checksum does not match")
partials.append(
CompletedPartial(
index,
LegacyArtifactReference(
artifact.artifact_id,
artifact.sha256,
artifact.media_type,
),
{},
)
)
result = self.workload.reduce(partials, context.task.parameters, workspace)
if not isinstance(result, FinalResult):
raise ValueError("legacy reducer must return a FinalResult")
path = self._find_planned_file(workspace, result.artifact.sha256, set())
sealed = context.sink.seal(
path,
declaration=self.output_port.schema,
)
if sealed.sha256 != result.artifact.sha256:
raise ValueError("artifact sink returned a checksum that differs from the legacy result")
return OutputManifest(
context.task.task_key,
{"result": ArtifactCollection.single(sealed)},
result.metrics,
context.provenance,
).validate_against(context.task.expected_outputs, max_output_bytes=self.manifest.limits.max_output_bytes)
File diff suppressed because it is too large Load Diff
+345
View File
@@ -0,0 +1,345 @@
"""Execution, retry, checkpoint, cancellation, and failure declarations."""
from __future__ import annotations
from dataclasses import dataclass
from enum import Enum
from types import MappingProxyType
from typing import Any, Mapping
from ._validation import (
enum_value,
freeze_json_mapping,
require_exact_keys,
require_identifier,
require_nonnegative_int,
require_safe_message,
require_positive_int,
require_string,
thaw_json,
)
from .identity import SchemaRef
from .resources import ResourceAllocation, ResourceRequirements
class ProcessModel(str, Enum):
SINGLE = "single"
PROCESS_POOL = "process_pool"
THREAD_POOL = "thread_pool"
EXTERNAL_RUNTIME = "external_runtime"
class NetworkPolicy(str, Enum):
NONE = "none"
COORDINATOR_ARTIFACTS_ONLY = "coordinator_artifacts_only"
ALLOWLISTED_EGRESS = "allowlisted_egress"
TRUSTED = "trusted"
class FailureCategory(str, Enum):
INPUT = "input"
SCIENTIFIC = "scientific"
RESOURCE = "resource"
INFRASTRUCTURE = "infrastructure"
LEASE = "lease"
VERIFICATION = "verification"
POLICY = "policy"
@dataclass(frozen=True, slots=True)
class RetryPolicy:
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"))
categories = tuple(
enum_value(FailureCategory, value, "retryable_category")
for value in self.retryable_categories
)
if len(categories) != len(set(categories)):
raise ValueError("retryable_categories must be unique")
object.__setattr__(self, "retryable_categories", categories)
object.__setattr__(
self,
"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"),
)
if self.max_backoff_seconds < self.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],
"initial_backoff_seconds": self.initial_backoff_seconds,
"max_backoff_seconds": self.max_backoff_seconds,
}
@classmethod
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"}
require_exact_keys(value, fields, "retry policy")
categories = value["retryable_categories"]
if not isinstance(categories, list):
raise ValueError("retryable_categories must be an array")
return cls(
max_attempts=value["max_attempts"], # type: ignore[arg-type]
retryable_categories=tuple(categories),
initial_backoff_seconds=value["initial_backoff_seconds"], # type: ignore[arg-type]
max_backoff_seconds=value["max_backoff_seconds"], # type: ignore[arg-type]
)
@dataclass(frozen=True, slots=True)
class CheckpointPolicy:
enabled: bool = False
schema: SchemaRef | None = None
compatibility_version: int | None = None
interval_seconds: int | None = None
def __post_init__(self) -> None:
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")
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")
object.__setattr__(
self,
"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"),
)
def to_dict(self) -> dict[str, object]:
return {
"enabled": self.enabled,
"schema": self.schema.canonical if self.schema is not None else None,
"compatibility_version": self.compatibility_version,
"interval_seconds": self.interval_seconds,
}
@classmethod
def from_dict(cls, value: object) -> "CheckpointPolicy":
if not isinstance(value, Mapping):
raise ValueError("checkpoint policy must be an object")
fields = {"enabled", "schema", "compatibility_version", "interval_seconds"}
require_exact_keys(value, fields, "checkpoint policy")
raw_schema = value["schema"]
return cls(
enabled=value["enabled"], # type: ignore[arg-type]
schema=None if raw_schema is None else SchemaRef.from_dict(raw_schema),
compatibility_version=value["compatibility_version"], # type: ignore[arg-type]
interval_seconds=value["interval_seconds"], # type: ignore[arg-type]
)
@dataclass(frozen=True, slots=True)
class ExecutionProfile:
profile: str
process_model: ProcessModel = ProcessModel.SINGLE
max_processes: int = 1
threads_per_process: int = 1
native_threads: int = 1
nested_parallelism: bool = False
network: NetworkPolicy = NetworkPolicy.NONE
timeout_seconds: int = 3600
cancellation_grace_seconds: int = 10
checkpoint: CheckpointPolicy = CheckpointPolicy()
allowed_egress: tuple[str, ...] = ()
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,
"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"))
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,
"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)
if len(egress) != len(set(egress)):
raise ValueError("allowed_egress must be unique")
if self.network is NetworkPolicy.ALLOWLISTED_EGRESS and not egress:
raise ValueError("allowlisted egress policy requires at least one target")
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)
if len(handles) != len(set(handles)):
raise ValueError("secret_handles must be unique")
if handles and self.network is NetworkPolicy.NONE:
raise ValueError("secret handles require an explicit network policy")
object.__setattr__(self, "secret_handles", handles)
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("nested thread pools require nested_parallelism=true")
@property
def maximum_cpu_threads(self) -> int:
return self.max_processes * self.threads_per_process * self.native_threads
def validate_resources(self, resources: ResourceRequirements) -> None:
if self.maximum_cpu_threads > resources.cpu_cores:
raise ValueError("execution profile can oversubscribe its CPU reservation")
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]:
"""Return only allocation-derived thread/device isolation variables."""
if not isinstance(allocation, ResourceAllocation):
raise ValueError("allocation must be a ResourceAllocation")
native = str(min(self.native_threads, allocation.cpu_cores))
values = {
"OMP_NUM_THREADS": native,
"OPENBLAS_NUM_THREADS": native,
"MKL_NUM_THREADS": native,
"NUMEXPR_NUM_THREADS": native,
"VECLIB_MAXIMUM_THREADS": native,
# Empty visibility explicitly prevents a CPU task from inheriting
# access to all host devices.
"CUDA_VISIBLE_DEVICES": ",".join(allocation.accelerator_ids),
"ROCR_VISIBLE_DEVICES": ",".join(allocation.accelerator_ids),
}
return MappingProxyType(values)
def to_dict(self) -> dict[str, object]:
return {
"profile": self.profile,
"process_model": self.process_model.value,
"max_processes": self.max_processes,
"threads_per_process": self.threads_per_process,
"native_threads": self.native_threads,
"nested_parallelism": self.nested_parallelism,
"network": self.network.value,
"timeout_seconds": self.timeout_seconds,
"cancellation_grace_seconds": self.cancellation_grace_seconds,
"checkpoint": self.checkpoint.to_dict(),
"allowed_egress": list(self.allowed_egress),
"secret_handles": list(self.secret_handles),
}
@classmethod
def from_dict(cls, value: object) -> "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",
}
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")
return cls(
profile=value["profile"], # type: ignore[arg-type]
process_model=value["process_model"], # type: ignore[arg-type]
max_processes=value["max_processes"], # type: ignore[arg-type]
threads_per_process=value["threads_per_process"], # type: ignore[arg-type]
native_threads=value["native_threads"], # type: ignore[arg-type]
nested_parallelism=value["nested_parallelism"], # type: ignore[arg-type]
network=value["network"], # type: ignore[arg-type]
timeout_seconds=value["timeout_seconds"], # type: ignore[arg-type]
cancellation_grace_seconds=value["cancellation_grace_seconds"], # type: ignore[arg-type]
checkpoint=CheckpointPolicy.from_dict(value["checkpoint"]),
allowed_egress=tuple(allowed_egress),
secret_handles=tuple(secret_handles),
)
@dataclass(frozen=True, slots=True)
class FailureReport:
code: str
category: FailureCategory
retryable: bool
message: str
evidence: Mapping[str, Any]
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"))
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)
import json
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)
def to_dict(self) -> dict[str, object]:
return {
"code": self.code,
"category": self.category.value,
"retryable": self.retryable,
"message": self.message,
"evidence": thaw_json(self.evidence),
}
def to_json(self) -> str:
from ._validation import canonical_json
return canonical_json(self.to_dict())
@classmethod
def from_dict(cls, value: object) -> "FailureReport":
if not isinstance(value, Mapping):
raise ValueError("failure report must be an object")
fields = {"code", "category", "retryable", "message", "evidence"}
require_exact_keys(value, fields, "failure report")
return cls(
code=value["code"], # type: ignore[arg-type]
category=value["category"], # type: ignore[arg-type]
retryable=value["retryable"], # type: ignore[arg-type]
message=value["message"], # type: ignore[arg-type]
evidence=value["evidence"], # type: ignore[arg-type]
)
@classmethod
def from_json(cls, value: str) -> "FailureReport":
import json
try:
decoded = json.loads(value)
except (TypeError, json.JSONDecodeError, RecursionError) as error:
raise ValueError("failure report must be valid JSON") from error
return cls.from_dict(decoded)
+169
View File
@@ -0,0 +1,169 @@
"""Versioned identities used across the SciMesh workload SDK."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Mapping
from ._validation import (
require_exact_keys,
require_identifier,
require_semver,
require_string,
require_workload_name,
validate_version_range,
version_in_range,
)
SDK_API_VERSION = "1.0.0"
MANIFEST_SCHEMA_VERSION = 1
WORKFLOW_SCHEMA_VERSION = 1
TASK_SCHEMA_VERSION = 1
OUTPUT_SCHEMA_VERSION = 1
@dataclass(frozen=True, slots=True)
class VersionRange:
"""A deliberately small, explicit compatibility range.
The v1 SDK accepts comma-separated comparisons such as ``>=1.0,<2.0``.
Wildcards and an omitted operator are rejected so a missing version can
never be interpreted as "latest".
"""
expression: str
def __post_init__(self) -> None:
object.__setattr__(self, "expression", validate_version_range(self.expression, "version range"))
def contains(self, version: str) -> bool:
return version_in_range(version, self.expression)
def to_dict(self) -> str:
return self.expression
@classmethod
def from_dict(cls, value: object) -> "VersionRange":
return cls(value) # type: ignore[arg-type]
@dataclass(frozen=True, slots=True)
class WorkloadId:
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"))
def to_dict(self) -> dict[str, str]:
return {"name": self.name, "version": self.version}
@classmethod
def from_dict(cls, value: object) -> "WorkloadId":
if not isinstance(value, Mapping):
raise ValueError("workload identity must be an object")
require_exact_keys(value, {"name", "version"}, "workload identity")
return cls(name=value["name"], version=value["version"]) # type: ignore[arg-type]
@dataclass(frozen=True, slots=True)
class SchemaRef:
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:
raise ValueError("schema.version must be a positive integer")
@property
def canonical(self) -> str:
return f"{self.name}@{self.version}"
def to_dict(self) -> dict[str, object]:
return {"name": self.name, "version": self.version}
@classmethod
def parse(cls, value: object, field: str = "schema") -> "SchemaRef":
text = require_string(value, field, max_length=160)
name, separator, raw_version = text.rpartition("@")
if not separator or not raw_version.isdigit():
raise ValueError(f"{field} must use the name@version form")
return cls(name=name, version=int(raw_version))
@classmethod
def from_dict(cls, value: object) -> "SchemaRef":
if isinstance(value, str):
return cls.parse(value)
if not isinstance(value, Mapping):
raise ValueError("schema reference must be a name@version string or object")
require_exact_keys(value, {"name", "version"}, "schema reference")
return cls(name=value["name"], version=value["version"]) # type: ignore[arg-type]
@dataclass(frozen=True, slots=True)
class ComponentRef:
"""Versioned, package-owned planner/runner/reducer/verifier identity."""
name: str
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:
raise ValueError("component.version must be a positive integer")
@property
def canonical(self) -> str:
return f"{self.name}@{self.version}"
def to_dict(self) -> dict[str, object]:
return {"name": self.name, "version": self.version}
@classmethod
def from_dict(cls, value: object) -> "ComponentRef":
if isinstance(value, str):
parsed = SchemaRef.parse(value, "component")
return cls(parsed.name, parsed.version)
if not isinstance(value, Mapping):
raise ValueError("component reference must be an object")
require_exact_keys(value, {"name", "version"}, "component reference")
return cls(name=value["name"], version=value["version"]) # type: ignore[arg-type]
@dataclass(frozen=True, slots=True)
class FeatureRequirement:
name: str
versions: VersionRange
fallback: str | None = None
def __post_init__(self) -> None:
object.__setattr__(self, "name", require_identifier(self.name, "feature.name"))
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"))
def to_dict(self) -> dict[str, object]:
result: dict[str, object] = {"name": self.name, "versions": self.versions.expression}
if self.fallback is not None:
result["fallback"] = self.fallback
return result
@classmethod
def from_dict(cls, value: object) -> "FeatureRequirement":
if not isinstance(value, Mapping):
raise ValueError("feature requirement must be an object")
require_exact_keys(
value,
{"name", "versions"},
"feature requirement",
optional={"fallback"},
)
return cls(
name=value["name"], # type: ignore[arg-type]
versions=VersionRange.from_dict(value["versions"]),
fallback=value.get("fallback"), # type: ignore[arg-type]
)
+140
View File
@@ -0,0 +1,140 @@
"""Independent installed-distribution content measurement for SDK allowlists."""
from __future__ import annotations
import hashlib
import importlib.util
from importlib import metadata
from pathlib import Path
def installed_distribution_digest(
distribution: metadata.Distribution | str,
*,
allow_editable: bool = False,
) -> str:
"""Hash installed package payload files using a stable path/length framing.
Distribution metadata is deliberately excluded: editable/non-editable
installers generate different RECORD and entry-point files for identical
package code. All source/native modules and package data below declared
top-level packages are included. Interpreter-generated ``__pycache__``
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
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()]
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)
if not top_levels:
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"
for item in declared_files
)
if editable_bootstrap and not allow_editable:
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")
if suffix in {".pyc", ".pyo"} and "__pycache__" not in relative.parts:
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))
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 ()
if len(locations) > 1:
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")
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")
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")
relative = f"{top_level}/{path.relative_to(root).as_posix()}"
selected.append((relative, path))
continue
module = Path(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")
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
# sourceless bytecode payloads were rejected above. Generated pycache
# entries are deliberately ignored and discovery imports from an empty
# cache prefix.
selected_names = {relative for relative, _ in selected}
metadata_root_names = {
Path(str(item)).parts[0]
for item in declared_files
if Path(str(item)).parts
and Path(str(item)).parts[0].endswith((".dist-info", ".egg-info"))
}
for item in declared_files:
relative = Path(str(item))
text = relative.as_posix()
if (
not relative.parts
or relative.parts[0] in metadata_root_names
or text.startswith("../../../bin/")
or "__pycache__" in relative.parts
):
continue
path = Path(installed.locate_file(item))
if path.is_symlink():
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))
selected_names.add(text)
entry_point_payloads = [
(
f".entry-points/{entry_point.group}/{entry_point.name}",
entry_point.value.encode("utf-8"),
)
for entry_point in installed.entry_points
]
if not selected:
raise ValueError("installed distribution has no measurable package payload")
digest = hashlib.sha256()
for relative, path in sorted(selected):
name = relative.encode("utf-8")
payload = path.read_bytes()
digest.update(len(name).to_bytes(4, "big"))
digest.update(name)
digest.update(len(payload).to_bytes(8, "big"))
digest.update(payload)
for relative, payload in sorted(entry_point_payloads):
name = relative.encode("utf-8")
digest.update(len(name).to_bytes(4, "big"))
digest.update(name)
digest.update(len(payload).to_bytes(8, "big"))
digest.update(payload)
return "sha256:" + digest.hexdigest()
+408
View File
@@ -0,0 +1,408 @@
"""Installed-package manifest and cross-component compatibility contract."""
from __future__ import annotations
import json
import re
from dataclasses import dataclass
from enum import Enum
from types import MappingProxyType
from typing import Any, Mapping
from ._validation import (
canonical_json,
enum_value,
freeze_json_mapping,
require_exact_keys,
require_identifier,
require_positive_int,
require_sha256,
require_schema_version,
require_string,
thaw_json,
)
from .artifacts import PortSpec
from .identity import (
MANIFEST_SCHEMA_VERSION,
ComponentRef,
FeatureRequirement,
VersionRange,
WorkloadId,
)
from .workflow import StageKind, WorkflowSpec
from .schema import validate_schema_definition
class DeterminismProfile(str, Enum):
BYTE_EXACT = "byte_exact"
CANONICAL_EXACT = "canonical_exact"
NUMERIC_TOLERANCE = "numeric_tolerance"
SEEDED_STOCHASTIC = "seeded_stochastic"
SEARCH_OR_OPTIMIZATION = "search_or_optimization"
SIDE_EFFECTING = "side_effecting"
class TrustMode(str, Enum):
TRUSTED = "trusted"
VERIFIED = "verified"
UNTRUSTED_QUORUM = "untrusted_quorum"
@dataclass(frozen=True, slots=True)
class PackageSpec:
distribution: str
digest: str
signature: str | None = None
def __post_init__(self) -> None:
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")
object.__setattr__(self, "distribution", distribution.replace("_", "-"))
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))
def to_dict(self) -> dict[str, object]:
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")
return cls(
distribution=value["distribution"], # type: ignore[arg-type]
digest=value["digest"], # type: ignore[arg-type]
signature=value["signature"], # type: ignore[arg-type]
)
@dataclass(frozen=True, slots=True)
class EnvironmentSpec:
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"))
def to_dict(self) -> dict[str, object]:
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")
return cls(
kind=value["kind"], # type: ignore[arg-type]
digest=value["digest"], # type: ignore[arg-type]
metadata=value["metadata"], # type: ignore[arg-type]
)
@dataclass(frozen=True, slots=True)
class VerifierSpec:
verifier: ComponentRef
configuration: Mapping[str, Any]
def __post_init__(self) -> None:
if not isinstance(self.verifier, ComponentRef):
raise ValueError("verifier must be a ComponentRef")
object.__setattr__(
self,
"configuration",
freeze_json_mapping(self.configuration, "verifier.configuration"),
)
def to_dict(self) -> dict[str, object]:
return {
"verifier": self.verifier.canonical,
"configuration": thaw_json(self.configuration),
}
@classmethod
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")
return cls(
verifier=ComponentRef.from_dict(value["verifier"]),
configuration=value["configuration"], # type: ignore[arg-type]
)
@dataclass(frozen=True, slots=True)
class WorkloadLimits:
max_input_bytes: int
max_tasks: int
max_output_bytes: int
max_parameter_bytes: int = 65_536
max_artifacts: int = 100_000
def __post_init__(self) -> None:
for field in (
"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}"))
def to_dict(self) -> dict[str, int]:
return {
"max_input_bytes": self.max_input_bytes,
"max_tasks": self.max_tasks,
"max_output_bytes": self.max_output_bytes,
"max_parameter_bytes": self.max_parameter_bytes,
"max_artifacts": self.max_artifacts,
}
@classmethod
def from_dict(cls, value: object) -> "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",
}
require_exact_keys(value, fields, "workload limits")
return cls(**value) # type: ignore[arg-type]
def _ports(
value: Mapping[str, PortSpec], field: str, *, allow_empty: bool = False
) -> Mapping[str, PortSpec]:
if not isinstance(value, Mapping) or (not value and not allow_empty):
qualifier = "an object" if allow_empty else "a non-empty object"
raise ValueError(f"{field} must be {qualifier}")
result: dict[str, PortSpec] = {}
for name, port in value.items():
canonical = require_identifier(name, f"{field} port")
if not isinstance(port, PortSpec):
raise ValueError(f"{field} values must be PortSpec values")
result[canonical] = port
return MappingProxyType(result)
@dataclass(frozen=True, slots=True)
class WorkloadManifest:
sdk_api: VersionRange
protocol: VersionRange
workload: WorkloadId
description: str
package: PackageSpec
environment: EnvironmentSpec
parameters_schema: Mapping[str, Any]
workflow: WorkflowSpec
inputs: Mapping[str, PortSpec]
outputs: Mapping[str, PortSpec]
determinism: DeterminismProfile
trust_modes: tuple[TrustMode, ...]
verifier: VerifierSpec
limits: WorkloadLimits
capabilities: tuple[str, ...]
conformance_profiles: tuple[str, ...]
required_features: tuple[FeatureRequirement, ...] = ()
optional_features: tuple[FeatureRequirement, ...] = ()
manifest_schema_version: int = MANIFEST_SCHEMA_VERSION
def __post_init__(self) -> None:
require_schema_version(
self.manifest_schema_version,
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.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")
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")
properties = schema.get("properties")
if not isinstance(properties, Mapping):
raise ValueError("parameters_schema.properties must be an object")
if len(canonical_json(schema).encode("utf-8")) > 1_048_576:
raise ValueError("parameters_schema exceeds 1 MiB")
validate_schema_definition(schema)
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, "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)
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)
manifest_mode_values = {mode.value for mode in modes}
terminal_stage_ids = {
reference.stage_id
for reference in self.workflow.outputs.values()
if reference.stage_id is not None
}
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")
if stage.verifier is None:
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 ()
)
if any(
resources.environment_digest not in {None, self.environment.digest}
for resources in resource_sets
):
raise ValueError(
"stage resource environment must match the manifest environment pin"
)
if not isinstance(self.verifier, VerifierSpec):
raise ValueError("verifier must be a VerifierSpec")
for stage in self.workflow.stages:
if (
stage.stage_id in terminal_stage_ids
and stage.verifier != self.verifier.verifier
):
raise ValueError(
"terminal stage verifier must match the manifest acceptance verifier"
)
if not isinstance(self.limits, WorkloadLimits):
raise ValueError("limits must be WorkloadLimits")
if self.workflow.max_tasks > self.limits.max_tasks:
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)
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)
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):
raise ValueError("features must contain FeatureRequirement values")
names = [item.name for item in required + optional]
if len(names) != len(set(names)):
raise ValueError("required and optional feature names must be unique")
object.__setattr__(self, "required_features", required)
object.__setattr__(self, "optional_features", optional)
self._validate_acceptance_policy()
def _validate_acceptance_policy(self) -> None:
verifier = self.verifier.verifier
exact = verifier == ComponentRef("exact-artifact", 1)
canonical = verifier == ComponentRef("canonical-record", 1)
numeric = verifier == ComponentRef("numeric-tolerance", 1)
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")
if self.determinism is DeterminismProfile.NUMERIC_TOLERANCE and not numeric:
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("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):
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]:
return {
"manifest_schema_version": self.manifest_schema_version,
"sdk_api": self.sdk_api.expression,
"protocol": self.protocol.expression,
"workload": self.workload.to_dict(),
"description": self.description,
"package": self.package.to_dict(),
"environment": self.environment.to_dict(),
"parameters_schema": thaw_json(self.parameters_schema),
"workflow": self.workflow.to_dict(),
"inputs": {name: port.to_dict() for name, port in self.inputs.items()},
"outputs": {name: port.to_dict() for name, port in self.outputs.items()},
"determinism": self.determinism.value,
"trust_modes": [mode.value for mode in self.trust_modes],
"verifier": self.verifier.to_dict(),
"limits": self.limits.to_dict(),
"capabilities": list(self.capabilities),
"conformance_profiles": list(self.conformance_profiles),
"required_features": [item.to_dict() for item in self.required_features],
"optional_features": [item.to_dict() for item in self.optional_features],
}
def to_json(self) -> str:
return canonical_json(self.to_dict())
@classmethod
def from_dict(cls, value: object) -> "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",
}
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"],
)
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")
return cls(
manifest_schema_version=value["manifest_schema_version"], # type: ignore[arg-type]
sdk_api=VersionRange.from_dict(value["sdk_api"]),
protocol=VersionRange.from_dict(value["protocol"]),
workload=WorkloadId.from_dict(value["workload"]),
description=value["description"], # type: ignore[arg-type]
package=PackageSpec.from_dict(value["package"]),
environment=EnvironmentSpec.from_dict(value["environment"]),
parameters_schema=value["parameters_schema"], # type: ignore[arg-type]
workflow=WorkflowSpec.from_dict(value["workflow"]),
inputs={name: PortSpec.from_dict(port) for name, port in inputs.items()},
outputs={name: PortSpec.from_dict(port) for name, port in outputs.items()},
determinism=value["determinism"], # type: ignore[arg-type]
trust_modes=tuple(value["trust_modes"]), # type: ignore[arg-type]
verifier=VerifierSpec.from_dict(value["verifier"]),
limits=WorkloadLimits.from_dict(value["limits"]),
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]
),
optional_features=tuple(
FeatureRequirement.from_dict(item) for item in value["optional_features"] # type: ignore[union-attr]
),
)
@classmethod
def from_json(cls, value: str) -> "WorkloadManifest":
try:
decoded = json.loads(value)
except (TypeError, json.JSONDecodeError, RecursionError) as error:
raise ValueError("workload manifest must be valid JSON") from error
return cls.from_dict(decoded)
+845
View File
@@ -0,0 +1,845 @@
"""Strict job, task, workflow-plan, and expansion value objects."""
from __future__ import annotations
import hashlib
import json
from dataclasses import dataclass, field
from types import MappingProxyType
from typing import Any, Mapping
from ._validation import (
canonical_json,
freeze_json_mapping,
require_exact_keys,
require_identifier,
require_nonnegative_int,
parse_release,
require_positive_int,
require_sha256,
require_schema_version,
require_string,
require_task_key,
require_uuid,
thaw_json,
)
from .artifacts import ArtifactCollection, Cardinality, CollectionKind, PortSpec
from .execution import ExecutionProfile
from .identity import ComponentRef, TASK_SCHEMA_VERSION, WorkloadId
from .manifest import TrustMode
from .resources import ResourceRequirements
from .workflow import StageKind, StageSpec, WorkflowSpec
def _collections(
value: Mapping[str, ArtifactCollection], field: str
) -> Mapping[str, ArtifactCollection]:
if not isinstance(value, Mapping):
raise ValueError(f"{field} must be an object")
result: dict[str, ArtifactCollection] = {}
for name, collection in value.items():
canonical = require_identifier(name, f"{field} port")
if not isinstance(collection, ArtifactCollection):
raise ValueError(f"{field} values must be ArtifactCollection values")
result[canonical] = collection
return MappingProxyType(result)
def _ports(value: Mapping[str, PortSpec], field: str) -> Mapping[str, PortSpec]:
if not isinstance(value, Mapping):
raise ValueError(f"{field} must be an object")
result: dict[str, PortSpec] = {}
for name, port in value.items():
canonical = require_identifier(name, f"{field} port")
if not isinstance(port, PortSpec):
raise ValueError(f"{field} values must be PortSpec values")
result[canonical] = port
return MappingProxyType(result)
def _feature_versions(value: object, field: str) -> Mapping[str, str]:
if not isinstance(value, Mapping):
raise ValueError(f"{field} must be an object")
result: dict[str, str] = {}
for name, version in value.items():
canonical = require_identifier(name, f"{field} feature")
text = require_string(version, f"{field} version", max_length=32)
parse_release(text, f"{field} version")
result[canonical] = text
return MappingProxyType(result)
def _fallbacks(value: object, field: str) -> Mapping[str, str]:
if not isinstance(value, Mapping):
raise ValueError(f"{field} must be an object")
return MappingProxyType(
{
require_identifier(name, f"{field} feature"): require_identifier(
fallback,
f"{field} fallback",
)
for name, fallback in value.items()
}
)
@dataclass(frozen=True, slots=True)
class JobRequest:
workload: WorkloadId
parameters: Mapping[str, Any]
inputs: Mapping[str, ArtifactCollection]
required_features: tuple[str, ...] = ()
trust_mode: TrustMode = TrustMode.TRUSTED
def __post_init__(self) -> None:
if not isinstance(self.workload, WorkloadId):
raise ValueError("job workload must be a WorkloadId")
object.__setattr__(
self,
"parameters",
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)
if len(features) != len(set(features)):
raise ValueError("required_features must be unique")
object.__setattr__(self, "required_features", features)
try:
trust_mode = TrustMode(self.trust_mode)
except (TypeError, ValueError) as error:
raise ValueError("job trust_mode is unsupported") from error
object.__setattr__(self, "trust_mode", trust_mode)
@property
def parameters_digest(self) -> str:
return hashlib.sha256(canonical_json(self.parameters).encode("utf-8")).hexdigest()
def to_dict(self) -> dict[str, object]:
return {
"workload": self.workload.to_dict(),
"parameters": thaw_json(self.parameters),
"inputs": {name: value.to_dict() for name, value in self.inputs.items()},
"required_features": list(self.required_features),
"trust_mode": self.trust_mode.value,
}
def to_json(self) -> str:
return canonical_json(self.to_dict())
@classmethod
def from_dict(cls, value: object) -> "JobRequest":
if not isinstance(value, Mapping):
raise ValueError("job request must be an object")
fields = {"workload", "parameters", "inputs", "required_features", "trust_mode"}
require_exact_keys(value, fields, "job request")
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")
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()},
required_features=tuple(features),
trust_mode=value["trust_mode"], # type: ignore[arg-type]
)
@classmethod
def from_json(cls, value: str) -> "JobRequest":
try:
decoded = json.loads(value)
except (TypeError, json.JSONDecodeError, RecursionError) as error:
raise ValueError("job request must be valid JSON") from error
return cls.from_dict(decoded)
@dataclass(frozen=True, slots=True)
class ValidatedJob:
request: JobRequest
resolved_parameters: Mapping[str, Any]
def __post_init__(self) -> None:
if not isinstance(self.request, JobRequest):
raise ValueError("validated job request must be a JobRequest")
object.__setattr__(
self,
"resolved_parameters",
freeze_json_mapping(
self.resolved_parameters,
"resolved_parameters",
forbid_locations=True,
),
)
@property
def parameters_digest(self) -> str:
return hashlib.sha256(canonical_json(self.resolved_parameters).encode("utf-8")).hexdigest()
@dataclass(frozen=True, slots=True)
class TaskSpec:
workload: WorkloadId
package_digest: str
manifest_digest: str
trust_mode: TrustMode
sdk_api_version: str
protocol_version: str
manifest_schema_version: int
workflow_schema_version: int
environment_digest: str
verifier: ComponentRef
selected_features: Mapping[str, str]
optional_fallbacks: Mapping[str, str]
task_key: str
stage_id: str
parameters: Mapping[str, Any]
inputs: Mapping[str, ArtifactCollection]
expected_outputs: Mapping[str, PortSpec]
resources: ResourceRequirements
execution: ExecutionProfile
expected_input_keys: Mapping[str, tuple[str, ...]] = field(default_factory=dict)
schema_version: int = TASK_SCHEMA_VERSION
def __post_init__(self) -> None:
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__(
self,
"package_digest",
require_sha256(self.package_digest, "task package_digest", prefixed=True),
)
object.__setattr__(
self,
"manifest_digest",
require_sha256(self.manifest_digest, "task manifest_digest"),
)
try:
trust_mode = TrustMode(self.trust_mode)
except (TypeError, ValueError) as error:
raise ValueError("task trust_mode is unsupported") from error
object.__setattr__(self, "trust_mode", trust_mode)
object.__setattr__(
self,
"sdk_api_version",
require_string(self.sdk_api_version, "task sdk_api_version", max_length=32),
)
object.__setattr__(
self,
"protocol_version",
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"),
)
object.__setattr__(
self,
"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),
)
if not isinstance(self.verifier, ComponentRef):
raise ValueError("task verifier must be a ComponentRef")
object.__setattr__(
self,
"selected_features",
_feature_versions(self.selected_features, "task selected_features"),
)
object.__setattr__(
self,
"optional_fallbacks",
_fallbacks(self.optional_fallbacks, "task optional_fallbacks"),
)
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,
"parameters",
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"))
if not self.expected_outputs:
raise ValueError("task expected_outputs must not be empty")
if not isinstance(self.resources, ResourceRequirements):
raise ValueError("task resources must be ResourceRequirements")
if not isinstance(self.execution, ExecutionProfile):
raise ValueError("task execution must be ExecutionProfile")
self.execution.validate_resources(self.resources)
if not isinstance(self.expected_input_keys, Mapping):
raise ValueError("expected_input_keys must be an object")
expected_keys: dict[str, tuple[str, ...]] = {}
for port_name, keys in self.expected_input_keys.items():
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
))
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
object.__setattr__(self, "expected_input_keys", MappingProxyType(expected_keys))
def validate_stage(self, stage: StageSpec) -> "TaskSpec":
if not isinstance(stage, StageSpec) or stage.stage_id != self.stage_id:
raise ValueError("task stage does not match its StageSpec")
if set(self.inputs) != set(stage.inputs):
raise ValueError("task input ports do not match the stage")
for name, declaration in stage.inputs.items():
declaration.validate_collection(self.inputs[name], f"task input {name}")
for name, expected_keys in self.expected_input_keys.items():
declaration = stage.inputs.get(name)
if (
declaration is None
or declaration.cardinality is not Cardinality.MANY
or declaration.collection is not CollectionKind.KEYED
):
raise ValueError("expected input keys require a keyed-many stage input")
actual_keys = tuple(
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")
keyed_many_ports = {
name
for name, declaration in stage.inputs.items()
if declaration.cardinality is Cardinality.MANY
and declaration.collection is CollectionKind.KEYED
}
if set(self.expected_input_keys) != keyed_many_ports:
raise ValueError("task must pin expected keys for every keyed-many input")
if dict(self.expected_outputs) != dict(stage.outputs):
raise ValueError("task expected outputs do not match the stage")
if not set(self.parameters).issubset(stage.parameter_names):
raise ValueError("task parameters are outside the stage projection")
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")
if self.trust_mode.value not in stage.trust_modes:
raise ValueError("task trust mode is not allowed by the stage")
return self
def to_dict(self) -> dict[str, object]:
return {
"schema_version": self.schema_version,
"workload": self.workload.to_dict(),
"package_digest": self.package_digest,
"manifest_digest": self.manifest_digest,
"trust_mode": self.trust_mode.value,
"sdk_api_version": self.sdk_api_version,
"protocol_version": self.protocol_version,
"manifest_schema_version": self.manifest_schema_version,
"workflow_schema_version": self.workflow_schema_version,
"environment_digest": self.environment_digest,
"verifier": self.verifier.canonical,
"selected_features": dict(self.selected_features),
"optional_fallbacks": dict(self.optional_fallbacks),
"task_key": self.task_key,
"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()},
"resources": self.resources.to_dict(),
"execution": self.execution.to_dict(),
"expected_input_keys": {
name: list(keys) for name, keys in self.expected_input_keys.items()
},
}
def to_json(self) -> str:
return canonical_json(self.to_dict())
@property
def digest(self) -> str:
"""Canonical digest used to pin a coordinator execution contract."""
return hashlib.sha256(self.to_json().encode("utf-8")).hexdigest()
@classmethod
def from_dict(cls, value: object) -> "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",
}
require_exact_keys(value, fields, "task specification")
inputs, outputs = value["inputs"], value["expected_outputs"]
if not isinstance(inputs, Mapping) or not isinstance(outputs, Mapping):
raise ValueError("task inputs and expected_outputs must be objects")
return cls(
schema_version=value["schema_version"], # type: ignore[arg-type]
workload=WorkloadId.from_dict(value["workload"]),
package_digest=value["package_digest"], # type: ignore[arg-type]
manifest_digest=value["manifest_digest"], # type: ignore[arg-type]
trust_mode=value["trust_mode"], # type: ignore[arg-type]
sdk_api_version=value["sdk_api_version"], # type: ignore[arg-type]
protocol_version=value["protocol_version"], # type: ignore[arg-type]
manifest_schema_version=value["manifest_schema_version"], # type: ignore[arg-type]
workflow_schema_version=value["workflow_schema_version"], # type: ignore[arg-type]
environment_digest=value["environment_digest"], # type: ignore[arg-type]
verifier=ComponentRef.from_dict(value["verifier"]),
selected_features=value["selected_features"], # type: ignore[arg-type]
optional_fallbacks=value["optional_fallbacks"], # type: ignore[arg-type]
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()},
resources=ResourceRequirements.from_dict(value["resources"]),
execution=ExecutionProfile.from_dict(value["execution"]),
expected_input_keys=value["expected_input_keys"], # type: ignore[arg-type]
)
@classmethod
def from_json(cls, value: str) -> "TaskSpec":
try:
decoded = json.loads(value)
except (TypeError, json.JSONDecodeError, RecursionError) as error:
raise ValueError("task specification must be valid JSON") from error
return cls.from_dict(decoded)
@dataclass(frozen=True, slots=True)
class WorkflowPlan:
workload: WorkloadId
package_digest: str
manifest_digest: str
trust_mode: TrustMode
sdk_api_version: str
protocol_version: str
manifest_schema_version: int
workflow_schema_version: int
environment_digest: str
verifier: ComponentRef
selected_features: Mapping[str, str]
optional_fallbacks: Mapping[str, str]
workflow_id: str
resolved_parameters: Mapping[str, Any]
tasks: tuple[TaskSpec, ...]
schema_version: int = 1
def __post_init__(self) -> None:
require_schema_version(self.schema_version, 1, "workflow plan schema_version")
if not isinstance(self.workload, WorkloadId):
raise ValueError("workflow plan workload must be a WorkloadId")
object.__setattr__(
self,
"package_digest",
require_sha256(self.package_digest, "plan package_digest", prefixed=True),
)
object.__setattr__(
self,
"manifest_digest",
require_sha256(self.manifest_digest, "plan manifest_digest"),
)
try:
trust_mode = TrustMode(self.trust_mode)
except (TypeError, ValueError) as error:
raise ValueError("plan trust_mode is unsupported") from error
object.__setattr__(self, "trust_mode", trust_mode)
object.__setattr__(
self,
"sdk_api_version",
require_string(self.sdk_api_version, "plan sdk_api_version", max_length=32),
)
object.__setattr__(
self,
"protocol_version",
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"),
)
object.__setattr__(
self,
"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),
)
if not isinstance(self.verifier, ComponentRef):
raise ValueError("plan verifier must be a ComponentRef")
object.__setattr__(
self,
"selected_features",
_feature_versions(self.selected_features, "plan selected_features"),
)
object.__setattr__(
self,
"optional_fallbacks",
_fallbacks(self.optional_fallbacks, "plan optional_fallbacks"),
)
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,
"resolved_parameters",
freeze_json_mapping(
self.resolved_parameters,
"resolved_parameters",
forbid_locations=True,
),
)
tasks = tuple(self.tasks)
if not tasks or any(not isinstance(task, TaskSpec) for task in tasks):
raise ValueError("workflow plan tasks must contain at least one TaskSpec")
keys = [task.task_key for task in tasks]
if keys != sorted(keys) or len(keys) != len(set(keys)):
raise ValueError("workflow plan task keys must be unique and ascending")
for task in tasks:
if (
task.workload != self.workload
or task.package_digest != self.package_digest
or task.manifest_digest != self.manifest_digest
or task.trust_mode is not self.trust_mode
or task.sdk_api_version != self.sdk_api_version
or task.protocol_version != self.protocol_version
or task.manifest_schema_version != self.manifest_schema_version
or task.workflow_schema_version != self.workflow_schema_version
or task.environment_digest != self.environment_digest
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")
object.__setattr__(self, "tasks", tasks)
def validate_workflow(self, workflow: WorkflowSpec) -> "WorkflowPlan":
if workflow.workflow_id != self.workflow_id:
raise ValueError("workflow plan references another workflow")
if len(self.tasks) > workflow.max_tasks:
raise ValueError("workflow plan exceeds max_tasks")
stages = {stage.stage_id: stage for stage in workflow.stages}
task_counts: dict[str, int] = {}
for task in self.tasks:
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
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}")
return self
@property
def digest(self) -> str:
return hashlib.sha256(self.to_json().encode("utf-8")).hexdigest()
def to_dict(self) -> dict[str, object]:
return {
"schema_version": self.schema_version,
"workload": self.workload.to_dict(),
"package_digest": self.package_digest,
"manifest_digest": self.manifest_digest,
"trust_mode": self.trust_mode.value,
"sdk_api_version": self.sdk_api_version,
"protocol_version": self.protocol_version,
"manifest_schema_version": self.manifest_schema_version,
"workflow_schema_version": self.workflow_schema_version,
"environment_digest": self.environment_digest,
"verifier": self.verifier.canonical,
"selected_features": dict(self.selected_features),
"optional_fallbacks": dict(self.optional_fallbacks),
"workflow_id": self.workflow_id,
"resolved_parameters": thaw_json(self.resolved_parameters),
"tasks": [task.to_dict() for task in self.tasks],
}
def to_json(self) -> str:
return canonical_json(self.to_dict())
@classmethod
def from_dict(cls, value: object) -> "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",
}
require_exact_keys(value, fields, "workflow plan")
tasks = value["tasks"]
if not isinstance(tasks, list):
raise ValueError("workflow plan tasks must be an array")
return cls(
schema_version=value["schema_version"], # type: ignore[arg-type]
workload=WorkloadId.from_dict(value["workload"]),
package_digest=value["package_digest"], # type: ignore[arg-type]
manifest_digest=value["manifest_digest"], # type: ignore[arg-type]
trust_mode=value["trust_mode"], # type: ignore[arg-type]
sdk_api_version=value["sdk_api_version"], # type: ignore[arg-type]
protocol_version=value["protocol_version"], # type: ignore[arg-type]
manifest_schema_version=value["manifest_schema_version"], # type: ignore[arg-type]
workflow_schema_version=value["workflow_schema_version"], # type: ignore[arg-type]
environment_digest=value["environment_digest"], # type: ignore[arg-type]
verifier=ComponentRef.from_dict(value["verifier"]),
selected_features=value["selected_features"], # type: ignore[arg-type]
optional_fallbacks=value["optional_fallbacks"], # type: ignore[arg-type]
workflow_id=value["workflow_id"], # type: ignore[arg-type]
resolved_parameters=value["resolved_parameters"], # type: ignore[arg-type]
tasks=tuple(TaskSpec.from_dict(task) for task in tasks),
)
@classmethod
def from_json(cls, value: str) -> "WorkflowPlan":
try:
decoded = json.loads(value)
except (TypeError, json.JSONDecodeError, RecursionError) as error:
raise ValueError("workflow plan must be valid JSON") from error
return cls.from_dict(decoded)
@dataclass(frozen=True, slots=True)
class ExpansionManifest:
job_id: str
parent_task_id: str
parent_task_key: str
parent_execution_contract_digest: str
tasks: tuple[TaskSpec, ...]
max_children: int
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"))
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_execution_contract_digest",
require_sha256(
self.parent_execution_contract_digest,
"expansion parent_execution_contract_digest",
),
)
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")
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")
first = tasks[0]
if any(
task.workload != first.workload
or task.package_digest != first.package_digest
or task.manifest_digest != first.manifest_digest
or task.trust_mode is not first.trust_mode
or task.sdk_api_version != first.sdk_api_version
or task.protocol_version != first.protocol_version
or task.manifest_schema_version != first.manifest_schema_version
or task.workflow_schema_version != first.workflow_schema_version
or task.environment_digest != first.environment_digest
or task.selected_features != first.selected_features
or task.optional_fallbacks != first.optional_fallbacks
for task in tasks[1:]
):
raise ValueError("expansion child tasks must carry one exact workload pin")
object.__setattr__(self, "tasks", tasks)
def validate_against(
self,
parent: TaskSpec,
workflow: WorkflowSpec,
*,
job_id: str,
parent_task_id: str,
declared_max_children: int,
remaining_tasks: int,
authorized_inputs: Mapping[str, Mapping[str, ArtifactCollection]],
existing_stage_task_counts: Mapping[str, int],
) -> "ExpansionManifest":
"""Validate an expansion against coordinator-owned durable state.
The IDs and remaining budget are deliberately supplied by the
coordinator rather than trusted from the package-produced manifest.
"""
if not isinstance(parent, TaskSpec):
raise ValueError("expansion parent must be a TaskSpec")
if not isinstance(workflow, WorkflowSpec):
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"):
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")
if self.parent_execution_contract_digest != parent.digest:
raise ValueError("expansion parent execution contract does not match")
remaining = require_nonnegative_int(remaining_tasks, "remaining_tasks")
stages = {stage.stage_id: stage for stage in workflow.stages}
try:
parent_stage = stages[parent.stage_id]
except KeyError as 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")
declared_limit = require_positive_int(
declared_max_children,
"declared_max_children",
)
allowed_children = min(declared_limit, remaining)
if self.max_children > declared_limit or len(self.tasks) > allowed_children:
raise ValueError("expansion exceeds the coordinator child task budget")
if not isinstance(authorized_inputs, Mapping):
raise ValueError("authorized_inputs must be an object")
allowed_by_target: dict[str, dict[str, ArtifactCollection]] = {}
for stage_id, ports in authorized_inputs.items():
canonical_stage = require_identifier(stage_id, "authorized input stage")
if canonical_stage not in stages or not isinstance(ports, Mapping):
raise ValueError("authorized_inputs references an unknown stage")
allowed_ports: dict[str, ArtifactCollection] = {}
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")
declaration.validate_collection(
collection,
f"authorized input {canonical_stage}.{canonical_port}",
)
allowed_ports[canonical_port] = collection
allowed_by_target[canonical_stage] = allowed_ports
raw_counts = existing_stage_task_counts
if not isinstance(raw_counts, Mapping):
raise ValueError("existing_stage_task_counts must be an object")
stage_counts: dict[str, int] = {}
for stage_id, count in raw_counts.items():
canonical = require_identifier(stage_id, "existing stage task count")
if canonical not in stages:
raise ValueError("existing task count references an unknown stage")
stage_counts[canonical] = require_nonnegative_int(
count,
"existing stage task count",
)
for task in self.tasks:
if (
task.workload != parent.workload
or task.package_digest != parent.package_digest
or task.manifest_digest != parent.manifest_digest
or task.trust_mode is not parent.trust_mode
or task.sdk_api_version != parent.sdk_api_version
or task.protocol_version != parent.protocol_version
or task.manifest_schema_version != parent.manifest_schema_version
or task.workflow_schema_version != parent.workflow_schema_version
or task.environment_digest != parent.environment_digest
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")
try:
stage = stages[task.stage_id]
except KeyError as 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")
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")
if collection.kind is CollectionKind.ORDERED:
cursor = 0
for item in collection.items:
while cursor < len(allowed.items) and allowed.items[cursor] != item:
cursor += 1
if cursor == len(allowed.items):
raise ValueError(
"expansion child input is not an authorized ordered subsequence"
)
cursor += 1
elif any(item not in allowed.items for item in collection.items):
raise ValueError(
"expansion child input artifact is not coordinator-authorized"
)
stage_counts[task.stage_id] = stage_counts.get(task.stage_id, 0) + 1
if stage_counts[task.stage_id] > stage.max_fan_out:
raise ValueError(
f"expansion exceeds max_fan_out for stage {task.stage_id}"
)
return self
@property
def digest(self) -> str:
return hashlib.sha256(canonical_json(self.to_dict()).encode("utf-8")).hexdigest()
def to_dict(self) -> dict[str, object]:
return {
"schema_version": self.schema_version,
"job_id": self.job_id,
"parent_task_id": self.parent_task_id,
"parent_task_key": self.parent_task_key,
"parent_execution_contract_digest": self.parent_execution_contract_digest,
"max_children": self.max_children,
"tasks": [task.to_dict() for task in self.tasks],
}
def to_json(self) -> str:
return canonical_json(self.to_dict())
@classmethod
def from_dict(cls, value: object) -> "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",
}
require_exact_keys(value, fields, "expansion manifest")
tasks = value["tasks"]
if not isinstance(tasks, list):
raise ValueError("expansion tasks must be an array")
return cls(
schema_version=value["schema_version"], # type: ignore[arg-type]
job_id=value["job_id"], # type: ignore[arg-type]
parent_task_id=value["parent_task_id"], # type: ignore[arg-type]
parent_task_key=value["parent_task_key"], # type: ignore[arg-type]
parent_execution_contract_digest=value["parent_execution_contract_digest"], # type: ignore[arg-type]
max_children=value["max_children"], # type: ignore[arg-type]
tasks=tuple(TaskSpec.from_dict(task) for task in tasks),
)
@classmethod
def from_json(cls, value: str) -> "ExpansionManifest":
try:
decoded = json.loads(value)
except (TypeError, json.JSONDecodeError, RecursionError) as error:
raise ValueError("expansion manifest must be valid JSON") from error
return cls.from_dict(decoded)
+108
View File
@@ -0,0 +1,108 @@
"""Author-facing planner, runner, reducer, and verifier protocols."""
from __future__ import annotations
from pathlib import Path
from typing import Mapping, Protocol, Sequence
from .artifacts import ArtifactCollection, ArtifactRef, ArtifactSchema, OutputManifest, Provenance
from .plans import JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
from .runtime import NegotiatedWorkload
from .verification import CandidateOutputs, VerificationDecision, VerifyContext
class ArtifactCatalog(Protocol):
"""Bridge-owned, read-only access to durable input artifacts."""
def materialize(self, artifact: ArtifactRef) -> Path:
"""Return an attempt-scoped verified local copy without exposing credentials."""
class ArtifactSink(Protocol):
"""Agent/bridge-owned sealing boundary for scientific output files."""
def seal(
self,
path: Path,
*,
declaration: ArtifactSchema,
records: int | None = None,
dimensions: tuple[int, ...] = (),
) -> ArtifactRef:
"""Validate/upload bytes and return coordinator-owned immutable metadata."""
class CancellationToken(Protocol):
def cancelled(self) -> bool: ...
def raise_if_cancelled(self) -> None: ...
class PlanningResources(Protocol):
"""Caller-provided catalog, sink, and workspace for registry planning."""
@property
def catalog(self) -> ArtifactCatalog: ...
@property
def sink(self) -> ArtifactSink: ...
@property
def workspace(self) -> Path: ...
class PlanningContext(PlanningResources, Protocol):
"""Planner-facing resources augmented by completed negotiation."""
@property
def negotiated(self) -> NegotiatedWorkload:
"""Resolved optional fallbacks and the exact negotiated manifest."""
class TaskContext(Protocol):
@property
def task(self) -> TaskSpec: ...
@property
def catalog(self) -> ArtifactCatalog: ...
@property
def sink(self) -> ArtifactSink: ...
@property
def workspace(self) -> Path: ...
@property
def cancellation(self) -> CancellationToken: ...
@property
def provenance(self) -> Provenance: ...
class ReduceContext(TaskContext, Protocol):
@property
def accepted_inputs(self) -> Mapping[str, ArtifactCollection]: ...
class Planner(Protocol):
entry_point: str
def validate(self, request: JobRequest) -> ValidatedJob: ...
def plan(self, job: ValidatedJob, context: PlanningContext) -> WorkflowPlan: ...
class Runner(Protocol):
def run(self, context: TaskContext) -> OutputManifest: ...
class Reducer(Protocol):
def reduce(self, context: ReduceContext) -> OutputManifest: ...
class Verifier(Protocol):
def verify(
self,
context: VerifyContext,
candidates: CandidateOutputs,
) -> VerificationDecision: ...
+526
View File
@@ -0,0 +1,526 @@
"""Explicit, digest-pinned workload package registry and safe discovery."""
from __future__ import annotations
import re
import sys
from dataclasses import dataclass
from importlib import machinery, util
from importlib import metadata
from pathlib import Path
from tempfile import TemporaryDirectory
from threading import RLock
from types import MappingProxyType
from typing import Any, Mapping
from ._validation import (
canonical_json,
require_semver,
require_sha256,
require_string,
require_workload_name,
)
from .identity import ComponentRef, WorkloadId
from .integrity import installed_distribution_digest
from .manifest import WorkloadManifest
from .plans import JobRequest, ValidatedJob, WorkflowPlan
from .protocols import Planner, PlanningContext, PlanningResources, Reducer, Runner, Verifier
from .runtime import CompatibilityError, NegotiatedWorkload, RuntimeCapabilities, negotiate_manifest
from .schema import validate_parameter_instance
from .workflow import StageKind
_DISCOVERY_IMPORT_LOCK = RLock()
def _normalized_distribution_name(value: str) -> str:
return re.sub(r"[-_.]+", "-", value).lower()
def _validate_entry_point_ownership(entry_point: metadata.EntryPoint) -> None:
"""Require the entry-point module to be payload of its own distribution."""
distribution = entry_point.dist
if distribution is None:
raise ValueError("workload entry point has no owning distribution")
module_name = getattr(entry_point, "module", None)
if not isinstance(module_name, str) or not module_name:
value = getattr(entry_point, "value", "")
module_name = value.partition(":")[0].strip() if isinstance(value, str) else ""
parts = module_name.split(".")
if not parts or any(not part.isidentifier() for part in parts):
raise ValueError("workload entry point has an invalid module path")
raw_top_level = distribution.read_text("top_level.txt")
declared = {
line.strip()
for line in raw_top_level.splitlines()
if line.strip()
} if raw_top_level is not None else set()
root_name = parts[0]
if root_name not in declared:
raise ValueError("workload entry point module is outside its distribution")
owners = metadata.packages_distributions().get(root_name, ())
normalized_owners = {_normalized_distribution_name(owner) for owner in owners}
expected_owner = _normalized_distribution_name(distribution.name)
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))
if not package_root.exists():
root_spec = util.find_spec(root_name)
locations = (
tuple(root_spec.submodule_search_locations or ())
if root_spec is not None
else ()
)
if len(locations) == 1:
package_root = Path(locations[0])
if package_root.is_dir():
module_base = package_root.joinpath(*parts[1:])
ownership_root = package_root.resolve()
candidates = [
*(Path(str(module_base) + suffix) for suffix in machinery.SOURCE_SUFFIXES),
*(Path(str(module_base) + suffix) for suffix in machinery.EXTENSION_SUFFIXES),
*(module_base / ("__init__" + suffix) for suffix in machinery.SOURCE_SUFFIXES),
*(module_base / ("__init__" + suffix) for suffix in machinery.EXTENSION_SUFFIXES),
]
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))
candidates = [
*(Path(str(module_base) + suffix) for suffix in machinery.SOURCE_SUFFIXES),
*(Path(str(module_base) + suffix) for suffix in machinery.EXTENSION_SUFFIXES),
]
existing = tuple(candidate for candidate in candidates if candidate.is_file())
if len(existing) != 1 or not existing[0].resolve().is_relative_to(ownership_root):
raise ValueError("workload entry point module is not an owned package payload")
@dataclass(frozen=True, slots=True)
class WorkloadDefinition:
manifest: WorkloadManifest
planner: Planner
runners: Mapping[str, Runner]
reducers: Mapping[str, Reducer]
verifiers: Mapping[str, Verifier]
def __post_init__(self) -> None:
if not isinstance(self.manifest, WorkloadManifest):
raise ValueError("definition manifest must be a WorkloadManifest")
if not callable(getattr(self.planner, "validate", None)) or not callable(
getattr(self.planner, "plan", None)
):
raise ValueError("definition planner must implement validate and plan")
collections: list[tuple[str, Mapping[str, Any], str]] = [
("runners", self.runners, "run"),
("reducers", self.reducers, "reduce"),
("verifiers", self.verifiers, "verify"),
]
for field, values, method in collections:
if not isinstance(values, Mapping):
raise ValueError(f"definition {field} must be an object")
copied: dict[str, Any] = {}
for name, handler in values.items():
canonical = require_string(name, f"{field} entry point", max_length=256)
if not callable(getattr(handler, method, None)):
raise ValueError(f"definition {field} handler must implement {method}")
copied[canonical] = handler
object.__setattr__(self, field, MappingProxyType(copied))
for stage in self.manifest.workflow.stages:
if stage.kind is StageKind.PLAN:
if getattr(self.planner, "entry_point", None) != stage.entry_point:
raise ValueError(
"PLAN stage entry point must match planner.entry_point"
)
continue
if stage.kind is StageKind.REDUCE:
handlers = self.reducers
else:
# A VERIFY node is still an executable DAG stage. Its
# ``entry_point`` is a Runner; ``stage.verifier`` selects the
# independent acceptance component applied to its output.
handlers = self.runners
if stage.entry_point not in handlers:
raise ValueError(
f"definition has no installed handler for stage entry point: {stage.entry_point}"
)
verifier_key = self.manifest.verifier.verifier.canonical
if verifier_key not in self.verifiers:
raise ValueError(f"definition has no installed manifest verifier: {verifier_key}")
for key, verifier in self.verifiers.items():
try:
declared_identity = ComponentRef.from_dict(key)
except ValueError as error:
raise ValueError("definition verifier keys must be component identities") from error
if declared_identity.canonical != key or getattr(verifier, "identity", None) != declared_identity:
raise ValueError("definition verifier handler identity does not match its key")
manifest_verifier = self.verifiers[verifier_key]
handler_configuration = getattr(manifest_verifier, "configuration", None)
if handler_configuration is None:
if self.manifest.verifier.configuration:
raise ValueError("manifest verifier configuration is not bound by its handler")
elif dict(handler_configuration) != dict(self.manifest.verifier.configuration):
raise ValueError("manifest verifier configuration does not match its handler")
for stage in self.manifest.workflow.stages:
if stage.verifier is not None and stage.verifier.canonical not in self.verifiers:
raise ValueError(
f"definition has no installed stage verifier: {stage.verifier.canonical}"
)
@dataclass(frozen=True, slots=True)
class AllowedPackage:
distribution: str
workload: WorkloadId
digest: str
def __post_init__(self) -> None:
distribution = require_string(self.distribution, "distribution", max_length=128).lower()
if not re.fullmatch(r"[a-z0-9]+(?:[-_.][a-z0-9]+)*", distribution):
raise ValueError("distribution must be a canonical Python distribution name")
object.__setattr__(self, "distribution", distribution.replace("_", "-"))
if not isinstance(self.workload, WorkloadId):
raise ValueError("allowed workload must be a WorkloadId")
object.__setattr__(self, "digest", require_sha256(self.digest, "allowed digest", prefixed=True))
@dataclass(frozen=True, slots=True)
class WorkloadDescription:
workload: WorkloadId
description: str
package_digest: str
enabled: bool
@dataclass(frozen=True, slots=True)
class _NegotiatedPlanningContext:
base: PlanningResources
negotiated: NegotiatedWorkload
@property
def catalog(self):
return self.base.catalog
@property
def sink(self):
return self.base.sink
@property
def workspace(self) -> Path:
return self.base.workspace
class WorkloadRegistry:
"""Registry keyed by exact workload version and immutable package digest."""
ENTRY_POINT_GROUP = "scimesh.workloads"
def __init__(self) -> None:
self._definitions: dict[tuple[str, str], WorkloadDefinition] = {}
self._enabled: set[tuple[str, str, str]] = set()
self._lock = RLock()
def register(self, definition: WorkloadDefinition, *, enabled: bool = False) -> None:
if not isinstance(definition, WorkloadDefinition):
raise ValueError("definition must be a WorkloadDefinition")
workload = definition.manifest.workload
key = (workload.name, workload.version)
with self._lock:
if key in self._definitions:
raise ValueError(f"workload version already registered: {workload.name}@{workload.version}")
self._definitions[key] = definition
if enabled:
self._enabled.add((*key, definition.manifest.package.digest))
def enable(self, name: str, version: str, package_digest: str) -> None:
digest = require_sha256(package_digest, "package_digest", prefixed=True)
with self._lock:
definition = self._registered(name, version)
if digest != definition.manifest.package.digest:
raise ValueError("package digest does not match the registered manifest")
self._enabled.add((definition.manifest.workload.name, version, digest))
def disable(self, name: str, version: str, package_digest: str) -> None:
canonical = require_workload_name(name)
version = require_semver(version, "workload.version")
digest = require_sha256(package_digest, "package_digest", prefixed=True)
with self._lock:
self._enabled.discard((canonical, version, digest))
def _registered(self, name: str, version: str) -> WorkloadDefinition:
canonical = require_workload_name(name)
version = require_semver(version, "workload.version")
with self._lock:
try:
return self._definitions[(canonical, version)]
except KeyError as error:
raise ValueError(f"unknown workload version: {canonical}@{version}") from error
def require(
self,
name: str,
version: str,
package_digest: str,
*,
runtime: RuntimeCapabilities | None = None,
) -> tuple[WorkloadDefinition, NegotiatedWorkload | None]:
digest = require_sha256(package_digest, "package_digest", prefixed=True)
with self._lock:
definition = self._registered(name, version)
identity = (definition.manifest.workload.name, definition.manifest.workload.version, digest)
if digest != definition.manifest.package.digest or identity not in self._enabled:
raise ValueError("workload package digest is not enabled")
negotiated = negotiate_manifest(definition.manifest, runtime) if runtime is not None else None
return definition, negotiated
def plan(
self,
request: JobRequest,
package_digest: str,
runtime: RuntimeCapabilities,
context: PlanningResources,
) -> WorkflowPlan:
"""Negotiate first, then invoke only the pre-registered planner object."""
if not isinstance(request, JobRequest):
raise ValueError("request must be a JobRequest")
definition, negotiated = self.require(
request.workload.name,
request.workload.version,
package_digest,
runtime=runtime,
)
assert negotiated is not None
self._validate_request_compatibility(request, definition.manifest, runtime, negotiated)
self._validate_request_shape(request, definition.manifest)
validated = definition.planner.validate(request)
if not isinstance(validated, ValidatedJob) or validated.request != request:
raise ValueError("planner.validate must return a ValidatedJob for the same request")
plan = definition.planner.plan(
validated,
_NegotiatedPlanningContext(context, negotiated),
)
if not isinstance(plan, WorkflowPlan) or plan.workload != request.workload:
raise ValueError("planner.plan must return a WorkflowPlan for the requested workload")
if (
plan.package_digest != definition.manifest.package.digest
or plan.manifest_digest != definition.manifest.digest
or plan.trust_mode is not request.trust_mode
or plan.sdk_api_version != runtime.sdk_api_version
or plan.protocol_version != runtime.protocol_version
or plan.manifest_schema_version != definition.manifest.manifest_schema_version
or plan.workflow_schema_version != definition.manifest.workflow.schema_version
or plan.environment_digest != definition.manifest.environment.digest
or plan.verifier != definition.manifest.verifier.verifier
or plan.selected_features != negotiated.selected_features
or plan.optional_fallbacks != negotiated.optional_fallbacks
):
raise ValueError("planner plan does not carry the selected immutable workload pin")
plan.validate_workflow(definition.manifest.workflow)
self._validate_plan_limits(request, plan, definition.manifest)
return WorkflowPlan.from_json(plan.to_json())
@staticmethod
def _validate_request_compatibility(
request: JobRequest,
manifest: WorkloadManifest,
runtime: RuntimeCapabilities,
negotiated: NegotiatedWorkload,
) -> None:
if request.trust_mode not in manifest.trust_modes:
raise CompatibilityError(
"trust-mode-undeclared",
"requested trust mode is not declared by the workload",
)
if request.trust_mode not in runtime.trust_modes:
raise CompatibilityError(
"trust-mode-unavailable",
"runtime cannot enforce the requested trust mode",
)
for stage in manifest.workflow.stages:
if request.trust_mode.value not in stage.trust_modes:
raise CompatibilityError(
"stage-trust-unavailable",
f"stage {stage.stage_id} does not support the requested trust mode",
)
declared = {
feature.name: feature
for feature in manifest.required_features + manifest.optional_features
}
for name in request.required_features:
requirement = declared.get(name)
if requirement is None:
raise CompatibilityError(
"feature-undeclared",
f"job requests a feature not declared by the workload: {name}",
)
version = runtime.features.get(name)
if version is None or not requirement.versions.contains(version):
raise CompatibilityError(
"feature-unavailable",
f"job-required feature is unavailable or incompatible: {name}",
)
if name in negotiated.optional_fallbacks:
raise CompatibilityError(
"feature-fallback-disallowed",
f"job-required feature cannot use its fallback: {name}",
)
@staticmethod
def _validate_request_shape(request: JobRequest, manifest: WorkloadManifest) -> None:
if set(request.inputs) != set(manifest.inputs):
raise ValueError("job input ports do not match the manifest")
total_bytes = 0
artifact_references: dict[str, object] = {}
for name, port in manifest.inputs.items():
port.validate_collection(request.inputs[name], f"job input {name}")
total_bytes += request.inputs[name].size_bytes
for item in request.inputs[name].items:
existing = artifact_references.get(item.artifact.artifact_id)
if existing is not None and existing != item.artifact:
raise ValueError("job reuses an artifact ID with conflicting metadata")
artifact_references[item.artifact.artifact_id] = item.artifact
if total_bytes > manifest.limits.max_input_bytes:
raise ValueError("job inputs exceed the manifest byte limit")
if len(artifact_references) > manifest.limits.max_artifacts:
raise ValueError("job inputs exceed the manifest artifact limit")
import json
from ._validation import thaw_json
if len(json.dumps(thaw_json(request.parameters), allow_nan=False).encode("utf-8")) > manifest.limits.max_parameter_bytes:
raise ValueError("job parameters exceed the manifest byte limit")
validate_parameter_instance(request.parameters, manifest.parameters_schema)
@staticmethod
def _validate_plan_limits(
request: JobRequest,
plan: WorkflowPlan,
manifest: WorkloadManifest,
) -> None:
references = {
item.artifact.artifact_id: item.artifact
for collection in request.inputs.values()
for item in collection.items
}
for task in plan.tasks:
for collection in task.inputs.values():
for item in collection.items:
existing = references.get(item.artifact.artifact_id)
if existing is not None and existing != item.artifact:
raise ValueError("workflow plan reuses an artifact ID with conflicting metadata")
references[item.artifact.artifact_id] = item.artifact
if len(canonical_json(task.parameters).encode("utf-8")) > manifest.limits.max_parameter_bytes:
raise ValueError("planned task parameters exceed the manifest byte limit")
if len(references) > manifest.limits.max_artifacts:
raise ValueError("workflow plan exceeds the manifest artifact limit")
if len(canonical_json(plan.resolved_parameters).encode("utf-8")) > manifest.limits.max_parameter_bytes:
raise ValueError("resolved parameters exceed the manifest byte limit")
def descriptions(self) -> tuple[WorkloadDescription, ...]:
with self._lock:
result = []
for key, definition in sorted(self._definitions.items()):
digest = definition.manifest.package.digest
result.append(
WorkloadDescription(
definition.manifest.workload,
definition.manifest.description,
digest,
(*key, digest) in self._enabled,
)
)
return tuple(result)
def discover_installed(self, allowlist: tuple[AllowedPackage, ...]) -> None:
"""Load only configured installed entry points; never accept job module paths."""
allowed: dict[tuple[str, str, str], AllowedPackage] = {}
for item in allowlist:
if not isinstance(item, AllowedPackage):
raise ValueError("allowlist must contain AllowedPackage values")
key = (item.distribution, item.workload.name, item.workload.version)
if key in allowed:
raise ValueError("allowlist identities must be unique")
allowed[key] = item
entry_points = metadata.entry_points()
selected = entry_points.select(group=self.ENTRY_POINT_GROUP)
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 ""
)
for key, approval in allowed.items():
if _normalized_distribution_name(key[0]) != distribution:
continue
if entry_point.name != f"{approval.workload.name}@{approval.workload.version}":
continue
_validate_entry_point_ownership(entry_point)
# Import policy is process-global, so installed discovery is
# serialized and intended for application startup. An empty
# cache prefix prevents pre-existing package pyc files from
# being consumed, while dont_write_bytecode keeps the
# measured source tree unchanged during both load and factory.
with _DISCOVERY_IMPORT_LOCK, TemporaryDirectory(
prefix="scimesh-discovery-cache-"
) as cache_prefix:
measured_before = installed_distribution_digest(entry_point.dist)
if measured_before != approval.digest:
raise ValueError(
"installed package content does not match its allowlist digest"
)
previous_bytecode_policy = sys.dont_write_bytecode
previous_cache_prefix = sys.pycache_prefix
sys.dont_write_bytecode = True
sys.pycache_prefix = cache_prefix
try:
loaded = entry_point.load()
definition = (
loaded()
if callable(loaded) and not isinstance(loaded, WorkloadDefinition)
else loaded
)
finally:
sys.pycache_prefix = previous_cache_prefix
sys.dont_write_bytecode = previous_bytecode_policy
if installed_distribution_digest(entry_point.dist) != measured_before:
raise ValueError(
"installed package content changed while loading its entry point"
)
if not isinstance(definition, WorkloadDefinition):
raise ValueError("workload entry point must provide a WorkloadDefinition")
if definition.manifest.workload != approval.workload:
raise ValueError("discovered workload identity does not match its allowlist entry")
if definition.manifest.package.distribution != approval.distribution:
raise ValueError("discovered package identity does not match its allowlist entry")
if definition.manifest.package.digest != approval.digest:
raise ValueError("discovered package digest does not match its allowlist entry")
if key in discovered:
raise ValueError("multiple installed entry points match one allowlist entry")
pending.append(definition)
discovered.add(key)
break
missing = sorted(set(allowed) - discovered)
if missing:
identities = ", ".join(f"{name}@{version}" for _, name, version in missing)
raise ValueError("allowlisted workload entry points were not installed: " + identities)
pending_keys = [
(definition.manifest.workload.name, definition.manifest.workload.version)
for definition in pending
]
if len(pending_keys) != len(set(pending_keys)):
raise ValueError("multiple allowlisted distributions provide one workload version")
with self._lock:
conflicts = [key for key in pending_keys if key in self._definitions]
if conflicts:
name, version = conflicts[0]
raise ValueError(f"workload version already registered: {name}@{version}")
definitions = dict(self._definitions)
enabled = set(self._enabled)
for key, definition in zip(pending_keys, pending):
definitions[key] = definition
enabled.add((*key, definition.manifest.package.digest))
self._definitions = definitions
self._enabled = enabled
+463
View File
@@ -0,0 +1,463 @@
"""Generic resource declarations, runtime inventory, and atomic local allocation."""
from __future__ import annotations
from dataclasses import dataclass
from enum import Enum
from threading import Lock
from types import MappingProxyType
from typing import Mapping
from uuid import uuid4
from ._validation import (
enum_value,
freeze_json_mapping,
require_exact_keys,
require_identifier,
require_nonnegative_int,
require_opaque_resource_id,
require_positive_int,
require_sha256,
require_string,
thaw_json,
)
class AcceleratorMode(str, Enum):
NONE = "none"
EXCLUSIVE_DEVICE = "exclusive_device"
FRACTIONAL = "fractional"
PARTITION = "partition"
def _resource_id(value: object, field: str) -> str:
return require_opaque_resource_id(value, field)
@dataclass(frozen=True, slots=True)
class AcceleratorDevice:
kind: str
vendor: str
device_id: str
model: str
memory_mb: int
modes: tuple[AcceleratorMode, ...]
capabilities: Mapping[str, str]
topology_group: str | None = None
partition_id: str | None = None
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)
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")
object.__setattr__(self, "modes", modes)
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"))
if self.partition_id is not None:
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")
if not isinstance(self.healthy, bool):
raise ValueError("accelerator.healthy must be a boolean")
@property
def allocation_id(self) -> str:
return self.partition_id or self.device_id
def to_dict(self) -> dict[str, object]:
return {
"kind": self.kind,
"vendor": self.vendor,
"device_id": self.device_id,
"model": self.model,
"memory_mb": self.memory_mb,
"modes": [mode.value for mode in self.modes],
"capabilities": thaw_json(self.capabilities),
"topology_group": self.topology_group,
"partition_id": self.partition_id,
"healthy": self.healthy,
}
@classmethod
def from_dict(cls, value: object) -> "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",
}
require_exact_keys(value, fields, "accelerator device")
modes = value["modes"]
if not isinstance(modes, list):
raise ValueError("accelerator modes must be an array")
return cls(
kind=value["kind"], # type: ignore[arg-type]
vendor=value["vendor"], # type: ignore[arg-type]
device_id=value["device_id"], # type: ignore[arg-type]
model=value["model"], # type: ignore[arg-type]
memory_mb=value["memory_mb"], # type: ignore[arg-type]
modes=tuple(modes),
capabilities=value["capabilities"], # type: ignore[arg-type]
topology_group=value["topology_group"], # type: ignore[arg-type]
partition_id=value["partition_id"], # type: ignore[arg-type]
healthy=value["healthy"], # type: ignore[arg-type]
)
@dataclass(frozen=True, slots=True)
class ResourceInventory:
cpu_cores: int
memory_mb: int
scratch_mb: int
architecture: str
accelerators: tuple[AcceleratorDevice, ...] = ()
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"))
devices = tuple(self.accelerators)
if any(not isinstance(device, AcceleratorDevice) for device in devices):
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")
object.__setattr__(self, "accelerators", devices)
digests = tuple(
require_sha256(value, "environment_digest", prefixed=True)
for value in self.environment_digests
)
if len(digests) != len(set(digests)):
raise ValueError("environment_digests must be unique")
object.__setattr__(self, "environment_digests", digests)
def to_dict(self) -> dict[str, object]:
return {
"cpu_cores": self.cpu_cores,
"memory_mb": self.memory_mb,
"scratch_mb": self.scratch_mb,
"architecture": self.architecture,
"accelerators": [device.to_dict() for device in self.accelerators],
"environment_digests": list(self.environment_digests),
}
@classmethod
def from_dict(cls, value: object) -> "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",
}
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")
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),
environment_digests=tuple(digests),
)
@dataclass(frozen=True, slots=True)
class ResourceRequirements:
profile: str
cpu_cores: int
memory_mb: int
scratch_mb: int
accelerator_count: int = 0
accelerator_kind: str | None = None
accelerator_memory_mb: int = 0
accelerator_mode: AcceleratorMode = AcceleratorMode.NONE
architecture: str | None = None
topology_group: str | None = None
environment_digest: str | None = None
estimated_input_bytes: int = 0
estimated_output_bytes: int = 0
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,
"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"),
)
object.__setattr__(
self,
"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.topology_group is not None:
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"))
if self.accelerator_mode is AcceleratorMode.NONE:
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"))
if self.topology_group is not None:
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),
)
object.__setattr__(
self,
"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"),
)
object.__setattr__(
self,
"max_duration_seconds",
require_positive_int(self.max_duration_seconds, "max_duration_seconds"),
)
def eligibility_errors(self, inventory: ResourceInventory) -> tuple[str, ...]:
errors: list[str] = []
if self.cpu_cores > inventory.cpu_cores:
errors.append("insufficient-cpu")
if self.memory_mb > inventory.memory_mb:
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:
errors.append("architecture-mismatch")
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:
errors.append("accelerator-unavailable")
return tuple(errors)
def _matching_devices(
self,
devices: tuple[AcceleratorDevice, ...],
unavailable: set[str] | None = None,
) -> tuple[AcceleratorDevice, ...]:
unavailable = unavailable or set()
if self.accelerator_count == 0:
return ()
matches = [
device
for device in devices
if device.healthy
and device.allocation_id not in unavailable
and device.kind == self.accelerator_kind
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)
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)
]
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]
if sufficiently_large:
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]:
return {
"profile": self.profile,
"cpu_cores": self.cpu_cores,
"memory_mb": self.memory_mb,
"scratch_mb": self.scratch_mb,
"accelerator_count": self.accelerator_count,
"accelerator_kind": self.accelerator_kind,
"accelerator_memory_mb": self.accelerator_memory_mb,
"accelerator_mode": self.accelerator_mode.value,
"architecture": self.architecture,
"topology_group": self.topology_group,
"environment_digest": self.environment_digest,
"estimated_input_bytes": self.estimated_input_bytes,
"estimated_output_bytes": self.estimated_output_bytes,
"max_duration_seconds": self.max_duration_seconds,
}
@classmethod
def from_dict(cls, value: object) -> "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",
}
require_exact_keys(value, fields, "resource requirements")
return cls(**value) # type: ignore[arg-type]
@dataclass(frozen=True, slots=True)
class ResourceAllocation:
allocation_id: str
owner_id: str
cpu_cores: int
memory_mb: int
scratch_mb: int
accelerator_ids: tuple[str, ...]
def __post_init__(self) -> None:
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)
if len(ids) != len(set(ids)):
raise ValueError("accelerator_ids must be unique")
object.__setattr__(self, "accelerator_ids", ids)
@property
def task_key(self) -> str:
"""Compatibility alias; new callers must supply a globally unique attempt owner."""
return self.owner_id
class ResourceUnavailableError(RuntimeError):
"""Raised before execution when a complete atomic reservation is unavailable."""
class ResourcePool:
"""Lock-protected local allocator used by an Agent execution layer.
This object is intentionally coordinator-independent. A protocol-v2 Agent
will bind its returned allocation ID to a coordinator-owned reservation
token; the current protocol must not enable concurrent claims based only on
this local state.
"""
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
self.max_concurrency = require_positive_int(max_concurrency, "max_concurrency")
self._lock = Lock()
self._allocations: dict[str, ResourceAllocation] = {}
self._allocated_devices: dict[str, tuple[AcceleratorDevice, ...]] = {}
@staticmethod
def _devices_conflict(left: AcceleratorDevice, right: AcceleratorDevice) -> bool:
if left.device_id != right.device_id:
return False
if left.partition_id is None or right.partition_id is None:
return True
return left.partition_id == right.partition_id
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 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())
if used_cpu + requirements.cpu_cores > self.inventory.cpu_cores:
raise ResourceUnavailableError("insufficient-cpu")
if used_memory + requirements.memory_mb > self.inventory.memory_mb:
raise ResourceUnavailableError("insufficient-memory")
if used_scratch + requirements.scratch_mb > self.inventory.scratch_mb:
raise ResourceUnavailableError("insufficient-scratch")
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 static_errors:
raise ResourceUnavailableError(static_errors[0])
reserved_devices = tuple(
device
for values in self._allocated_devices.values()
for device in values
)
available_devices = tuple(
device
for device in self.inventory.accelerators
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])
allocation = ResourceAllocation(
allocation_id=str(uuid4()),
owner_id=owner_id,
cpu_cores=requirements.cpu_cores,
memory_mb=requirements.memory_mb,
scratch_mb=requirements.scratch_mb,
accelerator_ids=selected,
)
self._allocations[allocation.allocation_id] = allocation
self._allocated_devices[allocation.allocation_id] = tuple(
devices[: requirements.accelerator_count]
)
return allocation
def release(self, allocation_id: str) -> bool:
allocation_id = _resource_id(allocation_id, "allocation_id")
with self._lock:
removed = self._allocations.pop(allocation_id, None)
self._allocated_devices.pop(allocation_id, None)
return removed is not None
def active_allocations(self) -> tuple[ResourceAllocation, ...]:
with self._lock:
return tuple(sorted(self._allocations.values(), key=lambda item: item.owner_id))
+269
View File
@@ -0,0 +1,269 @@
"""Fail-closed SDK/profile/feature/resource compatibility negotiation."""
from __future__ import annotations
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 .identity import SDK_API_VERSION
from .execution import NetworkPolicy, ProcessModel
from .manifest import TrustMode, WorkloadManifest
from .resources import AcceleratorMode, ResourceInventory
from .workflow import StageKind
class CompatibilityError(ValueError):
def __init__(self, code: str, message: str) -> None:
self.code = require_identifier(code, "compatibility error code")
super().__init__(message)
@dataclass(frozen=True, slots=True)
class RuntimeCapabilities:
sdk_api_version: str
protocol_version: str
profiles: tuple[str, ...]
features: Mapping[str, str]
workload_capabilities: tuple[str, ...]
inventory: ResourceInventory
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"))
# 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)
if len(profiles) != len(set(profiles)):
raise ValueError("runtime profiles must be unique")
object.__setattr__(self, "profiles", profiles)
if not isinstance(self.features, Mapping):
raise ValueError("runtime features must be an object")
features: dict[str, str] = {}
for name, version in self.features.items():
canonical = require_identifier(name, "runtime feature")
text = require_string(version, "runtime feature version", max_length=32)
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)
if len(capabilities) != len(set(capabilities)):
raise ValueError("workload_capabilities must be unique")
object.__setattr__(self, "workload_capabilities", capabilities)
if not isinstance(self.inventory, ResourceInventory):
raise ValueError("runtime inventory must be a ResourceInventory")
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
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)
@dataclass(frozen=True, slots=True)
class NegotiatedWorkload:
manifest: WorkloadManifest
optional_fallbacks: Mapping[str, str]
sdk_api_version: str
protocol_version: str
selected_features: Mapping[str, str]
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,
"sdk_api_version",
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),
)
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(
version,
"negotiated feature version",
max_length=32,
)
validate_version_range(
f"=={selected[name]}",
"negotiated feature version",
)
object.__setattr__(self, "selected_features", MappingProxyType(selected))
def negotiate_manifest(
manifest: WorkloadManifest,
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 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")
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))
if missing_profiles:
raise CompatibilityError(
"profile-unavailable",
"runtime does not support required profiles: " + ", ".join(missing_profiles),
)
if manifest.workload.name not in runtime.workload_capabilities:
raise CompatibilityError(
"workload-unavailable",
"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")
for feature in manifest.required_features:
version = runtime.features.get(feature.name)
if version is None or not feature.versions.contains(version):
raise CompatibilityError(
"feature-unavailable",
f"required feature is unavailable or incompatible: {feature.name}",
)
fallbacks: dict[str, str] = {}
selected_features: dict[str, str] = {}
for feature in manifest.required_features:
version = runtime.features.get(feature.name)
if version is not None and feature.versions.contains(version):
selected_features[feature.name] = version
for feature in manifest.optional_features:
version = runtime.features.get(feature.name)
if version is None or not feature.versions.contains(version):
if feature.fallback is None:
raise CompatibilityError(
"optional-feature-unavailable",
f"optional feature has no declared fallback: {feature.name}",
)
fallbacks[feature.name] = feature.fallback
else:
selected_features[feature.name] = version
required_by_shape: dict[StageKind, str] = {
StageKind.PLAN: "dynamic-expansion",
StageKind.LOOP_CONTROLLER: "bounded-loops",
StageKind.STREAM: "stream-checkpoints",
StageKind.SERVICE: "services",
StageKind.SIDE_EFFECT: "side-effect",
}
declared_required = {feature.name for feature in manifest.required_features}
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}")
for stage in manifest.workflow.stages:
shape_feature = required_by_shape.get(stage.kind)
if shape_feature is not None and shape_feature not in declared_required:
raise CompatibilityError(
"feature-undeclared",
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")
execution = stage.execution
require_declared(
execution.process_model is ProcessModel.PROCESS_POOL,
"process-pools",
f"stage {stage.stage_id} process pool",
)
require_declared(
execution.process_model is ProcessModel.THREAD_POOL,
"thread-pools",
f"stage {stage.stage_id} thread pool",
)
require_declared(
execution.process_model is ProcessModel.EXTERNAL_RUNTIME,
"external-runtimes",
f"stage {stage.stage_id} external runtime",
)
require_declared(
execution.max_processes > 1,
"multi-process",
f"stage {stage.stage_id} multi-process execution",
)
require_declared(
execution.threads_per_process > 1,
"python-threads",
f"stage {stage.stage_id} Python threading",
)
require_declared(
execution.native_threads > 1,
"native-threads",
f"stage {stage.stage_id} native threading",
)
require_declared(
execution.nested_parallelism,
"nested-parallelism",
f"stage {stage.stage_id} nested parallelism",
)
network_features = {
NetworkPolicy.NONE: "network-isolation",
NetworkPolicy.COORDINATOR_ARTIFACTS_ONLY: "artifact-network-policy",
NetworkPolicy.ALLOWLISTED_EGRESS: "egress-allowlist",
}
network_feature = network_features.get(execution.network)
if network_feature is not None:
require_declared(
True,
network_feature,
f"stage {stage.stage_id} network policy",
)
require_declared(
execution.checkpoint.enabled,
"checkpoints",
f"stage {stage.stage_id} checkpoint policy",
)
require_declared(
stage.retry.max_attempts > 1,
"retries",
f"stage {stage.stage_id} retry policy",
)
require_declared(
bool(execution.secret_handles),
"secret-injection",
f"stage {stage.stage_id} secret handles",
)
resource_sets = (stage.resources,) + (
(stage.gang.per_replica_resources,) if stage.gang is not None else ()
)
for resources in resource_sets:
if resources.accelerator_count:
if resources.accelerator_mode is AcceleratorMode.EXCLUSIVE_DEVICE:
feature = "gpu-exclusive"
elif resources.accelerator_mode is AcceleratorMode.PARTITION:
feature = "gpu-mig"
else:
feature = "accelerator-fractional"
if feature not in declared_required:
raise CompatibilityError(
"feature-undeclared",
f"accelerator stage requires declared feature {feature}",
)
errors = resources.eligibility_errors(runtime.inventory)
if errors:
raise CompatibilityError("resource-ineligible", errors[0])
return NegotiatedWorkload(
manifest,
fallbacks,
runtime.sdk_api_version,
runtime.protocol_version,
selected_features,
)
+380
View File
@@ -0,0 +1,380 @@
"""Bounded JSON Schema subset used for SDK v1 public parameters."""
from __future__ import annotations
import math
import re
from fractions import Fraction
from typing import Mapping, Sequence
_ANNOTATIONS = {
"$schema",
"title",
"description",
"default",
"examples",
"deprecated",
"readOnly",
"writeOnly",
}
_KEYWORDS = _ANNOTATIONS | {
"type",
"enum",
"const",
"properties",
"additionalProperties",
"required",
"minProperties",
"maxProperties",
"items",
"minItems",
"maxItems",
"uniqueItems",
"minLength",
"maxLength",
"pattern",
"minimum",
"maximum",
"exclusiveMinimum",
"exclusiveMaximum",
"multipleOf",
"allOf",
"anyOf",
"oneOf",
"not",
}
_TYPES = {"null", "boolean", "object", "array", "number", "integer", "string"}
class ParameterValidationError(ValueError):
"""Sanitized public-parameter schema failure."""
def _schema_error(message: str) -> ValueError:
return ValueError("unsupported or invalid parameters_schema: " + message)
def _json_equal(left: object, right: object) -> bool:
"""Compare values using the JSON data model rather than Python coercion.
Python considers ``True == 1`` while JSON has distinct boolean and number
types. JSON Schema does, however, treat integral and non-integral syntax for
the same mathematical number (for example ``1`` and ``1.0``) as equal.
"""
if isinstance(left, bool) or isinstance(right, bool):
return isinstance(left, bool) and isinstance(right, bool) and left is right
if isinstance(left, (int, float)) and isinstance(right, (int, float)):
return left == right
if left is None or right is None:
return left is None and right is None
if isinstance(left, str) or isinstance(right, str):
return isinstance(left, str) and isinstance(right, str) and left == right
if isinstance(left, Mapping) and isinstance(right, Mapping):
return set(left) == set(right) and all(
_json_equal(left[key], right[key]) for key in left
)
if isinstance(left, (list, tuple)) and isinstance(right, (list, tuple)):
return len(left) == len(right) and all(
_json_equal(left_item, right_item)
for left_item, right_item in zip(left, right)
)
return False
def _json_key(value: object, depth: int = 0) -> object:
"""Build a hashable JSON-type-aware key in linear time."""
if depth > 64:
raise ValueError("JSON value nesting exceeds 64 levels")
if value is None:
return ("null",)
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))
if isinstance(value, str):
return ("string", value)
if isinstance(value, Mapping):
return (
"object",
tuple(
(key, _json_key(child, depth + 1))
for key, child in sorted(value.items())
),
)
if isinstance(value, (list, tuple)):
return ("array", tuple(_json_key(child, depth + 1) for child in value))
raise ValueError("value is not JSON-compatible")
def _validate_safe_pattern(pattern: str) -> None:
"""Accept only the v1 linear-time regex subset.
Groups, alternation, backreferences, and repetition operators are excluded;
literals, anchors, character classes, escapes, and ``.`` remain available.
"""
escaped = False
in_class = False
for character in pattern:
if escaped:
if character.isdigit():
raise _schema_error("pattern backreferences are not supported")
escaped = False
continue
if character == "\\":
escaped = True
continue
if character == "[" and not in_class:
in_class = True
continue
if character == "]" and in_class:
in_class = False
continue
if not in_class and character in "()|*+?{}":
raise _schema_error("pattern uses an unbounded regex operator")
if escaped or in_class:
# ``re.compile`` will provide the canonical invalid-regex error below.
return
def _is_json_multiple(value: int | float, divisor: int | float) -> bool:
"""Evaluate ``multipleOf`` without converting arbitrary integers to float."""
if isinstance(value, int) and isinstance(divisor, int):
return value % divisor == 0
value_fraction = Fraction(value) if isinstance(value, int) else Fraction(str(value))
divisor_fraction = (
Fraction(divisor) if isinstance(divisor, int) else Fraction(str(divisor))
)
return (value_fraction / divisor_fraction).denominator == 1
def validate_schema_definition(schema: Mapping[str, object], *, _depth: int = 0) -> None:
if _depth > 64:
raise _schema_error("nesting exceeds 64 levels")
if not isinstance(schema, Mapping):
raise _schema_error("each schema node must be an object")
unknown = set(schema) - _KEYWORDS
if unknown:
raise _schema_error("unknown keyword " + sorted(unknown)[0])
raw_type = schema.get("type")
if raw_type is not None:
declared = (raw_type,) if isinstance(raw_type, str) else raw_type
if not isinstance(declared, (list, tuple)) or not declared:
raise _schema_error("type must be a string or non-empty array")
if any(not isinstance(value, str) or value not in _TYPES for value in declared):
raise _schema_error("type contains an unsupported JSON type")
if len(declared) != len(set(declared)):
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):
raise _schema_error("properties must be an object")
for child in properties.values():
validate_schema_definition(child, _depth=_depth + 1) # type: ignore[arg-type]
additional = schema.get("additionalProperties")
if additional is not None and not isinstance(additional, (bool, Mapping)):
raise _schema_error("additionalProperties must be a boolean or schema")
if isinstance(additional, Mapping):
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):
raise _schema_error("required must be an array of strings")
if len(required) != len(set(required)):
raise _schema_error("required names must be unique")
for keyword in ("items", "not"):
child = schema.get(keyword)
if child is not None:
validate_schema_definition(child, _depth=_depth + 1) # type: ignore[arg-type]
for keyword in ("allOf", "anyOf", "oneOf"):
children = schema.get(keyword)
if children is None:
continue
if not isinstance(children, (list, tuple)) or not children:
raise _schema_error(f"{keyword} must be a non-empty array")
for child in children:
validate_schema_definition(child, _depth=_depth + 1) # type: ignore[arg-type]
enum = schema.get("enum")
if enum is not None and (not isinstance(enum, (list, tuple)) or not enum):
raise _schema_error("enum must be a non-empty array")
if isinstance(enum, (list, tuple)):
seen_enum: set[object] = set()
for item in enum:
key = _json_key(item)
if key in seen_enum:
raise _schema_error("enum values must be unique")
seen_enum.add(key)
for keyword in (
"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):
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]
raise _schema_error(f"{minimum} must not exceed {maximum}")
for keyword in (
"minimum", "maximum", "exclusiveMinimum", "exclusiveMaximum", "multipleOf"
):
value = schema.get(keyword)
if value is not None and (
isinstance(value, bool)
or not isinstance(value, (int, float))
or (isinstance(value, float) and not math.isfinite(value))
):
raise _schema_error(f"{keyword} must be a finite number")
if "multipleOf" in schema and schema["multipleOf"] <= 0: # type: ignore[operator]
raise _schema_error("multipleOf must be positive")
pattern = schema.get("pattern")
if pattern is not None:
if not isinstance(pattern, str) or len(pattern) > 1024:
raise _schema_error("pattern must be a string of at most 1024 characters")
_validate_safe_pattern(pattern)
try:
re.compile(pattern)
except re.error as error:
raise _schema_error("pattern is not a valid regular expression") from error
for keyword in ("uniqueItems", "deprecated", "readOnly", "writeOnly"):
if keyword in schema and not isinstance(schema[keyword], bool):
raise _schema_error(f"{keyword} must be a boolean")
def _type_matches(value: object, expected: str) -> bool:
if expected == "null":
return value is None
if expected == "boolean":
return isinstance(value, bool)
if expected == "object":
return isinstance(value, Mapping)
if expected == "array":
return isinstance(value, (list, tuple))
if expected == "integer":
return isinstance(value, int) and not isinstance(value, bool)
if expected == "number":
return isinstance(value, (int, float)) and not isinstance(value, bool)
if expected == "string":
return isinstance(value, str)
return False
def _failure(path: str, reason: str) -> ParameterValidationError:
return ParameterValidationError(f"job parameters violate their schema at {path}: {reason}")
def validate_parameter_instance(
value: object,
schema: Mapping[str, object],
*,
path: str = "$",
_depth: int = 0,
) -> None:
if _depth > 64:
raise _failure(path, "nesting exceeds 64 levels")
raw_type = schema.get("type")
if raw_type is not None:
expected = (raw_type,) if isinstance(raw_type, str) else tuple(raw_type) # type: ignore[arg-type]
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]
):
raise _failure(path, "value is outside enum")
if "const" in schema and not _json_equal(value, schema["const"]):
raise _failure(path, "value does not match const")
for keyword in ("allOf", "anyOf", "oneOf"):
children = schema.get(keyword)
if children is None:
continue
matches = 0
for child in children: # type: ignore[union-attr]
try:
validate_parameter_instance(value, child, path=path, _depth=_depth + 1)
except ParameterValidationError:
continue
matches += 1
if keyword == "allOf" and matches != len(children): # type: ignore[arg-type]
raise _failure(path, "allOf did not match")
if keyword == "anyOf" and matches == 0:
raise _failure(path, "anyOf did not match")
if keyword == "oneOf" and matches != 1:
raise _failure(path, "oneOf did not match exactly once")
excluded = schema.get("not")
if excluded is not None:
try:
validate_parameter_instance(value, excluded, path=path, _depth=_depth + 1) # type: ignore[arg-type]
except ParameterValidationError:
pass
else:
raise _failure(path, "value matches a forbidden schema")
if isinstance(value, Mapping):
required = schema.get("required", ())
missing = set(required) - set(value) # type: ignore[arg-type]
if missing:
raise _failure(path, "missing required field " + sorted(missing)[0])
minimum = schema.get("minProperties")
maximum = schema.get("maxProperties")
if minimum is not None and len(value) < minimum: # type: ignore[operator]
raise _failure(path, "too few properties")
if maximum is not None and len(value) > maximum: # type: ignore[operator]
raise _failure(path, "too many properties")
properties = schema.get("properties", {})
additional = schema.get("additionalProperties", True)
for name, child in value.items():
if name in properties: # type: ignore[operator]
validate_parameter_instance(
child,
properties[name], # type: ignore[index]
path=f"{path}.{name}",
_depth=_depth + 1,
)
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)
if isinstance(value, (list, tuple)):
minimum = schema.get("minItems")
maximum = schema.get("maxItems")
if minimum is not None and len(value) < minimum: # type: ignore[operator]
raise _failure(path, "too few items")
if maximum is not None and len(value) > maximum: # type: ignore[operator]
raise _failure(path, "too many items")
if schema.get("uniqueItems"):
seen_items: set[object] = set()
for item in value:
key = _json_key(item)
if key in seen_items:
raise _failure(path, "items must be unique")
seen_items.add(key)
child_schema = schema.get("items")
if child_schema is not None:
for index, item in enumerate(value):
validate_parameter_instance(
item,
child_schema, # type: ignore[arg-type]
path=f"{path}[{index}]",
_depth=_depth + 1,
)
if isinstance(value, str):
if "minLength" in schema and len(value) < schema["minLength"]: # type: ignore[operator]
raise _failure(path, "string is too short")
if "maxLength" in schema and len(value) > schema["maxLength"]: # type: ignore[operator]
raise _failure(path, "string is too long")
if "pattern" in schema and re.search(schema["pattern"], value) is None: # type: ignore[arg-type]
raise _failure(path, "string does not match pattern")
if isinstance(value, (int, float)) and not isinstance(value, bool):
checks = (
("minimum", lambda actual, bound: actual >= bound),
("maximum", lambda actual, bound: actual <= bound),
("exclusiveMinimum", lambda actual, bound: actual > bound),
("exclusiveMaximum", lambda actual, bound: actual < bound),
)
for keyword, predicate in checks:
if keyword in schema and not predicate(value, schema[keyword]):
raise _failure(path, f"number violates {keyword}")
if "multipleOf" in schema:
if not _is_json_multiple(value, schema["multipleOf"]): # type: ignore[arg-type]
raise _failure(path, "number violates multipleOf")
File diff suppressed because it is too large Load Diff
+614
View File
@@ -0,0 +1,614 @@
"""Versioned workflow DAG and bounded advanced-stage declarations."""
from __future__ import annotations
from dataclasses import dataclass
from enum import Enum
from types import MappingProxyType
from typing import Mapping
from ._validation import (
enum_value,
require_entry_point,
require_exact_keys,
require_identifier,
require_nonnegative_int,
require_positive_int,
require_schema_version,
require_string,
)
from .artifacts import PortSpec
from .execution import ExecutionProfile, NetworkPolicy, RetryPolicy
from .identity import ComponentRef, SchemaRef, WORKFLOW_SCHEMA_VERSION
from .resources import ResourceRequirements
class StageKind(str, Enum):
PLAN = "plan"
MAP = "map"
REDUCE = "reduce"
VERIFY = "verify"
LOOP_CONTROLLER = "loop-controller"
STREAM = "stream"
SERVICE = "service"
SIDE_EFFECT = "side-effect"
class WorkflowFailurePolicy(str, Enum):
FAIL_FAST = "fail_fast"
CONTINUE_INDEPENDENT = "continue_independent"
ALLOW_PARTIAL = "allow_partial"
COMPENSATE = "compensate"
@dataclass(frozen=True, slots=True)
class LoopSpec:
state_schema: SchemaRef
max_iterations: int
max_wall_seconds: int
body_workflow: str
continue_when: ComponentRef
checkpoint_every: int
on_limit: str = "fail"
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"))
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"))
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")
def to_dict(self) -> dict[str, object]:
return {
"state_schema": self.state_schema.canonical,
"max_iterations": self.max_iterations,
"max_wall_seconds": self.max_wall_seconds,
"body_workflow": self.body_workflow,
"continue_when": self.continue_when.canonical,
"checkpoint_every": self.checkpoint_every,
"on_limit": self.on_limit,
}
@classmethod
def from_dict(cls, value: object) -> "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",
}
require_exact_keys(value, fields, "loop specification")
return cls(
state_schema=SchemaRef.from_dict(value["state_schema"]),
max_iterations=value["max_iterations"], # type: ignore[arg-type]
max_wall_seconds=value["max_wall_seconds"], # type: ignore[arg-type]
body_workflow=value["body_workflow"], # type: ignore[arg-type]
continue_when=ComponentRef.from_dict(value["continue_when"]),
checkpoint_every=value["checkpoint_every"], # type: ignore[arg-type]
on_limit=value["on_limit"], # type: ignore[arg-type]
)
@dataclass(frozen=True, slots=True)
class StreamSpec:
source: str
partitioning: str
checkpoint_schema: SchemaRef
window_seconds: int
watermark_seconds: int
backpressure_limit: int
delivery_guarantee: str
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"))
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,
"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"))
def to_dict(self) -> dict[str, object]:
return {
"source": self.source,
"partitioning": self.partitioning,
"checkpoint_schema": self.checkpoint_schema.canonical,
"window_seconds": self.window_seconds,
"watermark_seconds": self.watermark_seconds,
"backpressure_limit": self.backpressure_limit,
"delivery_guarantee": self.delivery_guarantee,
"max_windows": self.max_windows,
}
@classmethod
def from_dict(cls, value: object) -> "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",
}
require_exact_keys(value, fields, "stream specification")
return cls(
source=value["source"], # type: ignore[arg-type]
partitioning=value["partitioning"], # type: ignore[arg-type]
checkpoint_schema=SchemaRef.from_dict(value["checkpoint_schema"]),
window_seconds=value["window_seconds"], # type: ignore[arg-type]
watermark_seconds=value["watermark_seconds"], # type: ignore[arg-type]
backpressure_limit=value["backpressure_limit"], # type: ignore[arg-type]
delivery_guarantee=value["delivery_guarantee"], # type: ignore[arg-type]
max_windows=value["max_windows"], # type: ignore[arg-type]
)
@dataclass(frozen=True, slots=True)
class GangSpec:
replicas: int
per_replica_resources: ResourceRequirements
same_topology_group: bool = False
bandwidth_class: str | None = None
failure_mode: str = "fail_all"
def __post_init__(self) -> None:
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):
raise ValueError("gang per_replica_resources must be ResourceRequirements")
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"))
if self.failure_mode != "fail_all":
raise ValueError("SDK v1 gang failure_mode must be fail_all")
def to_dict(self) -> dict[str, object]:
return {
"replicas": self.replicas,
"per_replica_resources": self.per_replica_resources.to_dict(),
"same_topology_group": self.same_topology_group,
"bandwidth_class": self.bandwidth_class,
"failure_mode": self.failure_mode,
}
@classmethod
def from_dict(cls, value: object) -> "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",
}
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"]),
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]
)
@dataclass(frozen=True, slots=True)
class SideEffectSpec:
target: str
idempotency_key_parameter: str
credential_scope: str
compensation: str
manual_approval: bool = True
def __post_init__(self) -> None:
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"),
)
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")
def to_dict(self) -> dict[str, object]:
return {
"target": self.target,
"idempotency_key_parameter": self.idempotency_key_parameter,
"credential_scope": self.credential_scope,
"compensation": self.compensation,
"manual_approval": self.manual_approval,
}
@classmethod
def from_dict(cls, value: object) -> "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",
}
require_exact_keys(value, fields, "side-effect specification")
return cls(**value) # type: ignore[arg-type]
@dataclass(frozen=True, slots=True)
class PortRef:
"""A stage port, or an external workflow input when ``stage_id`` is None."""
port: str
stage_id: str | None = None
def __post_init__(self) -> None:
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"))
def to_dict(self) -> dict[str, object]:
return {"stage_id": self.stage_id, "port": self.port}
@classmethod
def from_dict(cls, value: object) -> "PortRef":
if not isinstance(value, Mapping):
raise ValueError("port reference must be an object")
require_exact_keys(value, {"stage_id", "port"}, "port reference")
return cls(stage_id=value["stage_id"], port=value["port"]) # type: ignore[arg-type]
@dataclass(frozen=True, slots=True)
class ArtifactEdge:
source: PortRef
target: PortRef
def __post_init__(self) -> None:
if not isinstance(self.source, PortRef) or not isinstance(self.target, PortRef):
raise ValueError("artifact edge endpoints must be PortRef values")
if self.target.stage_id is None:
raise ValueError("artifact edge target must be a stage input")
def to_dict(self) -> dict[str, object]:
return {"source": self.source.to_dict(), "target": self.target.to_dict()}
@classmethod
def from_dict(cls, value: object) -> "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"]))
def _port_mapping(value: Mapping[str, PortSpec], field: str) -> Mapping[str, PortSpec]:
if not isinstance(value, Mapping):
raise ValueError(f"{field} must be an object")
ports: dict[str, PortSpec] = {}
for name, port in value.items():
canonical = require_identifier(name, f"{field} port")
if not isinstance(port, PortSpec):
raise ValueError(f"{field} values must be PortSpec values")
ports[canonical] = port
return MappingProxyType(ports)
@dataclass(frozen=True, slots=True)
class StageSpec:
stage_id: str
kind: StageKind
entry_point: str
needs: tuple[str, ...]
inputs: Mapping[str, PortSpec]
outputs: Mapping[str, PortSpec]
parameter_names: tuple[str, ...]
resources: ResourceRequirements
execution: ExecutionProfile
retry: RetryPolicy
verifier: ComponentRef | None = None
trust_modes: tuple[str, ...] = ("trusted",)
max_fan_out: int = 1
cacheable: bool = False
loop: LoopSpec | None = None
stream: StreamSpec | None = None
gang: GangSpec | None = None
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, "kind", enum_value(StageKind, self.kind, "stage.kind"))
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"))
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)
if len(names) != len(set(names)):
raise ValueError("parameter_names must be unique")
object.__setattr__(self, "parameter_names", names)
if not isinstance(self.resources, ResourceRequirements):
raise ValueError("stage.resources must be ResourceRequirements")
if not isinstance(self.execution, ExecutionProfile):
raise ValueError("stage.execution must be ExecutionProfile")
self.execution.validate_resources(self.resources)
if not isinstance(self.retry, RetryPolicy):
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)
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"))
if not isinstance(self.cacheable, bool):
raise ValueError("stage.cacheable must be a boolean")
advanced = {
StageKind.LOOP_CONTROLLER: self.loop,
StageKind.STREAM: self.stream,
StageKind.SIDE_EFFECT: self.side_effect,
}
expected_types = {
StageKind.LOOP_CONTROLLER: LoopSpec,
StageKind.STREAM: StreamSpec,
StageKind.SIDE_EFFECT: SideEffectSpec,
}
for kind, declaration in advanced.items():
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 has the wrong type")
if self.gang is not None and not isinstance(self.gang, GangSpec):
raise ValueError("stage.gang must be a GangSpec")
if self.gang is not None:
self.execution.validate_resources(self.gang.per_replica_resources)
if self.kind is StageKind.SIDE_EFFECT:
raise ValueError("side-effect stages cannot use gang execution")
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}:
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")
def to_dict(self) -> dict[str, object]:
return {
"stage_id": self.stage_id,
"kind": self.kind.value,
"entry_point": self.entry_point,
"needs": list(self.needs),
"inputs": {name: port.to_dict() for name, port in self.inputs.items()},
"outputs": {name: port.to_dict() for name, port in self.outputs.items()},
"parameter_names": list(self.parameter_names),
"resources": self.resources.to_dict(),
"execution": self.execution.to_dict(),
"retry": self.retry.to_dict(),
"verifier": self.verifier.canonical if self.verifier is not None else None,
"trust_modes": list(self.trust_modes),
"max_fan_out": self.max_fan_out,
"cacheable": self.cacheable,
"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,
}
@classmethod
def from_dict(cls, value: object) -> "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",
}
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")
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")
return cls(
stage_id=value["stage_id"], # type: ignore[arg-type]
kind=value["kind"], # type: ignore[arg-type]
entry_point=value["entry_point"], # type: ignore[arg-type]
needs=tuple(value["needs"]), # type: ignore[arg-type]
inputs={name: PortSpec.from_dict(port) for name, port in inputs.items()},
outputs={name: PortSpec.from_dict(port) for name, port in outputs.items()},
parameter_names=tuple(value["parameter_names"]), # type: ignore[arg-type]
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"]),
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"]),
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"]),
)
@dataclass(frozen=True, slots=True)
class WorkflowSpec:
workflow_id: str
inputs: Mapping[str, PortSpec]
stages: tuple[StageSpec, ...]
edges: tuple[ArtifactEdge, ...]
outputs: Mapping[str, PortRef]
failure_policy: WorkflowFailurePolicy = WorkflowFailurePolicy.FAIL_FAST
max_tasks: int = 10_000
max_output_bytes: int = 10 * 1024 * 1024 * 1024
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"))
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")
stage_by_id = {stage.stage_id: stage for stage in stages}
if len(stage_by_id) != len(stages):
raise ValueError("workflow stage IDs must be unique")
object.__setattr__(self, "stages", stages)
edges = tuple(self.edges)
if any(not isinstance(edge, ArtifactEdge) for edge in edges):
raise ValueError("workflow edges must contain ArtifactEdge values")
if len({(edge.source, edge.target) for edge in edges}) != len(edges):
raise ValueError("workflow edges must be unique")
object.__setattr__(self, "edges", edges)
if not isinstance(self.outputs, Mapping) or not self.outputs:
raise ValueError("workflow outputs must be a non-empty object")
outputs: dict[str, PortRef] = {}
for name, reference in self.outputs.items():
canonical = require_identifier(name, "workflow output")
if not isinstance(reference, PortRef) or reference.stage_id is None:
raise ValueError("workflow outputs must reference stage output ports")
outputs[canonical] = reference
object.__setattr__(self, "outputs", MappingProxyType(outputs))
object.__setattr__(
self,
"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_output_bytes",
require_positive_int(self.max_output_bytes, "workflow.max_output_bytes"),
)
self._validate_graph(stage_by_id)
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
try:
stage = stages[reference.stage_id]
return stage.outputs[reference.port]
except KeyError as error:
raise ValueError(
f"unknown source stage output: {reference.stage_id}.{reference.port}"
) from error
def _validate_graph(self, stages: Mapping[str, StageSpec]) -> None:
incoming: dict[tuple[str, str], ArtifactEdge] = {}
dependencies: dict[str, set[str]] = {stage_id: set() for stage_id in stages}
for edge in self.edges:
source_port = self._source_port(edge.source, stages)
assert edge.target.stage_id is not None
try:
target_stage = stages[edge.target.stage_id]
target_port = target_stage.inputs[edge.target.port]
except KeyError as error:
raise ValueError(
f"unknown target stage input: {edge.target.stage_id}.{edge.target.port}"
) from error
target_key = (edge.target.stage_id, edge.target.port)
if target_key in incoming:
raise ValueError("each stage input must have exactly one artifact edge")
incoming[target_key] = edge
same_schema = source_port.schema == target_port.schema
direct_match = source_port == target_port
map_fan_in = (
source_port.cardinality.value == "one"
and target_port.cardinality.value == "many"
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")
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]
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")
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] = []
while ready:
current = ready.pop(0)
visited.append(current)
for name, values in remaining.items():
if current in values:
values.remove(current)
if not values and name not in visited and name not in ready:
ready.append(name)
ready.sort()
if len(visited) != len(stages):
raise ValueError("workflow graph must be acyclic")
for reference in self.outputs.values():
self._source_port(reference, stages)
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()
})
def to_dict(self) -> dict[str, object]:
return {
"schema_version": self.schema_version,
"workflow_id": self.workflow_id,
"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()},
"failure_policy": self.failure_policy.value,
"max_tasks": self.max_tasks,
"max_output_bytes": self.max_output_bytes,
}
@classmethod
def from_dict(cls, value: object) -> "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",
}
require_exact_keys(value, fields, "workflow specification")
inputs, outputs = value["inputs"], value["outputs"]
stages, edges = value["stages"], value["edges"]
if not isinstance(inputs, Mapping) or not isinstance(outputs, Mapping):
raise ValueError("workflow inputs and outputs must be objects")
if not isinstance(stages, list) or not isinstance(edges, list):
raise ValueError("workflow stages and edges must be arrays")
return cls(
schema_version=value["schema_version"], # type: ignore[arg-type]
workflow_id=value["workflow_id"], # type: ignore[arg-type]
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()},
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]
)