Add workload SDK foundation
This commit is contained in:
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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(),),
|
||||
),
|
||||
)
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Adapters for versioned pre-SDK SciMesh workload contracts."""
|
||||
|
||||
from .distributed_v1 import LegacyDistributedWorkloadAdapter
|
||||
|
||||
__all__ = ["LegacyDistributedWorkloadAdapter"]
|
||||
@@ -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
@@ -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)
|
||||
@@ -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]
|
||||
)
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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: ...
|
||||
@@ -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
|
||||
@@ -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))
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
@@ -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]
|
||||
)
|
||||
Reference in New Issue
Block a user