diff --git a/src/wf_client/__init__.py b/src/wf_client/__init__.py index d599f25d..43791cf4 100644 --- a/src/wf_client/__init__.py +++ b/src/wf_client/__init__.py @@ -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", diff --git a/src/wf_client/app.py b/src/wf_client/app.py index 94703d78..f4ec4ec7 100644 --- a/src/wf_client/app.py +++ b/src/wf_client/app.py @@ -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), + ) diff --git a/src/wf_client/codec.py b/src/wf_client/codec.py index a4d00f3b..443b44bc 100644 --- a/src/wf_client/codec.py +++ b/src/wf_client/codec.py @@ -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, ...]: diff --git a/src/wf_client/deployments.py b/src/wf_client/deployments.py new file mode 100644 index 00000000..7444dc70 --- /dev/null +++ b/src/wf_client/deployments.py @@ -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) diff --git a/src/wf_client/errors.py b/src/wf_client/errors.py index 7069bebd..a5af9e61 100644 --- a/src/wf_client/errors.py +++ b/src/wf_client/errors.py @@ -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): diff --git a/src/wf_client/runs.py b/src/wf_client/runs.py new file mode 100644 index 00000000..8d0bdbfb --- /dev/null +++ b/src/wf_client/runs.py @@ -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, + ) diff --git a/src/wf_client/workflows.py b/src/wf_client/workflows.py index 927567b8..2add4036 100644 --- a/src/wf_client/workflows.py +++ b/src/wf_client/workflows.py @@ -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, + ) diff --git a/tests/wf_client/test_deployments.py b/tests/wf_client/test_deployments.py new file mode 100644 index 00000000..261fb166 --- /dev/null +++ b/tests/wf_client/test_deployments.py @@ -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({}) diff --git a/tests/wf_client/test_runs.py b/tests/wf_client/test_runs.py new file mode 100644 index 00000000..68e183ff --- /dev/null +++ b/tests/wf_client/test_runs.py @@ -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 == []