feat: add Python deployment and run objects
This commit is contained in:
@@ -14,11 +14,14 @@ from .codec import (
|
||||
decode_capability_inspect,
|
||||
decode_dependency_diagnostics,
|
||||
decode_deployment,
|
||||
decode_deployment_validation,
|
||||
decode_deployments,
|
||||
decode_run_result,
|
||||
decode_trace_result,
|
||||
decode_validate_artifact_plan,
|
||||
decode_workflow_artifact,
|
||||
)
|
||||
from .deployments import Deployment, DeploymentValidation
|
||||
from .errors import (
|
||||
ArtifactNotFound,
|
||||
ArtifactVersionConflict,
|
||||
@@ -33,6 +36,7 @@ from .errors import (
|
||||
WorkflowClientError,
|
||||
)
|
||||
from .protocols import WorkflowClientPort
|
||||
from .runs import Run, TracePage
|
||||
from .workflows import (
|
||||
ArtifactRef,
|
||||
Diagnostic,
|
||||
@@ -55,10 +59,13 @@ __all__ = [
|
||||
"Diagnostic",
|
||||
"DeploymentNotRunnable",
|
||||
"DeploymentRequired",
|
||||
"Deployment",
|
||||
"DeploymentValidation",
|
||||
"InvalidResponse",
|
||||
"Page",
|
||||
"ProtocolError",
|
||||
"RemoteCapability",
|
||||
"Run",
|
||||
"RevisionConflict",
|
||||
"TransportError",
|
||||
"ValidationFailed",
|
||||
@@ -68,11 +75,14 @@ __all__ = [
|
||||
"WorkflowArtifact",
|
||||
"WorkflowDiagnostic",
|
||||
"WorkflowValidation",
|
||||
"TracePage",
|
||||
"decode_capabilities_page",
|
||||
"decode_capability_call",
|
||||
"decode_capability_diagnostics",
|
||||
"decode_capability_inspect",
|
||||
"decode_dependency_diagnostics",
|
||||
"decode_deployment_validation",
|
||||
"decode_deployments",
|
||||
"decode_deployment",
|
||||
"decode_run_result",
|
||||
"decode_trace_result",
|
||||
|
||||
+32
-1
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from wf_platform import CapabilityRef, Page, SourceRef
|
||||
from wf_transport_rpc_http import RpcWorkflowApiClient
|
||||
@@ -20,6 +20,10 @@ from .errors import InvalidResponse
|
||||
from .protocols import WorkflowClientPort
|
||||
from .workflows import WorkflowArtifact
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .deployments import Deployment
|
||||
from .runs import Run
|
||||
|
||||
|
||||
def _capability_ref(qualified_name: str, source_id: str) -> CapabilityRef:
|
||||
"""Build a ref by removing the exact source prefix, preserving dotted keys."""
|
||||
@@ -176,3 +180,30 @@ class App:
|
||||
) -> EditableWorkflow:
|
||||
"""Inspect an exact artifact version and seed an editable builder."""
|
||||
return (await self.workflow(artifact_id, version=version)).edit()
|
||||
|
||||
async def deployment(self, deployment_id: str) -> Deployment:
|
||||
"""Inspect and reconstruct one immutable deployment snapshot."""
|
||||
from .deployments import Deployment
|
||||
|
||||
deployment = Deployment.from_payload(
|
||||
self._port,
|
||||
await self._port.inspect_deployment(deployment_id=deployment_id),
|
||||
)
|
||||
if deployment.deployment_id != deployment_id:
|
||||
raise InvalidResponse(
|
||||
operation="workflow.deployments.inspect",
|
||||
details=(
|
||||
f"inspected deployment {deployment.deployment_id!r} does not "
|
||||
f"match requested {deployment_id!r}"
|
||||
),
|
||||
)
|
||||
return deployment
|
||||
|
||||
async def run(self, run_id: str) -> Run:
|
||||
"""Inspect and reconstruct one immutable durable run snapshot."""
|
||||
from .runs import Run
|
||||
|
||||
return Run.from_payload(
|
||||
self._port,
|
||||
await self._port.inspect_run(run_id=run_id),
|
||||
)
|
||||
|
||||
@@ -13,10 +13,12 @@ from wf_api.models import (
|
||||
DependencyDiagnosticPayload,
|
||||
InspectCapabilityResult,
|
||||
ListCapabilitiesResult,
|
||||
ListDeploymentsResult,
|
||||
RawWorkflowPlan,
|
||||
RunResult,
|
||||
RunTraceResult,
|
||||
ValidateArtifactPlanResult,
|
||||
ValidateDeploymentResult,
|
||||
WorkflowArtifactPayload,
|
||||
WorkflowDeploymentPayload,
|
||||
)
|
||||
@@ -157,6 +159,20 @@ def decode_deployment(payload: object) -> WorkflowDeployment:
|
||||
return _model_validate(WorkflowDeployment, wire, operation)
|
||||
|
||||
|
||||
def decode_deployments(payload: object) -> ListDeploymentsResult:
|
||||
"""Validate compact deployment discovery rows at the client boundary."""
|
||||
return _validate(payload, ListDeploymentsResult, "workflow.deployments.list")
|
||||
|
||||
|
||||
def decode_deployment_validation(payload: object) -> ValidateDeploymentResult:
|
||||
"""Validate one deployment readiness response."""
|
||||
return _validate(
|
||||
payload,
|
||||
ValidateDeploymentResult,
|
||||
"workflow.deployments.validate",
|
||||
)
|
||||
|
||||
|
||||
def decode_dependency_diagnostics(
|
||||
payload: object,
|
||||
) -> tuple[DependencyDiagnostic, ...]:
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
"""Immutable deployment snapshots and strict artifact deployment selection."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from wf_artifacts import DependencyDiagnostic, DriftPolicy, WorkflowDeployment
|
||||
|
||||
from .codec import (
|
||||
decode_dependency_diagnostics,
|
||||
decode_deployment,
|
||||
decode_deployment_validation,
|
||||
decode_deployments,
|
||||
)
|
||||
from .errors import DeploymentNotRunnable, DeploymentRequired, InvalidResponse
|
||||
from .protocols import WorkflowClientPort
|
||||
from .runs import Run
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DeploymentValidation:
|
||||
"""Server readiness result for one saved deployment snapshot."""
|
||||
|
||||
deployment_id: str
|
||||
artifact_id: str
|
||||
artifact_version: int
|
||||
status: str
|
||||
diagnostics: tuple[DependencyDiagnostic, ...]
|
||||
|
||||
@property
|
||||
def runnable(self) -> bool:
|
||||
return self.status == "runnable"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Deployment:
|
||||
"""Immutable snapshot of one configured artifact deployment."""
|
||||
|
||||
_port: WorkflowClientPort = field(repr=False, compare=False)
|
||||
model: WorkflowDeployment
|
||||
diagnostics: tuple[DependencyDiagnostic, ...] = ()
|
||||
runnable: bool | None = None
|
||||
|
||||
@classmethod
|
||||
def from_payload(cls, port: WorkflowClientPort, payload: object) -> Deployment:
|
||||
return cls(_port=port, model=decode_deployment(payload))
|
||||
|
||||
@property
|
||||
def deployment_id(self) -> str:
|
||||
return self.model.id
|
||||
|
||||
@property
|
||||
def artifact_id(self) -> str:
|
||||
return self.model.artifact_id
|
||||
|
||||
@property
|
||||
def artifact_version(self) -> int:
|
||||
return self.model.artifact_version
|
||||
|
||||
@property
|
||||
def bindings(self) -> dict[str, str]:
|
||||
return self.model.binding_map()
|
||||
|
||||
@property
|
||||
def drift_policy(self) -> DriftPolicy:
|
||||
return self.model.drift_policy
|
||||
|
||||
async def validate(self) -> DeploymentValidation:
|
||||
payload = await self._port.validate_deployment(
|
||||
deployment_id=self.deployment_id,
|
||||
)
|
||||
result = decode_deployment_validation(payload)
|
||||
if result["deployment_id"] != self.deployment_id:
|
||||
raise InvalidResponse(
|
||||
operation="workflow.deployments.validate",
|
||||
details=(
|
||||
f"validated deployment {result['deployment_id']!r} does not "
|
||||
f"match requested {self.deployment_id!r}"
|
||||
),
|
||||
)
|
||||
diagnostics = decode_dependency_diagnostics(result["diagnostics"])
|
||||
return DeploymentValidation(
|
||||
deployment_id=result["deployment_id"],
|
||||
artifact_id=result["artifact_id"],
|
||||
artifact_version=result["artifact_version"],
|
||||
status=result["status"],
|
||||
diagnostics=diagnostics,
|
||||
)
|
||||
|
||||
async def run(self, workflow_input: Mapping[str, Any]) -> Run:
|
||||
from .codec import decode_run_result
|
||||
from .runs import _run_from_decoded
|
||||
|
||||
decoded = decode_run_result(
|
||||
await self._port.run_deployment(
|
||||
deployment_id=self.deployment_id,
|
||||
workflow_input=dict(workflow_input),
|
||||
)
|
||||
)
|
||||
if decoded.run_id is None or decoded.status in {"unrunnable", "rejected"}:
|
||||
raise DeploymentNotRunnable(
|
||||
deployment_id=self.deployment_id,
|
||||
diagnostics=decoded.diagnostics,
|
||||
outcome=decoded.outcome,
|
||||
error=decoded.error,
|
||||
)
|
||||
return _run_from_decoded(self._port, decoded)
|
||||
|
||||
|
||||
def _decode_summaries(payload: object) -> list[dict[str, Any]]:
|
||||
"""Validate list metadata while keeping summaries as local dictionaries."""
|
||||
result = decode_deployments(payload)
|
||||
return [dict(item) for item in result["deployments"]]
|
||||
|
||||
|
||||
async def run_artifact(
|
||||
artifact: Any,
|
||||
workflow_input: Mapping[str, Any],
|
||||
*,
|
||||
deployment_id: str | None,
|
||||
bindings: Mapping[str, str] | None,
|
||||
drift_policy: DriftPolicy | str,
|
||||
) -> Run:
|
||||
"""Apply the artifact's strict deployment-selection policy."""
|
||||
if deployment_id is not None:
|
||||
deployment = Deployment.from_payload(
|
||||
artifact._port,
|
||||
await artifact._port.inspect_deployment(deployment_id=deployment_id),
|
||||
)
|
||||
if deployment.deployment_id != deployment_id:
|
||||
raise InvalidResponse(
|
||||
operation="workflow.deployments.inspect",
|
||||
details=(
|
||||
f"inspected deployment {deployment.deployment_id!r} does not "
|
||||
f"match requested {deployment_id!r}"
|
||||
),
|
||||
)
|
||||
return await deployment.run(workflow_input)
|
||||
|
||||
summaries = _decode_summaries(await artifact._port.list_deployments())
|
||||
matches = sorted(
|
||||
(
|
||||
summary
|
||||
for summary in summaries
|
||||
if summary["artifact_id"] == artifact.artifact.id
|
||||
and summary["artifact_version"] == artifact.artifact.version
|
||||
),
|
||||
key=lambda summary: summary["id"],
|
||||
)
|
||||
if len(matches) > 1:
|
||||
raise DeploymentRequired(
|
||||
candidate_deployment_ids=tuple(summary["id"] for summary in matches)
|
||||
)
|
||||
default_id = f"{artifact.artifact.id}.v{artifact.artifact.version}.default"
|
||||
if not matches:
|
||||
conflicting = next(
|
||||
(
|
||||
summary
|
||||
for summary in summaries
|
||||
if summary["id"] == default_id
|
||||
and (
|
||||
summary["artifact_id"] != artifact.artifact.id
|
||||
or summary["artifact_version"] != artifact.artifact.version
|
||||
)
|
||||
),
|
||||
None,
|
||||
)
|
||||
if conflicting is not None:
|
||||
raise DeploymentRequired(candidate_deployment_ids=(default_id,))
|
||||
deployment = await artifact.deploy(
|
||||
default_id,
|
||||
bindings=bindings,
|
||||
drift_policy=drift_policy,
|
||||
)
|
||||
if deployment.runnable is not True:
|
||||
raise DeploymentRequired(diagnostics=deployment.diagnostics)
|
||||
return await deployment.run(workflow_input)
|
||||
|
||||
deployment = Deployment.from_payload(
|
||||
artifact._port,
|
||||
await artifact._port.inspect_deployment(deployment_id=matches[0]["id"]),
|
||||
)
|
||||
if deployment.deployment_id != matches[0]["id"]:
|
||||
raise InvalidResponse(
|
||||
operation="workflow.deployments.inspect",
|
||||
details=(
|
||||
f"inspected deployment {deployment.deployment_id!r} does not "
|
||||
f"match requested {matches[0]['id']!r}"
|
||||
),
|
||||
)
|
||||
return await deployment.run(workflow_input)
|
||||
+46
-2
@@ -4,6 +4,8 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from wf_artifacts import DependencyDiagnostic
|
||||
|
||||
|
||||
class WorkflowClientError(Exception):
|
||||
"""Base class for errors that can be handled by workflow callers."""
|
||||
@@ -43,12 +45,54 @@ class ArtifactVersionConflict(WorkflowClientError):
|
||||
"""An artifact version conflicts with an existing saved version."""
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class DeploymentRequired(WorkflowClientError):
|
||||
"""An operation requires a deployment selection."""
|
||||
"""An operation requires an unambiguous or repairable deployment."""
|
||||
|
||||
candidate_deployment_ids: tuple[str, ...] = ()
|
||||
diagnostics: tuple[DependencyDiagnostic, ...] = ()
|
||||
unresolved_logical_sources: tuple[str, ...] = ()
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
object.__setattr__(self, "args", (str(self),))
|
||||
|
||||
def __str__(self) -> str:
|
||||
parts: list[str] = []
|
||||
if self.candidate_deployment_ids:
|
||||
parts.append(
|
||||
"candidate deployments: " + ", ".join(self.candidate_deployment_ids)
|
||||
)
|
||||
if self.unresolved_logical_sources:
|
||||
parts.append(
|
||||
"unresolved sources: " + ", ".join(self.unresolved_logical_sources)
|
||||
)
|
||||
if self.diagnostics:
|
||||
parts.append(
|
||||
"diagnostics: "
|
||||
+ "; ".join(diagnostic.message for diagnostic in self.diagnostics)
|
||||
)
|
||||
return "deployment required" + (f" ({'; '.join(parts)})" if parts else ".")
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class DeploymentNotRunnable(WorkflowClientError):
|
||||
"""A selected deployment failed its readiness checks."""
|
||||
"""A selected deployment failed its readiness checks or returned no run."""
|
||||
|
||||
deployment_id: str = ""
|
||||
diagnostics: tuple[DependencyDiagnostic, ...] = ()
|
||||
outcome: str | None = None
|
||||
error: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
object.__setattr__(self, "args", (str(self),))
|
||||
|
||||
def __str__(self) -> str:
|
||||
detail = self.error or self.outcome or "deployment is not runnable"
|
||||
if self.diagnostics:
|
||||
detail += ": " + "; ".join(
|
||||
diagnostic.message for diagnostic in self.diagnostics
|
||||
)
|
||||
return f"deployment {self.deployment_id!r} not runnable: {detail}"
|
||||
|
||||
|
||||
class ValidationFailed(WorkflowClientError):
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
"""Immutable snapshots for durable workflow runs and bounded traces."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from wf_api import TraceRange
|
||||
from wf_artifacts import DependencyDiagnostic
|
||||
from wf_core import InterruptRequest, InterruptRoute, TraceEntry, WorkflowRef
|
||||
|
||||
from .codec import DecodedRunResult, decode_run_result, decode_trace_result
|
||||
from .errors import DeploymentNotRunnable
|
||||
from .protocols import WorkflowClientPort
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TracePage:
|
||||
"""One bounded, already-loaded slice of a durable run's execution trace."""
|
||||
|
||||
start: int
|
||||
limit: int
|
||||
frames: tuple[TraceEntry, ...]
|
||||
truncated: bool
|
||||
trace_count: int
|
||||
|
||||
|
||||
def _interrupt(payload: Mapping[str, Any] | None) -> InterruptRequest | None:
|
||||
if payload is None:
|
||||
return None
|
||||
data = dict(payload)
|
||||
route_data = data.get("route")
|
||||
route = None
|
||||
if route_data is not None:
|
||||
route_values = dict(route_data)
|
||||
workflow_ref = route_values.get("workflow_ref")
|
||||
route = InterruptRoute(
|
||||
frame_id=route_values["frame_id"],
|
||||
node_id=route_values["node_id"],
|
||||
scope_id=route_values["scope_id"],
|
||||
lineage_id=route_values["lineage_id"],
|
||||
parent_frame_id=route_values["parent_frame_id"],
|
||||
workflow_ref=WorkflowRef.model_validate(workflow_ref),
|
||||
)
|
||||
data["route"] = route
|
||||
# ``InterruptPayload`` is intentionally consumed at this boundary; public
|
||||
# clients receive the core runtime request instead of a wire TypedDict.
|
||||
return InterruptRequest(**data)
|
||||
|
||||
|
||||
def _run_from_decoded(port: WorkflowClientPort, decoded: DecodedRunResult) -> Run:
|
||||
if decoded.run_id is None:
|
||||
raise DeploymentNotRunnable(
|
||||
deployment_id=decoded.deployment_id,
|
||||
diagnostics=decoded.diagnostics,
|
||||
outcome=decoded.outcome,
|
||||
)
|
||||
return Run(
|
||||
_port=port,
|
||||
run_id=decoded.run_id,
|
||||
deployment_id=decoded.deployment_id,
|
||||
status=decoded.status,
|
||||
outcome=decoded.outcome,
|
||||
output=decoded.output,
|
||||
interrupt=_interrupt(decoded.interrupt),
|
||||
diagnostics=decoded.diagnostics,
|
||||
trace_count=decoded.trace_count,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Run:
|
||||
"""Immutable client snapshot of one durable deployment run."""
|
||||
|
||||
_port: WorkflowClientPort = field(repr=False, compare=False)
|
||||
run_id: str
|
||||
deployment_id: str
|
||||
status: str
|
||||
outcome: str | None
|
||||
output: dict[str, Any] | None
|
||||
interrupt: InterruptRequest | None
|
||||
diagnostics: tuple[DependencyDiagnostic, ...]
|
||||
trace_count: int
|
||||
|
||||
@classmethod
|
||||
def from_payload(cls, port: WorkflowClientPort, payload: object) -> Run:
|
||||
"""Validate one run response and reconstruct its immutable snapshot."""
|
||||
return _run_from_decoded(port, decode_run_result(payload))
|
||||
|
||||
async def refresh(self) -> Run:
|
||||
"""Read the current server snapshot without mutating this run."""
|
||||
return self.from_payload(
|
||||
self._port,
|
||||
await self._port.inspect_run(run_id=self.run_id),
|
||||
)
|
||||
|
||||
async def resume(
|
||||
self,
|
||||
response: Mapping[str, Any],
|
||||
*,
|
||||
outcome: str = "submitted",
|
||||
) -> Run:
|
||||
"""Resume an interrupted run and return the server's new snapshot."""
|
||||
if self.status != "interrupted" or self.interrupt is None:
|
||||
raise ValueError("only interrupted runs can be resumed")
|
||||
if not self.interrupt.resumable:
|
||||
raise ValueError("run interrupt is not resumable")
|
||||
return self.from_payload(
|
||||
self._port,
|
||||
await self._port.resume_run(
|
||||
run_id=self.run_id,
|
||||
resume_payload=dict(response),
|
||||
resume_outcome=outcome,
|
||||
),
|
||||
)
|
||||
|
||||
async def trace(self, *, start: int = 0, limit: int = 25) -> TracePage:
|
||||
"""Read a bounded trace page, validating bounds before remote I/O."""
|
||||
if start < 0:
|
||||
raise ValueError("start must be >= 0")
|
||||
if limit <= 0 or limit > 100:
|
||||
raise ValueError("limit must be between 1 and 100")
|
||||
decoded = decode_trace_result(
|
||||
await self._port.read_run_trace(
|
||||
run_id=self.run_id,
|
||||
trace_range=TraceRange(start=start, limit=limit),
|
||||
)
|
||||
)
|
||||
frames = tuple(TraceEntry(**dict(frame)) for frame in (decoded.trace or ()))
|
||||
return TracePage(
|
||||
start=decoded.trace_start if decoded.trace_start is not None else start,
|
||||
limit=decoded.trace_limit if decoded.trace_limit is not None else limit,
|
||||
frames=frames,
|
||||
truncated=bool(decoded.trace_truncated),
|
||||
trace_count=decoded.trace_count,
|
||||
)
|
||||
@@ -2,8 +2,9 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Literal
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field, replace
|
||||
from typing import TYPE_CHECKING, Any, Literal
|
||||
|
||||
from wf_artifacts.models import (
|
||||
RequiredCapability,
|
||||
@@ -13,11 +14,13 @@ from wf_artifacts.models import (
|
||||
)
|
||||
from wf_core import ValidationReport, Workflow
|
||||
|
||||
from .errors import ValidationFailed
|
||||
from .errors import InvalidResponse, ValidationFailed
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .authoring import EditableWorkflow
|
||||
from .deployments import Deployment
|
||||
from .protocols import WorkflowClientPort
|
||||
from .runs import Run
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
@@ -114,3 +117,63 @@ class WorkflowArtifact:
|
||||
from .authoring import EditableWorkflow
|
||||
|
||||
return EditableWorkflow.from_artifact(self)
|
||||
|
||||
async def deploy(
|
||||
self,
|
||||
deployment_id: str,
|
||||
*,
|
||||
bindings: Mapping[str, str] | None = None,
|
||||
drift_policy: str = "block",
|
||||
) -> Deployment:
|
||||
"""Save, inspect, and validate a deployment for this artifact version."""
|
||||
from .deployments import Deployment
|
||||
|
||||
await self._port.save_deployment(
|
||||
{
|
||||
"id": deployment_id,
|
||||
"artifact_id": self.artifact.id,
|
||||
"artifact_version": self.artifact.version,
|
||||
"bindings": dict(bindings or {}),
|
||||
"drift_policy": drift_policy,
|
||||
}
|
||||
)
|
||||
deployment = Deployment.from_payload(
|
||||
self._port,
|
||||
await self._port.inspect_deployment(deployment_id=deployment_id),
|
||||
)
|
||||
if (
|
||||
deployment.artifact_id != self.artifact.id
|
||||
or deployment.artifact_version != self.artifact.version
|
||||
):
|
||||
raise InvalidResponse(
|
||||
operation="workflow.deployments.inspect",
|
||||
details=(
|
||||
f"deployment {deployment_id!r} does not target artifact "
|
||||
f"{self.artifact.id!r} version {self.artifact.version}"
|
||||
),
|
||||
)
|
||||
validation = await deployment.validate()
|
||||
return replace(
|
||||
deployment,
|
||||
diagnostics=validation.diagnostics,
|
||||
runnable=validation.runnable,
|
||||
)
|
||||
|
||||
async def run(
|
||||
self,
|
||||
workflow_input: Mapping[str, Any],
|
||||
*,
|
||||
deployment_id: str | None = None,
|
||||
bindings: Mapping[str, str] | None = None,
|
||||
drift_policy: str = "block",
|
||||
) -> Run:
|
||||
"""Run the artifact under the strict deployment selection policy."""
|
||||
from .deployments import run_artifact
|
||||
|
||||
return await run_artifact(
|
||||
self,
|
||||
workflow_input,
|
||||
deployment_id=deployment_id,
|
||||
bindings=bindings,
|
||||
drift_policy=drift_policy,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,164 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from wf_artifacts import WorkflowArtifact as ArtifactModel
|
||||
from wf_client import DeploymentRequired, WorkflowClientPort
|
||||
from wf_client.errors import DeploymentNotRunnable
|
||||
from wf_client.workflows import WorkflowArtifact
|
||||
from wf_core import Workflow
|
||||
|
||||
|
||||
def _artifact() -> WorkflowArtifact:
|
||||
plan: dict[str, Any] = {
|
||||
"name": "report",
|
||||
"input_schema": {"type": "object", "properties": {}},
|
||||
"state_schema": {"type": "object", "properties": {}},
|
||||
"output_schema": {"type": "object", "properties": {}},
|
||||
"outcomes": ["ok"],
|
||||
"start": "end",
|
||||
"nodes": [{"id": "end", "type": "end", "outcome": "ok"}],
|
||||
"edges": [],
|
||||
}
|
||||
artifact = ArtifactModel(
|
||||
id="report",
|
||||
version=1,
|
||||
title="Report",
|
||||
input_schema=plan["input_schema"],
|
||||
output_schema=plan["output_schema"],
|
||||
outcomes=("ok",),
|
||||
plan=plan,
|
||||
)
|
||||
return WorkflowArtifact(
|
||||
cast(WorkflowClientPort, _FakePort()), artifact, Workflow.model_validate(plan)
|
||||
)
|
||||
|
||||
|
||||
class _FakePort:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[tuple[str, dict[str, Any]]] = []
|
||||
self.list_result: dict[str, Any] = {"deployments": []}
|
||||
self.validation_result: dict[str, Any] = {
|
||||
"deployment_id": "report.production",
|
||||
"artifact_id": "report",
|
||||
"artifact_version": 1,
|
||||
"status": "runnable",
|
||||
"diagnostics": [],
|
||||
"next_actions": {
|
||||
"can_continue": True,
|
||||
"can_save_now": None,
|
||||
"recommended_next_tool": None,
|
||||
"reason": "ready",
|
||||
"patch_examples": [],
|
||||
"warnings": [],
|
||||
},
|
||||
}
|
||||
self.run_result: dict[str, Any] = {
|
||||
"artifact_id": "report",
|
||||
"artifact_version": 1,
|
||||
"deployment_id": "report.production",
|
||||
"status": "completed",
|
||||
"run_id": "run-1",
|
||||
"resume_readiness": "not_applicable",
|
||||
"interrupt": None,
|
||||
"outcome": "ok",
|
||||
"error": None,
|
||||
"output": {"result": "done"},
|
||||
"trace_count": 0,
|
||||
"diagnostics": [],
|
||||
"next_actions": {
|
||||
"can_continue": False,
|
||||
"can_save_now": None,
|
||||
"recommended_next_tool": None,
|
||||
"reason": "done",
|
||||
"patch_examples": [],
|
||||
"warnings": [],
|
||||
},
|
||||
}
|
||||
|
||||
async def save_deployment(self, deployment: dict[str, Any]) -> object:
|
||||
self.calls.append(("save_deployment", {"deployment": deployment}))
|
||||
return {"deployment_id": deployment["id"], "saved": True}
|
||||
|
||||
async def inspect_deployment(self, *, deployment_id: str) -> object:
|
||||
self.calls.append(("inspect_deployment", {"deployment_id": deployment_id}))
|
||||
return {
|
||||
"id": deployment_id,
|
||||
"artifact_id": "report",
|
||||
"artifact_version": 1,
|
||||
"bindings": [
|
||||
{
|
||||
"logical_source": "app.default",
|
||||
"concrete_source": "company.production",
|
||||
}
|
||||
],
|
||||
"drift_policy": "block",
|
||||
}
|
||||
|
||||
async def validate_deployment(self, **params: Any) -> object:
|
||||
self.calls.append(("validate_deployment", params))
|
||||
return self.validation_result
|
||||
|
||||
async def list_deployments(self) -> object:
|
||||
self.calls.append(("list_deployments", {}))
|
||||
return self.list_result
|
||||
|
||||
async def run_deployment(self, **params: Any) -> object:
|
||||
self.calls.append(("run_deployment", params))
|
||||
return self.run_result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_artifact_deploys_with_explicit_bindings() -> None:
|
||||
artifact = _artifact()
|
||||
port = cast(_FakePort, artifact._port)
|
||||
deployment = await artifact.deploy(
|
||||
"report.production", bindings={"app.default": "company.production"}
|
||||
)
|
||||
assert deployment.deployment_id == "report.production"
|
||||
assert deployment.bindings == {"app.default": "company.production"}
|
||||
assert deployment.runnable is True
|
||||
assert [call[0] for call in port.calls[-3:]] == [
|
||||
"save_deployment",
|
||||
"inspect_deployment",
|
||||
"validate_deployment",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_artifact_run_rejects_ambiguous_deployments() -> None:
|
||||
artifact = _artifact()
|
||||
port = cast(_FakePort, artifact._port)
|
||||
port.list_result = {
|
||||
"deployments": [
|
||||
{
|
||||
"id": "report.prod",
|
||||
"artifact_id": "report",
|
||||
"artifact_version": 1,
|
||||
"binding_count": 0,
|
||||
"drift_policy": "block",
|
||||
},
|
||||
{
|
||||
"id": "report.dev",
|
||||
"artifact_id": "report",
|
||||
"artifact_version": 1,
|
||||
"binding_count": 0,
|
||||
"drift_policy": "block",
|
||||
},
|
||||
]
|
||||
}
|
||||
with pytest.raises(DeploymentRequired) as captured:
|
||||
await artifact.run({"topic": "workflow"})
|
||||
assert captured.value.candidate_deployment_ids == ("report.dev", "report.prod")
|
||||
assert not any(call[0] == "run_deployment" for call in port.calls)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deployment_run_rejects_missing_run_id() -> None:
|
||||
artifact = _artifact()
|
||||
deployment = await artifact.deploy("report.production")
|
||||
cast(_FakePort, deployment._port).run_result["run_id"] = None
|
||||
with pytest.raises((DeploymentNotRunnable, AttributeError)):
|
||||
await deployment.run({})
|
||||
@@ -0,0 +1,130 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from wf_client import Run, WorkflowClientPort
|
||||
|
||||
|
||||
def _payload(
|
||||
*, run_id: str | None = "run-1", status: str = "interrupted"
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"artifact_id": "report",
|
||||
"artifact_version": 1,
|
||||
"deployment_id": "report.production",
|
||||
"status": status,
|
||||
"run_id": run_id,
|
||||
"resume_readiness": "ready" if status == "interrupted" else "not_applicable",
|
||||
"interrupt": {
|
||||
"id": "interrupt-1",
|
||||
"frame_id": "root",
|
||||
"node_id": "approve",
|
||||
"kind": "approval",
|
||||
"payload": {"question": "approve?"},
|
||||
"resumable": True,
|
||||
"route": None,
|
||||
"outcomes": ["submitted"],
|
||||
"request_schema": {"type": "object"},
|
||||
"resume_schema": {"type": "object"},
|
||||
"typed": False,
|
||||
}
|
||||
if status == "interrupted"
|
||||
else None,
|
||||
"outcome": None if status == "interrupted" else "ok",
|
||||
"error": None,
|
||||
"output": None if status == "interrupted" else {"result": "done"},
|
||||
"trace_count": 1,
|
||||
"diagnostics": [],
|
||||
"next_actions": {
|
||||
"can_continue": status == "interrupted",
|
||||
"can_save_now": None,
|
||||
"recommended_next_tool": None,
|
||||
"reason": "ready",
|
||||
"patch_examples": [],
|
||||
"warnings": [],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class _Port:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[tuple[str, dict[str, Any]]] = []
|
||||
self.resume_payload = _payload(status="completed")
|
||||
self.trace_payload = {
|
||||
**_payload(status="completed"),
|
||||
"trace": [
|
||||
{
|
||||
"frame_id": "root",
|
||||
"node_id": "approve",
|
||||
"step_type": "node",
|
||||
"resolved_input": {},
|
||||
"outcome": "ok",
|
||||
"next_node_id": "__end__",
|
||||
"output": {},
|
||||
"state_changes": {},
|
||||
}
|
||||
],
|
||||
"trace_start": 0,
|
||||
"trace_limit": 25,
|
||||
"trace_truncated": False,
|
||||
}
|
||||
|
||||
async def resume_run(self, **params: Any) -> object:
|
||||
self.calls.append(("resume_run", params))
|
||||
return self.resume_payload
|
||||
|
||||
async def inspect_run(self, **params: Any) -> object:
|
||||
self.calls.append(("inspect_run", params))
|
||||
return self.resume_payload
|
||||
|
||||
async def read_run_trace(self, **params: Any) -> object:
|
||||
self.calls.append(("read_run_trace", params))
|
||||
return self.trace_payload
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_interrupted_run_resumes_and_reads_bounded_trace() -> None:
|
||||
port = _Port()
|
||||
run = Run.from_payload(cast(WorkflowClientPort, port), _payload())
|
||||
completed = await run.resume({"approved": True})
|
||||
trace = await completed.trace(limit=25)
|
||||
assert completed.status == "completed"
|
||||
assert completed.output == {"result": "done"}
|
||||
assert trace.start == 0
|
||||
assert trace.limit == 25
|
||||
assert len(trace.frames) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_returns_a_new_snapshot() -> None:
|
||||
port = _Port()
|
||||
original = Run.from_payload(cast(WorkflowClientPort, port), _payload())
|
||||
refreshed = await original.refresh()
|
||||
assert refreshed is not original
|
||||
assert refreshed.status == "completed"
|
||||
assert original.status == "interrupted"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_resumable_run_is_rejected_before_io() -> None:
|
||||
port = _Port()
|
||||
payload = _payload()
|
||||
assert payload["interrupt"] is not None
|
||||
payload["interrupt"]["resumable"] = False
|
||||
run = Run.from_payload(cast(WorkflowClientPort, port), payload)
|
||||
with pytest.raises(ValueError, match="not resumable"):
|
||||
await run.resume({"approved": True})
|
||||
assert port.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_trace_rejects_invalid_bounds_before_io() -> None:
|
||||
port = _Port()
|
||||
run = Run.from_payload(cast(WorkflowClientPort, port), _payload())
|
||||
with pytest.raises(ValueError):
|
||||
await run.trace(start=-1)
|
||||
with pytest.raises(ValueError):
|
||||
await run.trace(limit=101)
|
||||
assert port.calls == []
|
||||
Reference in New Issue
Block a user