Serve documentation from the operator UI
This commit is contained in:
+290
-64
@@ -24,6 +24,14 @@ from .resources import ResourceRequirements
|
||||
|
||||
|
||||
class StageKind(str, Enum):
|
||||
"""The kind of a workflow stage.
|
||||
|
||||
``PLAN`` stages expand dynamically, ``MAP`` stages fan out, ``REDUCE``
|
||||
stages fan in, and the advanced kinds (loops, streams, services, side
|
||||
effects) are declared but fail negotiation unless the runtime advertises
|
||||
the corresponding features.
|
||||
"""
|
||||
|
||||
PLAN = "plan"
|
||||
MAP = "map"
|
||||
REDUCE = "reduce"
|
||||
@@ -35,6 +43,12 @@ class StageKind(str, Enum):
|
||||
|
||||
|
||||
class WorkflowFailurePolicy(str, Enum):
|
||||
"""How a workflow behaves when a stage fails.
|
||||
|
||||
``FAIL_FAST`` aborts on the first failure; the remaining policies require
|
||||
coordinator/runtime support and are fail-closed in v1.
|
||||
"""
|
||||
|
||||
FAIL_FAST = "fail_fast"
|
||||
CONTINUE_INDEPENDENT = "continue_independent"
|
||||
ALLOW_PARTIAL = "allow_partial"
|
||||
@@ -43,6 +57,11 @@ class WorkflowFailurePolicy(str, Enum):
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LoopSpec:
|
||||
"""Bounded loop declaration for a ``LOOP_CONTROLLER`` stage.
|
||||
|
||||
Declared but not executable until a runtime advertises ``bounded-loops``.
|
||||
"""
|
||||
|
||||
state_schema: SchemaRef
|
||||
max_iterations: int
|
||||
max_wall_seconds: int
|
||||
@@ -54,16 +73,34 @@ class LoopSpec:
|
||||
def __post_init__(self) -> None:
|
||||
if not isinstance(self.state_schema, SchemaRef):
|
||||
raise ValueError("loop state_schema must be a SchemaRef")
|
||||
object.__setattr__(self, "max_iterations", require_positive_int(self.max_iterations, "loop.max_iterations"))
|
||||
object.__setattr__(self, "max_wall_seconds", require_positive_int(self.max_wall_seconds, "loop.max_wall_seconds"))
|
||||
object.__setattr__(self, "body_workflow", require_identifier(self.body_workflow, "loop.body_workflow"))
|
||||
object.__setattr__(
|
||||
self,
|
||||
"max_iterations",
|
||||
require_positive_int(self.max_iterations, "loop.max_iterations"),
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"max_wall_seconds",
|
||||
require_positive_int(self.max_wall_seconds, "loop.max_wall_seconds"),
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"body_workflow",
|
||||
require_identifier(self.body_workflow, "loop.body_workflow"),
|
||||
)
|
||||
if not isinstance(self.continue_when, ComponentRef):
|
||||
raise ValueError("loop continue_when must be a ComponentRef")
|
||||
object.__setattr__(self, "checkpoint_every", require_positive_int(self.checkpoint_every, "loop.checkpoint_every"))
|
||||
object.__setattr__(
|
||||
self,
|
||||
"checkpoint_every",
|
||||
require_positive_int(self.checkpoint_every, "loop.checkpoint_every"),
|
||||
)
|
||||
if self.checkpoint_every > self.max_iterations:
|
||||
raise ValueError("loop checkpoint_every must not exceed max_iterations")
|
||||
if self.on_limit not in {"fail", "accept-best", "return-inconclusive"}:
|
||||
raise ValueError("loop on_limit must be fail, accept-best, or return-inconclusive")
|
||||
raise ValueError(
|
||||
"loop on_limit must be fail, accept-best, or return-inconclusive"
|
||||
)
|
||||
|
||||
def to_dict(self) -> dict[str, object]:
|
||||
return {
|
||||
@@ -81,8 +118,13 @@ class LoopSpec:
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("loop specification must be an object")
|
||||
fields = {
|
||||
"state_schema", "max_iterations", "max_wall_seconds", "body_workflow",
|
||||
"continue_when", "checkpoint_every", "on_limit",
|
||||
"state_schema",
|
||||
"max_iterations",
|
||||
"max_wall_seconds",
|
||||
"body_workflow",
|
||||
"continue_when",
|
||||
"checkpoint_every",
|
||||
"on_limit",
|
||||
}
|
||||
require_exact_keys(value, fields, "loop specification")
|
||||
return cls(
|
||||
@@ -98,6 +140,12 @@ class LoopSpec:
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StreamSpec:
|
||||
"""Bounded stream declaration for a ``STREAM`` stage.
|
||||
|
||||
Declared but not executable until a runtime advertises
|
||||
``stream-checkpoints``.
|
||||
"""
|
||||
|
||||
source: str
|
||||
partitioning: str
|
||||
checkpoint_schema: SchemaRef
|
||||
@@ -108,20 +156,40 @@ class StreamSpec:
|
||||
max_windows: int
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
object.__setattr__(self, "source", require_identifier(self.source, "stream.source"))
|
||||
object.__setattr__(self, "partitioning", require_identifier(self.partitioning, "stream.partitioning"))
|
||||
object.__setattr__(
|
||||
self, "source", require_identifier(self.source, "stream.source")
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"partitioning",
|
||||
require_identifier(self.partitioning, "stream.partitioning"),
|
||||
)
|
||||
if not isinstance(self.checkpoint_schema, SchemaRef):
|
||||
raise ValueError("stream checkpoint_schema must be a SchemaRef")
|
||||
object.__setattr__(self, "window_seconds", require_positive_int(self.window_seconds, "stream.window_seconds"))
|
||||
object.__setattr__(self, "watermark_seconds", require_nonnegative_int(self.watermark_seconds, "stream.watermark_seconds"))
|
||||
object.__setattr__(
|
||||
self,
|
||||
"window_seconds",
|
||||
require_positive_int(self.window_seconds, "stream.window_seconds"),
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"watermark_seconds",
|
||||
require_nonnegative_int(self.watermark_seconds, "stream.watermark_seconds"),
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"backpressure_limit",
|
||||
require_positive_int(self.backpressure_limit, "stream.backpressure_limit"),
|
||||
)
|
||||
if self.delivery_guarantee not in {"at_least_once", "exactly_once"}:
|
||||
raise ValueError("stream delivery_guarantee must be at_least_once or exactly_once")
|
||||
object.__setattr__(self, "max_windows", require_positive_int(self.max_windows, "stream.max_windows"))
|
||||
raise ValueError(
|
||||
"stream delivery_guarantee must be at_least_once or exactly_once"
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"max_windows",
|
||||
require_positive_int(self.max_windows, "stream.max_windows"),
|
||||
)
|
||||
|
||||
def to_dict(self) -> dict[str, object]:
|
||||
return {
|
||||
@@ -140,8 +208,14 @@ class StreamSpec:
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("stream specification must be an object")
|
||||
fields = {
|
||||
"source", "partitioning", "checkpoint_schema", "window_seconds",
|
||||
"watermark_seconds", "backpressure_limit", "delivery_guarantee", "max_windows",
|
||||
"source",
|
||||
"partitioning",
|
||||
"checkpoint_schema",
|
||||
"window_seconds",
|
||||
"watermark_seconds",
|
||||
"backpressure_limit",
|
||||
"delivery_guarantee",
|
||||
"max_windows",
|
||||
}
|
||||
require_exact_keys(value, fields, "stream specification")
|
||||
return cls(
|
||||
@@ -158,6 +232,11 @@ class StreamSpec:
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GangSpec:
|
||||
"""Co-scheduled replica group for one stage.
|
||||
|
||||
Declared but not executable until a runtime advertises ``gang-leases``.
|
||||
"""
|
||||
|
||||
replicas: int
|
||||
per_replica_resources: ResourceRequirements
|
||||
same_topology_group: bool = False
|
||||
@@ -165,7 +244,9 @@ class GangSpec:
|
||||
failure_mode: str = "fail_all"
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
object.__setattr__(self, "replicas", require_positive_int(self.replicas, "gang.replicas"))
|
||||
object.__setattr__(
|
||||
self, "replicas", require_positive_int(self.replicas, "gang.replicas")
|
||||
)
|
||||
if self.replicas < 2:
|
||||
raise ValueError("gang execution requires at least two replicas")
|
||||
if not isinstance(self.per_replica_resources, ResourceRequirements):
|
||||
@@ -173,7 +254,11 @@ class GangSpec:
|
||||
if not isinstance(self.same_topology_group, bool):
|
||||
raise ValueError("gang same_topology_group must be a boolean")
|
||||
if self.bandwidth_class is not None:
|
||||
object.__setattr__(self, "bandwidth_class", require_identifier(self.bandwidth_class, "gang.bandwidth_class"))
|
||||
object.__setattr__(
|
||||
self,
|
||||
"bandwidth_class",
|
||||
require_identifier(self.bandwidth_class, "gang.bandwidth_class"),
|
||||
)
|
||||
if self.failure_mode != "fail_all":
|
||||
raise ValueError("SDK v1 gang failure_mode must be fail_all")
|
||||
|
||||
@@ -191,13 +276,18 @@ class GangSpec:
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("gang specification must be an object")
|
||||
fields = {
|
||||
"replicas", "per_replica_resources", "same_topology_group",
|
||||
"bandwidth_class", "failure_mode",
|
||||
"replicas",
|
||||
"per_replica_resources",
|
||||
"same_topology_group",
|
||||
"bandwidth_class",
|
||||
"failure_mode",
|
||||
}
|
||||
require_exact_keys(value, fields, "gang specification")
|
||||
return cls(
|
||||
replicas=value["replicas"], # type: ignore[arg-type]
|
||||
per_replica_resources=ResourceRequirements.from_dict(value["per_replica_resources"]),
|
||||
per_replica_resources=ResourceRequirements.from_dict(
|
||||
value["per_replica_resources"]
|
||||
),
|
||||
same_topology_group=value["same_topology_group"], # type: ignore[arg-type]
|
||||
bandwidth_class=value["bandwidth_class"], # type: ignore[arg-type]
|
||||
failure_mode=value["failure_mode"], # type: ignore[arg-type]
|
||||
@@ -206,6 +296,12 @@ class GangSpec:
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SideEffectSpec:
|
||||
"""External side-effect declaration for a ``SIDE_EFFECT`` stage.
|
||||
|
||||
Declared but trusted-only and not executable until a runtime advertises
|
||||
``side-effect``; requires an idempotency key projected into the stage.
|
||||
"""
|
||||
|
||||
target: str
|
||||
idempotency_key_parameter: str
|
||||
credential_scope: str
|
||||
@@ -213,14 +309,26 @@ class SideEffectSpec:
|
||||
manual_approval: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
object.__setattr__(self, "target", require_identifier(self.target, "side_effect.target"))
|
||||
object.__setattr__(
|
||||
self, "target", require_identifier(self.target, "side_effect.target")
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"idempotency_key_parameter",
|
||||
require_identifier(self.idempotency_key_parameter, "side_effect.idempotency_key_parameter"),
|
||||
require_identifier(
|
||||
self.idempotency_key_parameter, "side_effect.idempotency_key_parameter"
|
||||
),
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"credential_scope",
|
||||
require_identifier(self.credential_scope, "side_effect.credential_scope"),
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"compensation",
|
||||
require_identifier(self.compensation, "side_effect.compensation"),
|
||||
)
|
||||
object.__setattr__(self, "credential_scope", require_identifier(self.credential_scope, "side_effect.credential_scope"))
|
||||
object.__setattr__(self, "compensation", require_identifier(self.compensation, "side_effect.compensation"))
|
||||
if not isinstance(self.manual_approval, bool):
|
||||
raise ValueError("side_effect.manual_approval must be a boolean")
|
||||
|
||||
@@ -238,7 +346,11 @@ class SideEffectSpec:
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("side-effect specification must be an object")
|
||||
fields = {
|
||||
"target", "idempotency_key_parameter", "credential_scope", "compensation", "manual_approval",
|
||||
"target",
|
||||
"idempotency_key_parameter",
|
||||
"credential_scope",
|
||||
"compensation",
|
||||
"manual_approval",
|
||||
}
|
||||
require_exact_keys(value, fields, "side-effect specification")
|
||||
return cls(**value) # type: ignore[arg-type]
|
||||
@@ -252,9 +364,13 @@ class PortRef:
|
||||
stage_id: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
object.__setattr__(self, "port", require_identifier(self.port, "port reference"))
|
||||
object.__setattr__(
|
||||
self, "port", require_identifier(self.port, "port reference")
|
||||
)
|
||||
if self.stage_id is not None:
|
||||
object.__setattr__(self, "stage_id", require_identifier(self.stage_id, "stage reference"))
|
||||
object.__setattr__(
|
||||
self, "stage_id", require_identifier(self.stage_id, "stage reference")
|
||||
)
|
||||
|
||||
def to_dict(self) -> dict[str, object]:
|
||||
return {"stage_id": self.stage_id, "port": self.port}
|
||||
@@ -269,6 +385,12 @@ class PortRef:
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ArtifactEdge:
|
||||
"""A typed data flow from a source port to a stage input port.
|
||||
|
||||
Endpoints must declare compatible schemas; every stage input receives
|
||||
exactly one edge.
|
||||
"""
|
||||
|
||||
source: PortRef
|
||||
target: PortRef
|
||||
|
||||
@@ -286,7 +408,10 @@ class ArtifactEdge:
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("artifact edge must be an object")
|
||||
require_exact_keys(value, {"source", "target"}, "artifact edge")
|
||||
return cls(source=PortRef.from_dict(value["source"]), target=PortRef.from_dict(value["target"]))
|
||||
return cls(
|
||||
source=PortRef.from_dict(value["source"]),
|
||||
target=PortRef.from_dict(value["target"]),
|
||||
)
|
||||
|
||||
|
||||
def _port_mapping(value: Mapping[str, PortSpec], field: str) -> Mapping[str, PortSpec]:
|
||||
@@ -303,6 +428,13 @@ def _port_mapping(value: Mapping[str, PortSpec], field: str) -> Mapping[str, Por
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StageSpec:
|
||||
"""One typed stage: kind, handler entry point, ports, and policy.
|
||||
|
||||
``entry_point`` must be an installed handler key (runner for map/verify
|
||||
stages, reducer for reduce stages); resources, execution profile, retry
|
||||
policy, verifier, and trust modes are validated at construction.
|
||||
"""
|
||||
|
||||
stage_id: str
|
||||
kind: StageKind
|
||||
entry_point: str
|
||||
@@ -323,18 +455,29 @@ class StageSpec:
|
||||
side_effect: SideEffectSpec | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
object.__setattr__(self, "stage_id", require_identifier(self.stage_id, "stage_id"))
|
||||
object.__setattr__(
|
||||
self, "stage_id", require_identifier(self.stage_id, "stage_id")
|
||||
)
|
||||
object.__setattr__(self, "kind", enum_value(StageKind, self.kind, "stage.kind"))
|
||||
object.__setattr__(self, "entry_point", require_entry_point(self.entry_point, "stage.entry_point"))
|
||||
object.__setattr__(
|
||||
self,
|
||||
"entry_point",
|
||||
require_entry_point(self.entry_point, "stage.entry_point"),
|
||||
)
|
||||
needs = tuple(require_identifier(value, "stage.needs") for value in self.needs)
|
||||
if self.stage_id in needs or len(needs) != len(set(needs)):
|
||||
raise ValueError("stage.needs must contain unique other stage IDs")
|
||||
object.__setattr__(self, "needs", needs)
|
||||
object.__setattr__(self, "inputs", _port_mapping(self.inputs, "stage.inputs"))
|
||||
object.__setattr__(self, "outputs", _port_mapping(self.outputs, "stage.outputs"))
|
||||
object.__setattr__(
|
||||
self, "outputs", _port_mapping(self.outputs, "stage.outputs")
|
||||
)
|
||||
if not self.outputs:
|
||||
raise ValueError("a stage must declare at least one output port")
|
||||
names = tuple(require_identifier(value, "parameter_name") for value in self.parameter_names)
|
||||
names = tuple(
|
||||
require_identifier(value, "parameter_name")
|
||||
for value in self.parameter_names
|
||||
)
|
||||
if len(names) != len(set(names)):
|
||||
raise ValueError("parameter_names must be unique")
|
||||
object.__setattr__(self, "parameter_names", names)
|
||||
@@ -347,13 +490,19 @@ class StageSpec:
|
||||
raise ValueError("stage.retry must be RetryPolicy")
|
||||
if self.verifier is not None and not isinstance(self.verifier, ComponentRef):
|
||||
raise ValueError("stage.verifier must be a ComponentRef")
|
||||
modes = tuple(require_identifier(value, "trust_mode") for value in self.trust_modes)
|
||||
modes = tuple(
|
||||
require_identifier(value, "trust_mode") for value in self.trust_modes
|
||||
)
|
||||
if not modes or len(modes) != len(set(modes)):
|
||||
raise ValueError("stage.trust_modes must be non-empty and unique")
|
||||
if not set(modes).issubset({"trusted", "verified", "untrusted_quorum"}):
|
||||
raise ValueError("stage.trust_modes contains an unsupported trust mode")
|
||||
object.__setattr__(self, "trust_modes", modes)
|
||||
object.__setattr__(self, "max_fan_out", require_positive_int(self.max_fan_out, "stage.max_fan_out"))
|
||||
object.__setattr__(
|
||||
self,
|
||||
"max_fan_out",
|
||||
require_positive_int(self.max_fan_out, "stage.max_fan_out"),
|
||||
)
|
||||
if not isinstance(self.cacheable, bool):
|
||||
raise ValueError("stage.cacheable must be a boolean")
|
||||
advanced = {
|
||||
@@ -370,8 +519,12 @@ class StageSpec:
|
||||
if self.kind is kind and declaration is None:
|
||||
raise ValueError(f"{kind.value} stage requires its bounded declaration")
|
||||
if self.kind is not kind and declaration is not None:
|
||||
raise ValueError(f"{kind.value} declaration is valid only for a {kind.value} stage")
|
||||
if declaration is not None and not isinstance(declaration, expected_types[kind]):
|
||||
raise ValueError(
|
||||
f"{kind.value} declaration is valid only for a {kind.value} stage"
|
||||
)
|
||||
if declaration is not None and not isinstance(
|
||||
declaration, expected_types[kind]
|
||||
):
|
||||
raise ValueError(f"{kind.value} declaration has the wrong type")
|
||||
if self.gang is not None and not isinstance(self.gang, GangSpec):
|
||||
raise ValueError("stage.gang must be a GangSpec")
|
||||
@@ -382,11 +535,16 @@ class StageSpec:
|
||||
if self.kind is StageKind.SIDE_EFFECT:
|
||||
if self.cacheable:
|
||||
raise ValueError("side-effect stages cannot be cached")
|
||||
if self.execution.network not in {NetworkPolicy.ALLOWLISTED_EGRESS, NetworkPolicy.TRUSTED}:
|
||||
if self.execution.network not in {
|
||||
NetworkPolicy.ALLOWLISTED_EGRESS,
|
||||
NetworkPolicy.TRUSTED,
|
||||
}:
|
||||
raise ValueError("side-effect stages require explicit egress")
|
||||
assert self.side_effect is not None
|
||||
if self.side_effect.idempotency_key_parameter not in self.parameter_names:
|
||||
raise ValueError("side-effect idempotency key must be projected into the stage")
|
||||
raise ValueError(
|
||||
"side-effect idempotency key must be projected into the stage"
|
||||
)
|
||||
|
||||
def to_dict(self) -> dict[str, object]:
|
||||
return {
|
||||
@@ -407,7 +565,9 @@ class StageSpec:
|
||||
"loop": self.loop.to_dict() if self.loop is not None else None,
|
||||
"stream": self.stream.to_dict() if self.stream is not None else None,
|
||||
"gang": self.gang.to_dict() if self.gang is not None else None,
|
||||
"side_effect": self.side_effect.to_dict() if self.side_effect is not None else None,
|
||||
"side_effect": self.side_effect.to_dict()
|
||||
if self.side_effect is not None
|
||||
else None,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
@@ -415,14 +575,31 @@ class StageSpec:
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("stage specification must be an object")
|
||||
fields = {
|
||||
"stage_id", "kind", "entry_point", "needs", "inputs", "outputs",
|
||||
"parameter_names", "resources", "execution", "retry", "verifier",
|
||||
"trust_modes", "max_fan_out", "cacheable", "loop", "stream", "gang", "side_effect",
|
||||
"stage_id",
|
||||
"kind",
|
||||
"entry_point",
|
||||
"needs",
|
||||
"inputs",
|
||||
"outputs",
|
||||
"parameter_names",
|
||||
"resources",
|
||||
"execution",
|
||||
"retry",
|
||||
"verifier",
|
||||
"trust_modes",
|
||||
"max_fan_out",
|
||||
"cacheable",
|
||||
"loop",
|
||||
"stream",
|
||||
"gang",
|
||||
"side_effect",
|
||||
}
|
||||
require_exact_keys(value, fields, "stage specification")
|
||||
arrays = (value["needs"], value["parameter_names"], value["trust_modes"])
|
||||
if any(not isinstance(item, list) for item in arrays):
|
||||
raise ValueError("stage needs, parameter_names, and trust_modes must be arrays")
|
||||
raise ValueError(
|
||||
"stage needs, parameter_names, and trust_modes must be arrays"
|
||||
)
|
||||
inputs, outputs = value["inputs"], value["outputs"]
|
||||
if not isinstance(inputs, Mapping) or not isinstance(outputs, Mapping):
|
||||
raise ValueError("stage inputs and outputs must be objects")
|
||||
@@ -437,19 +614,32 @@ class StageSpec:
|
||||
resources=ResourceRequirements.from_dict(value["resources"]),
|
||||
execution=ExecutionProfile.from_dict(value["execution"]),
|
||||
retry=RetryPolicy.from_dict(value["retry"]),
|
||||
verifier=None if value["verifier"] is None else ComponentRef.from_dict(value["verifier"]),
|
||||
verifier=None
|
||||
if value["verifier"] is None
|
||||
else ComponentRef.from_dict(value["verifier"]),
|
||||
trust_modes=tuple(value["trust_modes"]), # type: ignore[arg-type]
|
||||
max_fan_out=value["max_fan_out"], # type: ignore[arg-type]
|
||||
cacheable=value["cacheable"], # type: ignore[arg-type]
|
||||
loop=None if value["loop"] is None else LoopSpec.from_dict(value["loop"]),
|
||||
stream=None if value["stream"] is None else StreamSpec.from_dict(value["stream"]),
|
||||
stream=None
|
||||
if value["stream"] is None
|
||||
else StreamSpec.from_dict(value["stream"]),
|
||||
gang=None if value["gang"] is None else GangSpec.from_dict(value["gang"]),
|
||||
side_effect=None if value["side_effect"] is None else SideEffectSpec.from_dict(value["side_effect"]),
|
||||
side_effect=None
|
||||
if value["side_effect"] is None
|
||||
else SideEffectSpec.from_dict(value["side_effect"]),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class WorkflowSpec:
|
||||
"""A versioned acyclic workflow: inputs, stages, edges, and outputs.
|
||||
|
||||
Construction validates complete input bindings, matching ``needs``
|
||||
declarations, edge schema compatibility, acyclicity, and output port
|
||||
resolution.
|
||||
"""
|
||||
|
||||
workflow_id: str
|
||||
inputs: Mapping[str, PortSpec]
|
||||
stages: tuple[StageSpec, ...]
|
||||
@@ -461,9 +651,15 @@ class WorkflowSpec:
|
||||
schema_version: int = WORKFLOW_SCHEMA_VERSION
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
require_schema_version(self.schema_version, WORKFLOW_SCHEMA_VERSION, "workflow schema_version")
|
||||
object.__setattr__(self, "workflow_id", require_identifier(self.workflow_id, "workflow_id"))
|
||||
object.__setattr__(self, "inputs", _port_mapping(self.inputs, "workflow.inputs"))
|
||||
require_schema_version(
|
||||
self.schema_version, WORKFLOW_SCHEMA_VERSION, "workflow schema_version"
|
||||
)
|
||||
object.__setattr__(
|
||||
self, "workflow_id", require_identifier(self.workflow_id, "workflow_id")
|
||||
)
|
||||
object.__setattr__(
|
||||
self, "inputs", _port_mapping(self.inputs, "workflow.inputs")
|
||||
)
|
||||
stages = tuple(self.stages)
|
||||
if not stages or any(not isinstance(stage, StageSpec) for stage in stages):
|
||||
raise ValueError("workflow stages must contain at least one StageSpec")
|
||||
@@ -489,9 +685,15 @@ class WorkflowSpec:
|
||||
object.__setattr__(
|
||||
self,
|
||||
"failure_policy",
|
||||
enum_value(WorkflowFailurePolicy, self.failure_policy, "workflow.failure_policy"),
|
||||
enum_value(
|
||||
WorkflowFailurePolicy, self.failure_policy, "workflow.failure_policy"
|
||||
),
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"max_tasks",
|
||||
require_positive_int(self.max_tasks, "workflow.max_tasks"),
|
||||
)
|
||||
object.__setattr__(self, "max_tasks", require_positive_int(self.max_tasks, "workflow.max_tasks"))
|
||||
object.__setattr__(
|
||||
self,
|
||||
"max_output_bytes",
|
||||
@@ -499,12 +701,16 @@ class WorkflowSpec:
|
||||
)
|
||||
self._validate_graph(stage_by_id)
|
||||
|
||||
def _source_port(self, reference: PortRef, stages: Mapping[str, StageSpec]) -> PortSpec:
|
||||
def _source_port(
|
||||
self, reference: PortRef, stages: Mapping[str, StageSpec]
|
||||
) -> PortSpec:
|
||||
if reference.stage_id is None:
|
||||
try:
|
||||
return self.inputs[reference.port]
|
||||
except KeyError as error:
|
||||
raise ValueError(f"unknown workflow input port: {reference.port}") from error
|
||||
raise ValueError(
|
||||
f"unknown workflow input port: {reference.port}"
|
||||
) from error
|
||||
try:
|
||||
stage = stages[reference.stage_id]
|
||||
return stage.outputs[reference.port]
|
||||
@@ -538,17 +744,23 @@ class WorkflowSpec:
|
||||
and target_port.collection.value in {"ordered", "keyed", "set"}
|
||||
)
|
||||
if not same_schema or not (direct_match or map_fan_in):
|
||||
raise ValueError("artifact edge source and target port declarations are incompatible")
|
||||
raise ValueError(
|
||||
"artifact edge source and target port declarations are incompatible"
|
||||
)
|
||||
if edge.source.stage_id is not None:
|
||||
dependencies[edge.target.stage_id].add(edge.source.stage_id)
|
||||
for stage in stages.values():
|
||||
missing = [name for name in stage.inputs if (stage.stage_id, name) not in incoming]
|
||||
missing = [
|
||||
name for name in stage.inputs if (stage.stage_id, name) not in incoming
|
||||
]
|
||||
if missing:
|
||||
raise ValueError(
|
||||
f"stage {stage.stage_id} has unbound inputs: {', '.join(sorted(missing))}"
|
||||
)
|
||||
if dependencies[stage.stage_id] != set(stage.needs):
|
||||
raise ValueError(f"stage {stage.stage_id} needs do not match its artifact edges")
|
||||
raise ValueError(
|
||||
f"stage {stage.stage_id} needs do not match its artifact edges"
|
||||
)
|
||||
remaining = {name: set(values) for name, values in dependencies.items()}
|
||||
ready = sorted(name for name, values in remaining.items() if not values)
|
||||
visited: list[str] = []
|
||||
@@ -568,10 +780,12 @@ class WorkflowSpec:
|
||||
|
||||
def output_ports(self) -> Mapping[str, PortSpec]:
|
||||
stages = {stage.stage_id: stage for stage in self.stages}
|
||||
return MappingProxyType({
|
||||
name: self._source_port(reference, stages)
|
||||
for name, reference in self.outputs.items()
|
||||
})
|
||||
return MappingProxyType(
|
||||
{
|
||||
name: self._source_port(reference, stages)
|
||||
for name, reference in self.outputs.items()
|
||||
}
|
||||
)
|
||||
|
||||
def to_dict(self) -> dict[str, object]:
|
||||
return {
|
||||
@@ -580,7 +794,9 @@ class WorkflowSpec:
|
||||
"inputs": {name: port.to_dict() for name, port in self.inputs.items()},
|
||||
"stages": [stage.to_dict() for stage in self.stages],
|
||||
"edges": [edge.to_dict() for edge in self.edges],
|
||||
"outputs": {name: reference.to_dict() for name, reference in self.outputs.items()},
|
||||
"outputs": {
|
||||
name: reference.to_dict() for name, reference in self.outputs.items()
|
||||
},
|
||||
"failure_policy": self.failure_policy.value,
|
||||
"max_tasks": self.max_tasks,
|
||||
"max_output_bytes": self.max_output_bytes,
|
||||
@@ -591,8 +807,15 @@ class WorkflowSpec:
|
||||
if not isinstance(value, Mapping):
|
||||
raise ValueError("workflow specification must be an object")
|
||||
fields = {
|
||||
"schema_version", "workflow_id", "inputs", "stages", "edges",
|
||||
"outputs", "failure_policy", "max_tasks", "max_output_bytes",
|
||||
"schema_version",
|
||||
"workflow_id",
|
||||
"inputs",
|
||||
"stages",
|
||||
"edges",
|
||||
"outputs",
|
||||
"failure_policy",
|
||||
"max_tasks",
|
||||
"max_output_bytes",
|
||||
}
|
||||
require_exact_keys(value, fields, "workflow specification")
|
||||
inputs, outputs = value["inputs"], value["outputs"]
|
||||
@@ -607,7 +830,10 @@ class WorkflowSpec:
|
||||
inputs={name: PortSpec.from_dict(port) for name, port in inputs.items()},
|
||||
stages=tuple(StageSpec.from_dict(stage) for stage in stages),
|
||||
edges=tuple(ArtifactEdge.from_dict(edge) for edge in edges),
|
||||
outputs={name: PortRef.from_dict(reference) for name, reference in outputs.items()},
|
||||
outputs={
|
||||
name: PortRef.from_dict(reference)
|
||||
for name, reference in outputs.items()
|
||||
},
|
||||
failure_policy=value["failure_policy"], # type: ignore[arg-type]
|
||||
max_tasks=value["max_tasks"], # type: ignore[arg-type]
|
||||
max_output_bytes=value["max_output_bytes"], # type: ignore[arg-type]
|
||||
|
||||
Reference in New Issue
Block a user