From 4f946a35ef8ec11be34338780a5fafd0a1584a09 Mon Sep 17 00:00:00 2001 From: lda Date: Mon, 31 Aug 2026 19:51:43 +0700 Subject: [PATCH] fix: harden Python workflow client boundary --- .../final-fix-report.md | 47 +++ src/wf_api/deployments.py | 8 +- src/wf_artifacts/factory.py | 4 +- src/wf_client/__init__.py | 32 -- src/wf_client/_http_port.py | 302 ++++++++++++++++++ src/wf_client/_identity.py | 31 ++ src/wf_client/app.py | 16 +- src/wf_client/authoring.py | 45 ++- src/wf_client/capabilities.py | 19 ++ src/wf_client/codec.py | 20 ++ src/wf_client/deployments.py | 87 +++-- src/wf_client/errors.py | 23 +- src/wf_client/protocols.py | 21 -- src/wf_client/runs.py | 102 +++++- src/wf_client/workflows.py | 107 ++++--- src/wf_contract_manifest/__init__.py | 1 + src/wf_contract_manifest/__main__.py | 4 +- src/wf_transport_rpc_http/client/base.py | 4 +- .../test_lda_report_workflow_example.py | 4 +- tests/wf_api/test_capability_api.py | 3 +- tests/wf_api/test_composite_input_workflow.py | 3 +- tests/wf_api/test_drafts_service.py | 62 ++-- tests/wf_cli/test_context.py | 8 +- tests/wf_cli/test_remote_target.py | 68 ++-- tests/wf_client/test_app.py | 156 ++++++++- tests/wf_client/test_authoring.py | 87 ++++- tests/wf_client/test_capabilities.py | 73 ++++- tests/wf_client/test_deployments.py | 57 +++- tests/wf_client/test_repr.py | 5 +- tests/wf_client/test_runs.py | 75 ++++- tests/wf_contract_manifest/test_cli.py | 9 +- tests/wf_contract_manifest/test_generate.py | 4 +- tests/wf_mcp/test_saved_subgraphs.py | 4 +- tests/wf_transport_rpc_http/test_client.py | 44 +-- .../test_mcp_backed_server_rpc.py | 24 +- 35 files changed, 1257 insertions(+), 302 deletions(-) create mode 100644 .superpowers/sdd/python-workflow-client/final-fix-report.md create mode 100644 src/wf_client/_http_port.py create mode 100644 src/wf_client/_identity.py diff --git a/.superpowers/sdd/python-workflow-client/final-fix-report.md b/.superpowers/sdd/python-workflow-client/final-fix-report.md new file mode 100644 index 00000000..cecf4b98 --- /dev/null +++ b/.superpowers/sdd/python-workflow-client/final-fix-report.md @@ -0,0 +1,47 @@ +# Final fix recovery report + +## Scope + +Audited the uncommitted patch on `b377311a` against the final review findings +and the Python workflow-client design/plan. The generated `.wf_mcp_store/` and +`test-artifacts/` directories were left untouched and unstaged. + +## Fixes completed + +- Opted the remaining draft-focused RPC client test into `drafts=True`; default + server/storage composition remains draft-free. +- Added the public HTTP port adapter used by `App.from_http_jsonrpc()`. HTTP, + connection, malformed JSON, and malformed JSON-RPC response failures become + `WorkflowClientError` subclasses; known workflow error codes map to stable + subclasses and unknown codes remain inspectable `ProtocolError` values with + code/message/data preserved. +- Added strict identity validation for artifact inspection/save, capability + calls, deployment lifecycle, run lifecycle, and bounded trace pages. +- Made `WorkflowArtifact`, `Deployment`, and `Run` retain deep private copies + and expose defensive copies for nested mutable values. +- Narrowed workflow-plan reconstruction handling to Pydantic `ValidationError`. +- Typed deployment drift policy with the existing `wf_artifacts.DriftPolicy` + enum and removed unused internal client exports/protocol operations. +- Hardened the underlying RPC client against valid JSON values that are not + JSON-RPC objects, and removed `frozen=True` from `RpcProtocolError` so Python + can attach exception traceback state. + +## Verification + +Commands were run from the feature worktree. + +| Command | Result | +| --- | --- | +| `uv run pytest -q tests/wf_client tests/wf_transport_rpc_http` | **292 passed**, 257 warnings | +| `uv run pytest -q tests/wf_client tests/authoring/test_builder.py tests/authoring/test_subgraph.py tests/wf_api/test_artifact_api.py tests/wf_transport_rpc_http/test_client.py tests/wf_transport_rpc_http/test_app.py tests/wf_transport_rpc_http/test_openrpc_contract.py tests/wf_contract_manifest/test_generate.py tests/wf_contract_manifest/test_committed_manifest.py tests/wf_api/test_stores.py tests/wf_server/test_local_static_server.py tests/wf_mcp/test_mcp_workflow_server.py tests/wf_mcp/server/test_tools.py tests/wf_mcp/workflow_surface tests/wf_cli/test_context.py tests/wf_server/test_cli.py` | **429 passed**, 201 warnings | +| `uv run ruff check` | **All checks passed** | +| `uv run ruff format --check` | **677 files already formatted** | +| `uv run basedpyright --level error` | **0 errors, 0 warnings, 0 notes** | +| `uv run python -m wf_contract_manifest check` | **checked** `contracts/workflow-api.manifest.json` | +| `pnpm --dir web --filter @lda/workflow-rpc contract:check` | **passed** | +| `pnpm --dir web --filter @lda/workflow-rpc test` | **151 passed**, 3 skipped | +| `git diff --check` | **passed** | +| `uv run pytest -q` | **2644 passed**, 1 skipped, 1 xfailed; 1 known baseline failure: `tests/docs/test_big_doc_links.py::test_thesis_bundle_has_reproducible_agent_evaluation_assets` (missing generated thesis figure PDFs) | + +The full-suite failure is the documented pre-existing missing-asset failure; +no thesis assets were generated or added. diff --git a/src/wf_api/deployments.py b/src/wf_api/deployments.py index 83665f0c..bb1d15f8 100644 --- a/src/wf_api/deployments.py +++ b/src/wf_api/deployments.py @@ -62,9 +62,7 @@ class WorkflowDeploymentApi: .model_dump(mode="json"), ) - async def save_deployment( - self, deployment: dict[str, Any] - ) -> SaveDeploymentResult: + async def save_deployment(self, deployment: dict[str, Any]) -> SaveDeploymentResult: workflow_deployment = WorkflowDeployment.model_validate(deployment) self._artifact_store().save_deployment(workflow_deployment) self.context.events.record_workflow_event( @@ -83,9 +81,7 @@ class WorkflowDeploymentApi: "saved": True, } - async def delete_deployment( - self, *, deployment_id: str - ) -> DeleteDeploymentResult: + async def delete_deployment(self, *, deployment_id: str) -> DeleteDeploymentResult: """Delete one mutable deployment environment binding.""" self._artifact_store().delete_deployment(deployment_id) self.context.events.record_workflow_event( diff --git a/src/wf_artifacts/factory.py b/src/wf_artifacts/factory.py index b298dd51..a3d165c6 100644 --- a/src/wf_artifacts/factory.py +++ b/src/wf_artifacts/factory.py @@ -2,6 +2,8 @@ from __future__ import annotations from collections.abc import Mapping +from pydantic import ValidationError + from wf_core import ReducerRef, Workflow from wf_core.models.workflow_refs import WorkflowRef from wf_platform import CapabilityRef, NodeSpecInventory, hash_json_schema @@ -69,7 +71,7 @@ def _required_object_field(plan: JsonObject, field_name: str) -> JsonObject: def _validate_workflow_plan(plan: JsonObject) -> None: try: workflow = Workflow.model_validate(plan) - except Exception as exc: + except ValidationError as exc: raise WorkflowPlanValidationError(f"invalid workflow plan: {exc}") from exc node_ids = {node.id for node in workflow.nodes} diff --git a/src/wf_client/__init__.py b/src/wf_client/__init__.py index 43791cf4..91b42b82 100644 --- a/src/wf_client/__init__.py +++ b/src/wf_client/__init__.py @@ -5,22 +5,6 @@ from wf_platform import CapabilityRef, Page from .app import App from .authoring import EditableWorkflow from .capabilities import CapabilityResult, CapabilitySummary, RemoteCapability -from .codec import ( - DecodedRunResult, - DecodedTracePage, - decode_capabilities_page, - decode_capability_call, - decode_capability_diagnostics, - 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, @@ -35,7 +19,6 @@ from .errors import ( ValidationFailed, WorkflowClientError, ) -from .protocols import WorkflowClientPort from .runs import Run, TracePage from .workflows import ( ArtifactRef, @@ -54,8 +37,6 @@ __all__ = [ "CapabilityRef", "CapabilityResult", "CapabilitySummary", - "DecodedRunResult", - "DecodedTracePage", "Diagnostic", "DeploymentNotRunnable", "DeploymentRequired", @@ -70,22 +51,9 @@ __all__ = [ "TransportError", "ValidationFailed", "WorkflowClientError", - "WorkflowClientPort", "EditableWorkflow", "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", - "decode_validate_artifact_plan", - "decode_workflow_artifact", ] diff --git a/src/wf_client/_http_port.py b/src/wf_client/_http_port.py new file mode 100644 index 00000000..d6a19f22 --- /dev/null +++ b/src/wf_client/_http_port.py @@ -0,0 +1,302 @@ +"""HTTP client-port adapter that exposes only public ``wf_client`` errors.""" + +from __future__ import annotations + +from collections.abc import Awaitable, Callable, Sequence +from dataclasses import dataclass +from json import JSONDecodeError +from typing import Any, Literal, TypeVar + +import httpx + +from wf_api.models import ( + CapabilityCallResult, + InspectCapabilityResult, + ListCapabilitiesResult, + ListDeploymentsResult, + RunResult, + RunTraceResult, + SaveArtifactResult, + SaveDeploymentResult, + ValidateArtifactPlanResult, + ValidateDeploymentResult, + WorkflowArtifactPayload, + WorkflowDeploymentPayload, +) +from wf_api.runs import TraceRangeLike +from wf_transport_rpc_http import RpcWorkflowApiClient +from wf_transport_rpc_http.client.base import RpcProtocolError + +from .errors import ( + ArtifactNotFound, + ArtifactVersionConflict, + CapabilityNotFound, + DeploymentNotRunnable, + DeploymentRequired, + ProtocolError, + RevisionConflict, + TransportError, + WorkflowClientError, +) + +_ResultT = TypeVar("_ResultT") + + +def _server_detail(error: RpcProtocolError) -> tuple[str | None, str]: + data = error.data + if not isinstance(data, dict): + return error.code if isinstance(error.code, str) else None, error.message + code = data.get("code") + detail = data.get("message") + return ( + ( + code + if isinstance(code, str) + else error.code + if isinstance(error.code, str) + else None + ), + detail if isinstance(detail, str) else error.message, + ) + + +def _known_protocol_error( + operation: str, + error: RpcProtocolError, +) -> WorkflowClientError | None: + """Translate only stable codes or exact legacy missing-resource signals.""" + code, detail = _server_detail(error) + normalized = code.casefold() if code is not None else "" + if normalized in {"capability_not_found", "capabilitynotfound"}: + return CapabilityNotFound(detail) + if normalized in {"artifact_not_found", "artifactnotfound"}: + return ArtifactNotFound(detail) + if normalized in {"artifact_version_conflict", "artifactversionconflict"}: + return ArtifactVersionConflict(detail) + if normalized in {"revision_conflict", "revisionconflict"}: + return RevisionConflict(detail) + if normalized in {"deployment_required", "deploymentrequired"}: + return DeploymentRequired() + if normalized in {"deployment_not_runnable", "deploymentnotrunnable"}: + return DeploymentNotRunnable(error=detail) + + # The current RPC server reports expected application exception class names + # in ``data.code``. A generic KeyError is safe to specialize only when both + # the operation and its exact resource phrase agree. + if code == "KeyError": + if operation.startswith("workflow.capabilities.") and ( + "unknown workflow capability" in detail + ): + return CapabilityNotFound(detail) + if operation == "workflow.artifacts.inspect" and ( + "unknown workflow artifact" in detail + ): + return ArtifactNotFound(detail) + return None + + +@dataclass(slots=True) +class PublicErrorWorkflowClientPort: + """Delegate RPC operations while preventing transport exception leakage.""" + + _rpc: RpcWorkflowApiClient + + async def _invoke( + self, + operation: str, + call: Callable[..., Awaitable[_ResultT]], + /, + **params: Any, + ) -> _ResultT: + try: + return await call(**params) + except RpcProtocolError as exc: + known = _known_protocol_error(operation, exc) + if known is not None: + raise known from exc + raise ProtocolError(exc.code, exc.message, exc.data) from exc + except (httpx.TransportError, httpx.HTTPStatusError, JSONDecodeError) as exc: + raise TransportError(f"{operation} transport failed: {exc}") from exc + except RuntimeError as exc: + # The RPC transport uses RuntimeError only when a decoded JSON-RPC + # result is not an object. That is a protocol failure, not a public + # transport implementation detail. + raise ProtocolError(None, f"{operation}: {exc}") from exc + + async def list_capabilities( + self, + *, + query: str | None = None, + source_id: str | None = None, + cursor: str | None = None, + limit: int = 50, + ) -> ListCapabilitiesResult: + return await self._invoke( + "workflow.capabilities.list", + self._rpc.list_capabilities, + query=query, + source_id=source_id, + cursor=cursor, + limit=limit, + ) + + async def inspect_capability( + self, *, qualified_name: str + ) -> InspectCapabilityResult: + return await self._invoke( + "workflow.capabilities.inspect", + self._rpc.inspect_capability, + qualified_name=qualified_name, + ) + + async def call_capability( + self, + *, + qualified_name: str, + payload: dict[str, Any], + deployment_id: str | None = None, + ) -> CapabilityCallResult: + return await self._invoke( + "workflow.capabilities.call", + self._rpc.call_capability, + qualified_name=qualified_name, + payload=payload, + deployment_id=deployment_id, + ) + + async def inspect_artifact( + self, *, artifact_id: str, version: int + ) -> WorkflowArtifactPayload: + return await self._invoke( + "workflow.artifacts.inspect", + self._rpc.inspect_artifact, + artifact_id=artifact_id, + version=version, + ) + + async def create_artifact_from_plan( + self, + *, + artifact_id: str, + version: int, + title: str, + plan: dict[str, Any], + outcomes: Sequence[str], + kind: Literal["workflow", "wrapper"] = "workflow", + description: str | None = None, + required_capabilities: dict[str, dict[str, Any]] | None = None, + source_bindings: dict[str, str] | None = None, + created_from_catalog_version: str | None = None, + ) -> SaveArtifactResult: + return await self._invoke( + "workflow.artifacts.create_from_plan", + self._rpc.create_artifact_from_plan, + artifact_id=artifact_id, + version=version, + title=title, + plan=plan, + outcomes=outcomes, + kind=kind, + description=description, + required_capabilities=required_capabilities, + source_bindings=source_bindings, + created_from_catalog_version=created_from_catalog_version, + ) + + async def validate_artifact_plan( + self, + *, + plan: dict[str, Any], + outcomes: Sequence[str], + required_capabilities: dict[str, dict[str, Any]] | None = None, + source_bindings: dict[str, str] | None = None, + ) -> ValidateArtifactPlanResult: + return await self._invoke( + "workflow.artifacts.validate_plan", + self._rpc.validate_artifact_plan, + plan=plan, + outcomes=outcomes, + required_capabilities=required_capabilities, + source_bindings=source_bindings, + ) + + async def list_deployments(self) -> ListDeploymentsResult: + return await self._invoke( + "workflow.deployments.list", + self._rpc.list_deployments, + ) + + async def inspect_deployment( + self, *, deployment_id: str + ) -> WorkflowDeploymentPayload: + return await self._invoke( + "workflow.deployments.inspect", + self._rpc.inspect_deployment, + deployment_id=deployment_id, + ) + + async def save_deployment(self, deployment: dict[str, Any]) -> SaveDeploymentResult: + return await self._invoke( + "workflow.deployments.save", + self._rpc.save_deployment, + deployment=deployment, + ) + + async def validate_deployment( + self, *, deployment_id: str, live_check: bool = False + ) -> ValidateDeploymentResult: + return await self._invoke( + "workflow.deployments.validate", + self._rpc.validate_deployment, + deployment_id=deployment_id, + live_check=live_check, + ) + + async def run_deployment( + self, + *, + deployment_id: str, + workflow_input: dict[str, Any], + trace_range: TraceRangeLike | None = None, + ) -> RunResult: + return await self._invoke( + "workflow.runs.start", + self._rpc.run_deployment, + deployment_id=deployment_id, + workflow_input=workflow_input, + trace_range=trace_range, + ) + + async def inspect_run(self, *, run_id: str) -> RunResult: + return await self._invoke( + "workflow.runs.inspect", + self._rpc.inspect_run, + run_id=run_id, + ) + + async def resume_run( + self, + *, + run_id: str, + resume_payload: dict[str, Any], + resume_outcome: str = "submitted", + trace_range: TraceRangeLike | None = None, + ) -> RunResult: + return await self._invoke( + "workflow.runs.resume", + self._rpc.resume_run, + run_id=run_id, + resume_payload=resume_payload, + resume_outcome=resume_outcome, + trace_range=trace_range, + ) + + async def read_run_trace( + self, *, run_id: str, trace_range: TraceRangeLike + ) -> RunTraceResult: + return await self._invoke( + "workflow.runs.trace", + self._rpc.read_run_trace, + run_id=run_id, + trace_range=trace_range, + ) diff --git a/src/wf_client/_identity.py b/src/wf_client/_identity.py new file mode 100644 index 00000000..a49b095d --- /dev/null +++ b/src/wf_client/_identity.py @@ -0,0 +1,31 @@ +"""Strict identity checks for reconstructing public client snapshots.""" + +from __future__ import annotations + +from collections.abc import Mapping + +from .errors import InvalidResponse + + +def require_response_identity( + *, + operation: str, + actual: Mapping[str, object], + expected: Mapping[str, object], +) -> None: + """Reject a validly shaped response that belongs to another resource. + + Shape validation alone cannot prevent a server, proxy, or cache from + returning the wrong resource. Keeping this check centralized makes every + public reconstruction boundary report the operation and mismatched field. + """ + for field, expected_value in expected.items(): + actual_value = actual.get(field) + if actual_value != expected_value: + raise InvalidResponse( + operation=operation, + details=( + f"response {field} {actual_value!r} does not match " + f"requested {expected_value!r}" + ), + ) diff --git a/src/wf_client/app.py b/src/wf_client/app.py index 0672f57a..1849efd4 100644 --- a/src/wf_client/app.py +++ b/src/wf_client/app.py @@ -9,6 +9,8 @@ from typing import TYPE_CHECKING, Any from wf_platform import CapabilityRef, Page, SourceRef from wf_transport_rpc_http import RpcWorkflowApiClient +from ._http_port import PublicErrorWorkflowClientPort +from ._identity import require_response_identity from .authoring import EditableWorkflow from .capabilities import CapabilitySummary, RemoteCapability from .codec import ( @@ -82,9 +84,11 @@ class App: ) -> App: """Configure a lazy HTTP JSON-RPC connection without performing I/O.""" return cls( - _port=RpcWorkflowApiClient( - url=url, - timeout_seconds=timeout_seconds, + _port=PublicErrorWorkflowClientPort( + RpcWorkflowApiClient( + url=url, + timeout_seconds=timeout_seconds, + ) ), endpoint=url, ) @@ -118,6 +122,7 @@ class App: output_schema=dict(wire["output_schema"]), outcomes=tuple(wire["outcomes"]), is_async=wire["is_async"], + _kind=wire["kind"], ) async def capabilities( @@ -170,6 +175,11 @@ class App: version=version, ) ) + require_response_identity( + operation="workflow.artifacts.inspect", + actual={"artifact_id": artifact.id, "version": artifact.version}, + expected={"artifact_id": artifact_id, "version": version}, + ) return WorkflowArtifact(self._port, artifact, workflow) async def edit_workflow( diff --git a/src/wf_client/authoring.py b/src/wf_client/authoring.py index 7660db8c..20b0b1bd 100644 --- a/src/wf_client/authoring.py +++ b/src/wf_client/authoring.py @@ -19,8 +19,13 @@ from wf_core import ( Workflow, ) +from ._identity import require_response_identity from .capabilities import RemoteCapability -from .codec import decode_validate_artifact_plan, decode_workflow_artifact +from .codec import ( + decode_save_artifact, + decode_validate_artifact_plan, + decode_workflow_artifact, +) from .protocols import WorkflowClientPort from .workflows import ( ArtifactRef, @@ -39,9 +44,7 @@ class EditableWorkflow(WorkflowBuilder): artifact_title: str | None = field(default=None, kw_only=True) artifact_description: str | None = field(default=None, kw_only=True) _source_plan: dict[str, Any] | None = field(default=None, repr=False, kw_only=True) - _source_workflow: Workflow | None = field( - default=None, repr=False, kw_only=True - ) + _source_workflow: Workflow | None = field(default=None, repr=False, kw_only=True) _permissive_node_defs: set[str] = field( default_factory=set, repr=False, kw_only=True ) @@ -191,18 +194,31 @@ class EditableWorkflow(WorkflowBuilder): validation = await self.validate() validation.raise_for_errors() _workflow, plan = self._plan() - saved_id = artifact_id or (self.based_on.artifact_id if self.based_on else self.name) + saved_id = artifact_id or ( + self.based_on.artifact_id if self.based_on else self.name + ) saved_title = title if title is not None else self.artifact_title or self.name saved_description = ( description if description is not None else self.artifact_description ) - await self._port.create_artifact_from_plan( - artifact_id=saved_id, - version=version, - title=saved_title, - plan=plan, - outcomes=tuple(self.outcomes), - description=saved_description, + acknowledgement = decode_save_artifact( + await self._port.create_artifact_from_plan( + artifact_id=saved_id, + version=version, + title=saved_title, + plan=plan, + outcomes=tuple(self.outcomes), + description=saved_description, + ) + ) + require_response_identity( + operation="workflow.artifacts.create_from_plan", + actual={ + "artifact_id": acknowledgement["artifact_id"], + "version": acknowledgement["version"], + "saved": acknowledgement["saved"], + }, + expected={"artifact_id": saved_id, "version": version, "saved": True}, ) # The acknowledgement is only an identity signal. Inspecting the exact # requested version ensures server normalization is retained losslessly. @@ -211,6 +227,11 @@ class EditableWorkflow(WorkflowBuilder): version=version, ) artifact, workflow = decode_workflow_artifact(inspected) + require_response_identity( + operation="workflow.artifacts.inspect", + actual={"artifact_id": artifact.id, "version": artifact.version}, + expected={"artifact_id": saved_id, "version": version}, + ) return WorkflowArtifact(self._port, artifact, workflow) diff --git a/src/wf_client/capabilities.py b/src/wf_client/capabilities.py index 96aabcca..83156042 100644 --- a/src/wf_client/capabilities.py +++ b/src/wf_client/capabilities.py @@ -13,6 +13,7 @@ from wf_artifacts.models import DependencyDiagnostic from wf_core.models.schemas import NodeDef, SchemaRef from wf_platform import CapabilityRef +from ._identity import require_response_identity from ._repr import html_repr, short_repr from .codec import decode_capability_call, decode_capability_diagnostics from .errors import InvalidResponse @@ -113,6 +114,7 @@ class RemoteCapability: output_schema: dict[str, Any] outcomes: tuple[str, ...] is_async: bool + _kind: str = field(default="node_spec", repr=False, compare=False) def __post_init__(self) -> None: # Freeze the public container shape at construction. The nested JSON @@ -200,6 +202,23 @@ class RemoteCapability: f"match requested {self.qualified_name!r}" ), ) + require_response_identity( + operation="workflow.capabilities.call", + actual={ + "source_id": wire["source_id"], + "kind": wire["kind"], + "deployment_id": wire["deployment_id"], + }, + expected={ + "source_id": str(self.ref.source), + "kind": self._kind, + # Direct node calls intentionally ignore deployment ids; saved + # wrapper capabilities echo the selected deployment exactly. + "deployment_id": ( + deployment_id if self._kind == "wrapper_artifact" else None + ), + }, + ) if wire["outcome"] not in self.outcomes: raise InvalidResponse( operation="workflow.capabilities.call", diff --git a/src/wf_client/codec.py b/src/wf_client/codec.py index 56c43ed5..e9a6edf1 100644 --- a/src/wf_client/codec.py +++ b/src/wf_client/codec.py @@ -17,6 +17,8 @@ from wf_api.models import ( RawWorkflowPlan, RunResult, RunTraceResult, + SaveArtifactResult, + SaveDeploymentResult, ValidateArtifactPlanResult, ValidateDeploymentResult, WorkflowArtifactPayload, @@ -121,6 +123,24 @@ def decode_validate_artifact_plan(payload: object) -> ValidateArtifactPlanResult ) +def decode_save_artifact(payload: object) -> SaveArtifactResult: + """Validate an artifact creation acknowledgement.""" + return _validate( + payload, + SaveArtifactResult, + "workflow.artifacts.create_from_plan", + ) + + +def decode_save_deployment(payload: object) -> SaveDeploymentResult: + """Validate a deployment save acknowledgement.""" + return _validate( + payload, + SaveDeploymentResult, + "workflow.deployments.save", + ) + + _ModelT = TypeVar("_ModelT", bound=BaseModel) diff --git a/src/wf_client/deployments.py b/src/wf_client/deployments.py index 3535ac59..e2039083 100644 --- a/src/wf_client/deployments.py +++ b/src/wf_client/deployments.py @@ -51,22 +51,64 @@ class DeploymentValidation: ) -@dataclass(frozen=True, slots=True) +@dataclass(frozen=True, slots=True, init=False) class Deployment: """Immutable snapshot of one configured artifact deployment.""" _port: WorkflowClientPort = field(repr=False, compare=False) - model: WorkflowDeployment - diagnostics: tuple[DependencyDiagnostic, ...] = () + _model: WorkflowDeployment = field(repr=False) + _diagnostics: tuple[DependencyDiagnostic, ...] = field(repr=False) runnable: bool | None = None + def __init__( + self, + _port: WorkflowClientPort, + model: WorkflowDeployment, + diagnostics: tuple[DependencyDiagnostic, ...] = (), + runnable: bool | None = None, + ) -> None: + # Pydantic models remain mutable even inside a frozen dataclass. Keep + # private copies so public inspection cannot retarget later calls. + object.__setattr__(self, "_port", _port) + object.__setattr__(self, "_model", model.model_copy(deep=True)) + object.__setattr__( + self, + "_diagnostics", + tuple(item.model_copy(deep=True) for item in diagnostics), + ) + object.__setattr__(self, "runnable", runnable) + @classmethod def from_payload(cls, port: WorkflowClientPort, payload: object) -> Deployment: return cls(_port=port, model=decode_deployment(payload)) + def with_validation( + self, + *, + diagnostics: tuple[DependencyDiagnostic, ...], + runnable: bool, + ) -> Deployment: + """Return a new snapshot enriched with one validation result.""" + return type(self)( + _port=self._port, + model=self._model, + diagnostics=diagnostics, + runnable=runnable, + ) + + @property + def model(self) -> WorkflowDeployment: + """Return a defensive copy of the deployment domain model.""" + return self._model.model_copy(deep=True) + + @property + def diagnostics(self) -> tuple[DependencyDiagnostic, ...]: + """Return defensive copies of loaded dependency diagnostics.""" + return tuple(item.model_copy(deep=True) for item in self._diagnostics) + @property def deployment_id(self) -> str: - return self.model.id + return self._model.id def __repr__(self) -> str: return short_repr( @@ -89,19 +131,19 @@ class Deployment: @property def artifact_id(self) -> str: - return self.model.artifact_id + return self._model.artifact_id @property def artifact_version(self) -> int: - return self.model.artifact_version + return self._model.artifact_version @property def bindings(self) -> dict[str, str]: - return self.model.binding_map() + return self._model.binding_map() @property def drift_policy(self) -> DriftPolicy: - return self.model.drift_policy + return self._model.drift_policy async def validate(self) -> DeploymentValidation: payload = await self._port.validate_deployment( @@ -171,7 +213,10 @@ class Deployment: if decoded.run_id is None or decoded.status in {"unrunnable", "rejected"}: raise DeploymentNotRunnable( deployment_id=self.deployment_id, - diagnostics=decoded.diagnostics, + diagnostics=tuple( + diagnostic.model_copy(deep=True) + for diagnostic in decoded.diagnostics + ), outcome=decoded.outcome, error=decoded.error, ) @@ -194,7 +239,7 @@ async def run_artifact( *, deployment_id: str | None, bindings: Mapping[str, str] | None, - drift_policy: DriftPolicy | str, + drift_policy: DriftPolicy, ) -> Run: """Apply the artifact's strict deployment-selection policy.""" if deployment_id is not None: @@ -211,14 +256,14 @@ async def run_artifact( ), ) if ( - deployment.artifact_id != artifact.artifact.id - or deployment.artifact_version != artifact.artifact.version + deployment.artifact_id != artifact.ref.artifact_id + or deployment.artifact_version != artifact.ref.version ): raise InvalidResponse( operation="workflow.deployments.inspect", details=( f"deployment {deployment_id!r} does not target artifact " - f"{artifact.artifact.id!r} version {artifact.artifact.version}" + f"{artifact.ref.artifact_id!r} version {artifact.ref.version}" ), ) return await deployment.run(workflow_input) @@ -228,8 +273,8 @@ async def run_artifact( ( summary for summary in summaries - if summary["artifact_id"] == artifact.artifact.id - and summary["artifact_version"] == artifact.artifact.version + if summary["artifact_id"] == artifact.ref.artifact_id + and summary["artifact_version"] == artifact.ref.version ), key=lambda summary: summary["id"], ) @@ -237,7 +282,7 @@ async def run_artifact( raise DeploymentRequired( candidate_deployment_ids=tuple(summary["id"] for summary in matches) ) - default_id = f"{artifact.artifact.id}.v{artifact.artifact.version}.default" + default_id = f"{artifact.ref.artifact_id}.v{artifact.ref.version}.default" if not matches: conflicting = next( ( @@ -245,8 +290,8 @@ async def run_artifact( for summary in summaries if summary["id"] == default_id and ( - summary["artifact_id"] != artifact.artifact.id - or summary["artifact_version"] != artifact.artifact.version + summary["artifact_id"] != artifact.ref.artifact_id + or summary["artifact_version"] != artifact.ref.version ) ), None, @@ -275,14 +320,14 @@ async def run_artifact( ), ) if ( - deployment.artifact_id != artifact.artifact.id - or deployment.artifact_version != artifact.artifact.version + deployment.artifact_id != artifact.ref.artifact_id + or deployment.artifact_version != artifact.ref.version ): raise InvalidResponse( operation="workflow.deployments.inspect", details=( f"deployment {matches[0]['id']!r} does not target artifact " - f"{artifact.artifact.id!r} version {artifact.artifact.version}" + f"{artifact.ref.artifact_id!r} version {artifact.ref.version}" ), ) return await deployment.run(workflow_input) diff --git a/src/wf_client/errors.py b/src/wf_client/errors.py index a5af9e61..48fc4db9 100644 --- a/src/wf_client/errors.py +++ b/src/wf_client/errors.py @@ -2,6 +2,7 @@ from __future__ import annotations +from copy import deepcopy from dataclasses import dataclass from wf_artifacts import DependencyDiagnostic @@ -16,7 +17,27 @@ class TransportError(WorkflowClientError): class ProtocolError(WorkflowClientError): - """The service returned a response that violates its protocol contract.""" + """An inspectable JSON-RPC error not covered by a stable public subclass.""" + + code: int | str | None + message: str + data: object + + def __init__( + self, + code: int | str | None, + message: str, + data: object = None, + ) -> None: + self.code = code + self.message = message + self.data = deepcopy(data) + super().__init__(str(self)) + + def __str__(self) -> str: + if isinstance(self.data, dict) and isinstance(self.data.get("message"), str): + return f"{self.message}: {self.data['message']}" + return self.message @dataclass(slots=True) diff --git a/src/wf_client/protocols.py b/src/wf_client/protocols.py index b294e333..795edbea 100644 --- a/src/wf_client/protocols.py +++ b/src/wf_client/protocols.py @@ -8,10 +8,8 @@ from typing import Any, Literal, Protocol from wf_api.models import ( CapabilityCallResult, InspectCapabilityResult, - ListArtifactsResult, ListCapabilitiesResult, ListDeploymentsResult, - ListRunsResult, RunResult, RunTraceResult, SaveArtifactResult, @@ -55,15 +53,6 @@ class WorkflowClientPort(Protocol): deployment_id: str | None = None, ) -> CapabilityCallResult: ... - async def list_artifacts( - self, - *, - query: str | None = None, - kind: Literal["workflow", "wrapper"] | None = None, - cursor: str | None = None, - limit: int = 50, - ) -> ListArtifactsResult: ... - async def inspect_artifact( self, *, @@ -71,8 +60,6 @@ class WorkflowClientPort(Protocol): version: int, ) -> WorkflowArtifactPayload: ... - async def save_artifact(self, artifact: dict[str, Any]) -> SaveArtifactResult: ... - async def create_artifact_from_plan( self, *, @@ -117,14 +104,6 @@ class WorkflowClientPort(Protocol): live_check: bool = False, ) -> ValidateDeploymentResult: ... - async def list_runs( - self, - *, - status: str | None = None, - cursor: str | None = None, - limit: int = 50, - ) -> ListRunsResult: ... - async def run_deployment( self, *, diff --git a/src/wf_client/runs.py b/src/wf_client/runs.py index c2dadb2b..a389f44e 100644 --- a/src/wf_client/runs.py +++ b/src/wf_client/runs.py @@ -3,6 +3,7 @@ from __future__ import annotations from collections.abc import Mapping +from copy import deepcopy from dataclasses import dataclass, field from typing import Any @@ -12,6 +13,7 @@ from wf_api import TraceRange from wf_artifacts import DependencyDiagnostic from wf_core import InterruptRequest, InterruptRoute, TraceEntry, WorkflowRef +from ._identity import require_response_identity from ._repr import html_repr, short_repr from .codec import DecodedRunResult, decode_run_result, decode_trace_result from .errors import DeploymentNotRunnable, InvalidResponse @@ -85,15 +87,20 @@ def _run_from_decoded( decoded: DecodedRunResult, *, expected_run_id: str | None = None, + expected_deployment_id: str | None = None, operation: str = "workflow.runs.inspect", ) -> Run: if expected_run_id is not None and decoded.run_id != expected_run_id: - raise InvalidResponse( + require_response_identity( operation=operation, - details=( - f"returned run {decoded.run_id!r} does not match requested " - f"{expected_run_id!r}" - ), + actual={"run_id": decoded.run_id}, + expected={"run_id": expected_run_id}, + ) + if expected_deployment_id is not None: + require_response_identity( + operation=operation, + actual={"deployment_id": decoded.deployment_id}, + expected={"deployment_id": expected_deployment_id}, ) if decoded.run_id is None: raise DeploymentNotRunnable( @@ -115,7 +122,7 @@ def _run_from_decoded( ) -@dataclass(frozen=True, slots=True) +@dataclass(frozen=True, slots=True, init=False) class Run: """Immutable client snapshot of one durable deployment run.""" @@ -124,11 +131,53 @@ class Run: deployment_id: str status: str outcome: str | None - output: dict[str, Any] | None - interrupt: InterruptRequest | None - diagnostics: tuple[DependencyDiagnostic, ...] + _output: dict[str, Any] | None = field(repr=False) + _interrupt: InterruptRequest | None = field(repr=False) + _diagnostics: tuple[DependencyDiagnostic, ...] = field(repr=False) trace_count: int + def __init__( + self, + *, + _port: WorkflowClientPort, + 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, + ) -> None: + object.__setattr__(self, "_port", _port) + object.__setattr__(self, "run_id", run_id) + object.__setattr__(self, "deployment_id", deployment_id) + object.__setattr__(self, "status", status) + object.__setattr__(self, "outcome", outcome) + object.__setattr__(self, "_output", deepcopy(output)) + object.__setattr__(self, "_interrupt", deepcopy(interrupt)) + object.__setattr__( + self, + "_diagnostics", + tuple(item.model_copy(deep=True) for item in diagnostics), + ) + object.__setattr__(self, "trace_count", trace_count) + + @property + def output(self) -> dict[str, Any] | None: + """Return a defensive copy of the already-loaded workflow output.""" + return deepcopy(self._output) + + @property + def interrupt(self) -> InterruptRequest | None: + """Return a defensive copy of the already-loaded interrupt contract.""" + return deepcopy(self._interrupt) + + @property + def diagnostics(self) -> tuple[DependencyDiagnostic, ...]: + """Return defensive copies of loaded dependency diagnostics.""" + return tuple(item.model_copy(deep=True) for item in self._diagnostics) + def __repr__(self) -> str: return short_repr( type(self).__name__, @@ -136,8 +185,8 @@ class Run: deployment_id=self.deployment_id, status=self.status, outcome=self.outcome, - output=self.output, - diagnostics=f"{len(self.diagnostics)} diagnostics", + output=self._output, + diagnostics=f"{len(self._diagnostics)} diagnostics", trace=f"{self.trace_count} frames", ) @@ -148,8 +197,8 @@ class Run: deployment_id=self.deployment_id, status=self.status, outcome=self.outcome, - output=self.output, - diagnostics=f"{len(self.diagnostics)} diagnostics", + output=self._output, + diagnostics=f"{len(self._diagnostics)} diagnostics", trace=f"{self.trace_count} frames (use trace() for a bounded page)", ) @@ -160,6 +209,7 @@ class Run: payload: object, *, expected_run_id: str | None = None, + expected_deployment_id: str | None = None, operation: str = "workflow.runs.inspect", ) -> Run: """Validate one run response and reconstruct its immutable snapshot.""" @@ -167,6 +217,7 @@ class Run: port, decode_run_result(payload, operation=operation), expected_run_id=expected_run_id, + expected_deployment_id=expected_deployment_id, operation=operation, ) @@ -176,6 +227,7 @@ class Run: self._port, await self._port.inspect_run(run_id=self.run_id), expected_run_id=self.run_id, + expected_deployment_id=self.deployment_id, operation="workflow.runs.inspect", ) @@ -186,9 +238,9 @@ class Run: outcome: str = "submitted", ) -> Run: """Resume an interrupted run and return the server's new snapshot.""" - if self.status != "interrupted" or self.interrupt is None: + if self.status != "interrupted" or self._interrupt is None: raise ValueError("only interrupted runs can be resumed") - if not self.interrupt.resumable: + if not self._interrupt.resumable: raise ValueError("run interrupt is not resumable") return self.from_payload( self._port, @@ -198,6 +250,7 @@ class Run: resume_outcome=outcome, ), expected_run_id=self.run_id, + expected_deployment_id=self.deployment_id, operation="workflow.runs.resume", ) @@ -213,10 +266,25 @@ class Run: trace_range=TraceRange(start=start, limit=limit), ) ) + require_response_identity( + operation="workflow.runs.trace", + actual={ + "run_id": decoded.run_id, + "deployment_id": decoded.deployment_id, + "trace_start": decoded.trace_start, + "trace_limit": decoded.trace_limit, + }, + expected={ + "run_id": self.run_id, + "deployment_id": self.deployment_id, + "trace_start": start, + "trace_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, + start=start, + limit=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 74a6755e..8b366a76 100644 --- a/src/wf_client/workflows.py +++ b/src/wf_client/workflows.py @@ -3,10 +3,11 @@ from __future__ import annotations from collections.abc import Mapping -from dataclasses import dataclass, field, replace +from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Literal from wf_artifacts.models import ( + DriftPolicy, RequiredCapability, ) from wf_artifacts.models import ( @@ -14,7 +15,9 @@ from wf_artifacts.models import ( ) from wf_core import ValidationReport, Workflow +from ._identity import require_response_identity from ._repr import html_repr, short_repr +from .codec import decode_save_deployment from .errors import InvalidResponse, ValidationFailed if TYPE_CHECKING: @@ -119,17 +122,39 @@ class WorkflowValidation: ) -@dataclass(frozen=True, slots=True) +@dataclass(frozen=True, slots=True, init=False) class WorkflowArtifact: """Immutable client snapshot retaining the validated artifact and workflow.""" _port: WorkflowClientPort = field(repr=False, compare=False) - artifact: ArtifactDomainModel - workflow: Workflow + _artifact: ArtifactDomainModel = field(repr=False) + _workflow: Workflow = field(repr=False) + + def __init__( + self, + port: WorkflowClientPort, + artifact: ArtifactDomainModel, + workflow: Workflow, + ) -> None: + # Frozen dataclasses do not recursively freeze Pydantic models. Retain + # private deep copies and expose only defensive projections below. + object.__setattr__(self, "_port", port) + object.__setattr__(self, "_artifact", artifact.model_copy(deep=True)) + object.__setattr__(self, "_workflow", workflow.model_copy(deep=True)) + + @property + def artifact(self) -> ArtifactDomainModel: + """Return a defensive copy of the validated artifact envelope.""" + return self._artifact.model_copy(deep=True) + + @property + def workflow(self) -> Workflow: + """Return a defensive copy of the executable workflow.""" + return self._workflow.model_copy(deep=True) @property def ref(self) -> ArtifactRef: - return ArtifactRef(self.artifact.id, self.artifact.version) + return ArtifactRef(self._artifact.id, self._artifact.version) def __repr__(self) -> str: return short_repr( @@ -150,28 +175,28 @@ class WorkflowArtifact: @property def title(self) -> str: - return self.artifact.title + return self._artifact.title @property def description(self) -> str | None: - return self.artifact.description + return self._artifact.description @property def required_capabilities(self) -> tuple[RequiredCapability, ...]: return tuple( - capability + capability.model_copy(deep=True) if isinstance(capability, RequiredCapability) else RequiredCapability.model_validate(capability) - for capability in self.artifact.required_capabilities + for capability in self._artifact.required_capabilities ) @property def workflow_dependencies(self) -> dict[str, int]: - return dict(self.artifact.workflow_dependencies) + return dict(self._artifact.workflow_dependencies) def inspect(self) -> Workflow: """Return a deep copy so inspecting an artifact cannot mutate its snapshot.""" - return self.workflow.model_copy(deep=True) + return self._workflow.model_copy(deep=True) def edit(self) -> EditableWorkflow: """Seed an editable builder from this exact immutable artifact version.""" @@ -185,34 +210,37 @@ class WorkflowArtifact: deployment_id: str, *, bindings: Mapping[str, str] | None = None, - drift_policy: str = "block", + drift_policy: DriftPolicy = DriftPolicy.BLOCK, ) -> Deployment: """Save, inspect, and validate a deployment for this artifact version.""" from .deployments import Deployment - saved = 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, - } + saved = decode_save_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, + } + ) + ) + require_response_identity( + operation="workflow.deployments.save", + actual={ + "deployment_id": saved["deployment_id"], + "artifact_id": saved["artifact_id"], + "artifact_version": saved["artifact_version"], + "saved": saved["saved"], + }, + expected={ + "deployment_id": deployment_id, + "artifact_id": self._artifact.id, + "artifact_version": self._artifact.version, + "saved": True, + }, ) - if ( - not isinstance(saved, Mapping) - or saved.get("deployment_id") != deployment_id - ): - saved_id = ( - saved.get("deployment_id") if isinstance(saved, Mapping) else None - ) - raise InvalidResponse( - operation="workflow.deployments.save", - details=( - f"saved deployment id {saved_id!r} does not match requested " - f"{deployment_id!r}" - ), - ) deployment = Deployment.from_payload( self._port, await self._port.inspect_deployment(deployment_id=deployment_id), @@ -226,19 +254,18 @@ class WorkflowArtifact: ), ) if ( - deployment.artifact_id != self.artifact.id - or deployment.artifact_version != self.artifact.version + 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}" + f"{self._artifact.id!r} version {self._artifact.version}" ), ) validation = await deployment.validate() - return replace( - deployment, + return deployment.with_validation( diagnostics=validation.diagnostics, runnable=validation.runnable, ) @@ -249,7 +276,7 @@ class WorkflowArtifact: *, deployment_id: str | None = None, bindings: Mapping[str, str] | None = None, - drift_policy: str = "block", + drift_policy: DriftPolicy = DriftPolicy.BLOCK, ) -> Run: """Run the artifact under the strict deployment selection policy.""" from .deployments import run_artifact diff --git a/src/wf_contract_manifest/__init__.py b/src/wf_contract_manifest/__init__.py index 9384210b..73b18c73 100644 --- a/src/wf_contract_manifest/__init__.py +++ b/src/wf_contract_manifest/__init__.py @@ -22,6 +22,7 @@ def __getattr__(name: str) -> object: return generate_manifest raise AttributeError(name) + __all__ = [ "ContractManifest", "JsonSchema", diff --git a/src/wf_contract_manifest/__main__.py b/src/wf_contract_manifest/__main__.py index f8892551..a508a87c 100644 --- a/src/wf_contract_manifest/__main__.py +++ b/src/wf_contract_manifest/__main__.py @@ -15,7 +15,9 @@ from .model import ManifestError def _parser() -> argparse.ArgumentParser: - parser = argparse.ArgumentParser(description="Manage the checked workflow API contract manifest.") + parser = argparse.ArgumentParser( + description="Manage the checked workflow API contract manifest." + ) parser.add_argument("command", choices=("write", "check")) return parser diff --git a/src/wf_transport_rpc_http/client/base.py b/src/wf_transport_rpc_http/client/base.py index 6dd4b20c..d2df52d1 100644 --- a/src/wf_transport_rpc_http/client/base.py +++ b/src/wf_transport_rpc_http/client/base.py @@ -7,7 +7,7 @@ from uuid import uuid4 import httpx -@dataclass(frozen=True, slots=True) +@dataclass(slots=True) class RpcProtocolError(RuntimeError): """Structured JSON-RPC application error returned by a remote endpoint. @@ -64,6 +64,8 @@ class RpcClientTransport: response = await self.http_client.post(self.url, json=request) response.raise_for_status() payload = response.json() + if not isinstance(payload, dict): + raise RuntimeError("JSON-RPC response must be an object") if "error" in payload: error = payload["error"] if not isinstance(error, dict): diff --git a/tests/examples/test_lda_report_workflow_example.py b/tests/examples/test_lda_report_workflow_example.py index a097677c..89591f76 100644 --- a/tests/examples/test_lda_report_workflow_example.py +++ b/tests/examples/test_lda_report_workflow_example.py @@ -256,9 +256,7 @@ async def test_lda_report_workflow_artifact_interrupt_resume_path( "approved", "selected_issue_ids", } - proposed_ids = [ - issue["id"] for issue in interrupt["payload"]["proposed_issues"] - ] + proposed_ids = [issue["id"] for issue in interrupt["payload"]["proposed_issues"]] assert proposed_ids started_run_id = started["run_id"] assert isinstance(started_run_id, str) diff --git a/tests/wf_api/test_capability_api.py b/tests/wf_api/test_capability_api.py index 383e30bd..0f2ad7d0 100644 --- a/tests/wf_api/test_capability_api.py +++ b/tests/wf_api/test_capability_api.py @@ -39,7 +39,7 @@ def _capability_api( ) service.register_specs("demo.personal", failing_tool) context = context_from_service(service) - return WorkflowCapabilityApi(context), service + return WorkflowCapabilityApi(context, drafts=True), service @pytest.mark.asyncio @@ -226,6 +226,7 @@ async def test_create_draft_workspace_from_capability(tmp_path: Path) -> None: assert "next_actions" in result assert result["wrapper_hints"]["capability_name"] == "demo.personal.echo_tool" + assert api.drafts is not None fetched = await api.drafts.get_draft_workspace( workspace_id="echo_ws", include_draft=True ) diff --git a/tests/wf_api/test_composite_input_workflow.py b/tests/wf_api/test_composite_input_workflow.py index d7583a1f..f976e484 100644 --- a/tests/wf_api/test_composite_input_workflow.py +++ b/tests/wf_api/test_composite_input_workflow.py @@ -42,12 +42,13 @@ def _composite_concat_draft() -> dict[str, object]: async def test_composite_concat_runs_through_the_platform_registry( tmp_path: Path, ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) await server.api.create_draft_workspace( workspace_id="composite_concat", draft=_composite_concat_draft(), ) + assert server.api.draft_authoring is not None authored = await server.api.draft_authoring.set_step_input_bindings( workspace_id="composite_concat", revision=1, diff --git a/tests/wf_api/test_drafts_service.py b/tests/wf_api/test_drafts_service.py index 2336d827..5d6a0391 100644 --- a/tests/wf_api/test_drafts_service.py +++ b/tests/wf_api/test_drafts_service.py @@ -171,7 +171,7 @@ async def test_inspect_draft_authoring_contract_projects_selected_capability( "properties": {"echoed": {"type": "string"}}, } await draft_api.create_draft_workspace(workspace_id="authoring", draft=draft) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) inventory = _authoring_inventory( await api.inspect_draft_authoring_contract( @@ -219,7 +219,7 @@ async def test_inspect_draft_authoring_contract_tolerates_invalid_workflow_schem "properties": {"echoed": {"type": "string"}}, } await draft_api.create_draft_workspace(workspace_id="authoring", draft=draft) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) inventory = _authoring_inventory( await api.inspect_draft_authoring_contract( @@ -256,7 +256,7 @@ async def test_inspect_draft_authoring_contract_resolves_saved_wrapper_capabilit "properties": {"echoed": {"type": "string"}}, } await draft_api.create_draft_workspace(workspace_id="authoring", draft=draft) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) inventory = _authoring_inventory( await api.inspect_draft_authoring_contract( @@ -296,7 +296,7 @@ async def test_inspect_draft_authoring_contract_preserves_empty_capability_schem "properties": {"echoed": {"type": "string"}}, } await draft_api.create_draft_workspace(workspace_id="authoring", draft=draft) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) inventory = _authoring_inventory( await api.inspect_draft_authoring_contract( @@ -344,7 +344,7 @@ async def test_inspect_draft_authoring_contract_warns_for_invalid_capability_sch workspace_id="authoring", draft=_echo_draft(), ) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) inventory = _authoring_inventory( await api.inspect_draft_authoring_contract( @@ -375,7 +375,7 @@ async def test_inspect_draft_authoring_contract_rejects_unknown_selected_step( await draft_api.create_draft_workspace( workspace_id="authoring", draft=_echo_draft() ) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) with pytest.raises(KeyError, match="unknown draft step"): await api.inspect_draft_authoring_contract( @@ -397,7 +397,7 @@ async def test_inspect_draft_authoring_contract_tolerates_invalid_selected_step( draft["steps"] = {"broken": {"unknown_kind": {}}} draft["start"] = "broken" await draft_api.create_draft_workspace(workspace_id="authoring", draft=draft) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) inventory = _authoring_inventory( await api.inspect_draft_authoring_contract( @@ -426,7 +426,7 @@ async def test_inspect_draft_authoring_contract_stale_revision_is_read_only( await draft_api.create_draft_workspace( workspace_id="authoring", draft=_echo_draft() ) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) changed = await api.set_draft_name( workspace_id="authoring", revision=1, @@ -1023,7 +1023,7 @@ async def _create_structured_binding_api( workspace_id=workspace_id, draft=_structured_report_draft(), ) - return draft_api, service, WorkflowApi(authoring.context) + return draft_api, service, WorkflowApi(authoring.context, drafts=True) async def _create_nested_output_binding_api( @@ -1058,7 +1058,7 @@ async def _create_nested_output_binding_api( workspace_id=workspace_id, draft=_nested_report_draft(), ) - return draft_api, service, WorkflowApi(authoring.context) + return draft_api, service, WorkflowApi(authoring.context, drafts=True) @pytest.mark.asyncio @@ -1163,7 +1163,7 @@ async def test_create_empty_draft_workspace_persists_invalid_skeleton( ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / "drafts_create_empty") _drafts, _service, authoring = _draft_api(artifact_store) - facade = WorkflowApi(authoring.context) + facade = WorkflowApi(authoring.context, drafts=True) created = await facade.create_empty_draft_workspace( workspace_id="control_first", @@ -1198,7 +1198,7 @@ async def test_create_empty_draft_workspace_preserves_custom_contract( ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / "drafts_create_contract") _drafts, _service, authoring = _draft_api(artifact_store) - facade = WorkflowApi(authoring.context) + facade = WorkflowApi(authoring.context, drafts=True) input_schema = { "type": "object", "properties": {"topic": {"type": "string"}}, @@ -1249,7 +1249,7 @@ async def test_create_empty_draft_workspace_isolates_default_schemas( ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / "drafts_schema_isolation") _drafts, _service, authoring = _draft_api(artifact_store) - facade = WorkflowApi(authoring.context) + facade = WorkflowApi(authoring.context, drafts=True) input_schema = { "type": "object", "properties": {"topic": {"type": "string"}}, @@ -1283,7 +1283,7 @@ async def test_create_empty_draft_workspace_reports_duplicate_conflict( ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / "drafts_create_conflict") _drafts, _service, authoring = _draft_api(artifact_store) - facade = WorkflowApi(authoring.context) + facade = WorkflowApi(authoring.context, drafts=True) await facade.create_empty_draft_workspace( workspace_id="control_first", name="control_first", @@ -1317,7 +1317,7 @@ async def test_create_empty_draft_workspace_rejects_invalid_contract_before_muta ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / "drafts_create_rejected") _drafts, _service, authoring = _draft_api(artifact_store) - facade = WorkflowApi(authoring.context) + facade = WorkflowApi(authoring.context, drafts=True) with pytest.raises(ValueError): await facade.create_empty_draft_workspace( @@ -1335,7 +1335,7 @@ async def test_set_draft_start_and_contract_replace_top_level_fields_atomically( ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / "drafts_set_lifecycle") _drafts, _service, authoring = _draft_api(artifact_store) - facade = WorkflowApi(authoring.context) + facade = WorkflowApi(authoring.context, drafts=True) await facade.create_empty_draft_workspace( workspace_id="control_first", name="control_first", @@ -1412,7 +1412,7 @@ async def test_lifecycle_edits_reject_invalid_envelopes_without_mutation( tmp_path / f"drafts_lifecycle_rejected_{operation}" ) _drafts, _service, authoring = _draft_api(artifact_store) - facade = WorkflowApi(authoring.context) + facade = WorkflowApi(authoring.context, drafts=True) await facade.create_empty_draft_workspace( workspace_id="control_first", name="control_first", @@ -1453,7 +1453,7 @@ async def test_lifecycle_edits_report_stale_revision_without_mutation( tmp_path / f"drafts_lifecycle_stale_{operation}" ) _drafts, _service, authoring = _draft_api(artifact_store) - facade = WorkflowApi(authoring.context) + facade = WorkflowApi(authoring.context, drafts=True) await facade.create_empty_draft_workspace( workspace_id="control_first", name="control_first", @@ -2187,7 +2187,7 @@ async def test_facade_delegates_semantic_authoring_to_authoring_service( service.register_specs("demo.personal", echo_tool, _snapshot_tool) context = context_from_service(service) - facade = WorkflowApi(context) + facade = WorkflowApi(context, drafts=True) assert facade.draft_authoring is not None assert isinstance(facade.draft_authoring, WorkflowDraftAuthoringApi) @@ -3469,7 +3469,7 @@ async def test_set_step_input_bindings_rejects_remote_target_reference_without_m workspace_id="remote_target", draft=_structured_report_draft("demo.personal.remote_structured_report"), ) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) before = await draft_api.get_draft_workspace( workspace_id="remote_target", include_draft=True, @@ -3506,7 +3506,7 @@ async def test_set_step_input_bindings_rejects_non_capability_step_without_mutat draft = _structured_report_draft() draft["steps"]["report"] = {"join": {}} await draft_api.create_draft_workspace(workspace_id="non_capability", draft=draft) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) before = await draft_api.get_draft_workspace( workspace_id="non_capability", include_draft=True, @@ -4007,7 +4007,7 @@ async def test_add_step_accepts_every_typed_draft_step( ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / f"draft_add_{step_name}") draft_api, _service, authoring = _draft_api(artifact_store, register_echo=True) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) await draft_api.create_draft_workspace(workspace_id="draft_ws", draft=_echo_draft()) step = TypeAdapter(DraftStep).validate_python(step_payload) @@ -4033,7 +4033,7 @@ async def test_add_step_routes_incoming_and_outgoing_edges_atomically( ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / "draft_add_routes") draft_api, _service, authoring = _draft_api(artifact_store, register_echo=True) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) await draft_api.create_draft_workspace(workspace_id="draft_ws", draft=_echo_draft()) step = TypeAdapter(DraftStep).validate_python( @@ -4062,7 +4062,7 @@ async def test_add_step_stale_revision_wins_over_content_preflight( ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / "draft_add_stale") draft_api, _service, authoring = _draft_api(artifact_store, register_echo=True) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) await draft_api.create_draft_workspace(workspace_id="draft_ws", draft=_echo_draft()) before = await draft_api.get_draft_workspace( workspace_id="draft_ws", include_draft=True @@ -4135,7 +4135,7 @@ async def test_add_step_adds_missing_incoming_route_parent_atomically( ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / "draft_add_missing_parent") draft_api, _service, authoring = _draft_api(artifact_store, register_echo=True) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) draft = _echo_draft() draft["routes"] = {} await draft_api.create_draft_workspace(workspace_id="draft_ws", draft=draft) @@ -4164,7 +4164,7 @@ async def test_add_step_distinguishes_missing_and_explicit_empty_routes( ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / "draft_add_empty_routes") draft_api, _service, authoring = _draft_api(artifact_store, register_echo=True) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) await draft_api.create_draft_workspace(workspace_id="draft_ws", draft=_echo_draft()) step_adapter = TypeAdapter(DraftStep) @@ -4207,7 +4207,7 @@ async def test_add_step_rejects_unknown_incoming_outcome_without_mutation( ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / "draft_add_bad_incoming") draft_api, _service, authoring = _draft_api(artifact_store, register_echo=True) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) await draft_api.create_draft_workspace(workspace_id="draft_ws", draft=_echo_draft()) draft_store = authoring.drafts._draft_store() @@ -4277,7 +4277,7 @@ async def test_add_step_rejects_invalid_routing_inputs_atomically( ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / "draft_add_errors") draft_api, _service, authoring = _draft_api(artifact_store, register_echo=True) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) await draft_api.create_draft_workspace(workspace_id="draft_ws", draft=_echo_draft()) await _assert_add_step_rejected_without_mutation( @@ -4338,7 +4338,7 @@ async def test_add_step_rejects_routes_for_non_routable_steps_atomically( ) -> None: artifact_store = FileWorkflowArtifactStore(tmp_path / "draft_add_forbidden_routes") draft_api, _service, authoring = _draft_api(artifact_store, register_echo=True) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) await draft_api.create_draft_workspace(workspace_id="draft_ws", draft=_echo_draft()) before = await draft_api.get_draft_workspace( @@ -4367,7 +4367,7 @@ async def test_add_step_accepts_incomplete_declared_route_subset( artifact_store = FileWorkflowArtifactStore(tmp_path / "draft_add_partial_routes") draft_api, service, authoring = _draft_api(artifact_store, register_echo=True) service.register_specs("demo.personal", echo_tool, _snapshot_tool) - api = WorkflowApi(authoring.context) + api = WorkflowApi(authoring.context, drafts=True) await draft_api.create_draft_workspace(workspace_id="draft_ws", draft=_echo_draft()) step = TypeAdapter(DraftStep).validate_python( @@ -6300,7 +6300,7 @@ def _browser_click_api( _collect_snapshots, ) context = context_from_service(service) - return WorkflowApi(context), service + return WorkflowApi(context, drafts=True), service @pytest.mark.asyncio diff --git a/tests/wf_cli/test_context.py b/tests/wf_cli/test_context.py index 5571572a..d7a0dc02 100644 --- a/tests/wf_cli/test_context.py +++ b/tests/wf_cli/test_context.py @@ -8,6 +8,7 @@ import typer from typer.core import TyperCommand from wf_api import WorkflowApi +from wf_artifacts import FileWorkflowArtifactStore from wf_cli.context import ( CliTyperState, config_path_from_context, @@ -128,10 +129,9 @@ def test_load_cli_context_local_uses_workflow_store_override( assert context.service is None assert isinstance(context.handlers, WorkflowApi) assert context.handlers.drafts_enabled is True - assert ( - context.handlers.context.artifact_store.root - == (tmp_path / ".workflow").resolve() - ) + artifact_store = context.handlers.context.artifact_store + assert isinstance(artifact_store, FileWorkflowArtifactStore) + assert artifact_store.root == (tmp_path / ".workflow").resolve() @pytest.mark.asyncio diff --git a/tests/wf_cli/test_remote_target.py b/tests/wf_cli/test_remote_target.py index 2eb1c858..4233b54a 100644 --- a/tests/wf_cli/test_remote_target.py +++ b/tests/wf_cli/test_remote_target.py @@ -329,7 +329,7 @@ def _patch_rpc_client_to_server(monkeypatch, server) -> None: url=url, timeout_seconds=timeout_seconds, http_client=httpx.AsyncClient( - transport=httpx.ASGITransport(app=create_rpc_app(server)), + transport=httpx.ASGITransport(app=create_rpc_app(server, drafts=True)), base_url="http://test", ), ) @@ -340,7 +340,7 @@ def _patch_rpc_client_to_server(monkeypatch, server) -> None: def test_wf_cap_commands_use_rpc_url_override(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -433,7 +433,7 @@ def test_wf_cap_commands_use_rpc_url_override(monkeypatch, tmp_path) -> None: def test_wf_source_commands_use_rpc_url_override(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -511,7 +511,7 @@ def test_wf_remote_source_inspect_formats_expected_rpc_error( monkeypatch, tmp_path, ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -607,7 +607,7 @@ def test_wf_verbose_shows_full_traceback_for_unexpected_error( def test_wf_admin_commands_use_rpc_url_override(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) server.events.record_workflow_event( "workflow_test_event", capability_id="workflow.demo.v1", @@ -632,7 +632,7 @@ def test_wf_admin_commands_use_rpc_url_override(monkeypatch, tmp_path) -> None: def test_wf_remote_draft_artifact_deploy_lifecycle(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -760,7 +760,7 @@ def test_wf_remote_draft_artifact_deploy_lifecycle(monkeypatch, tmp_path) -> Non def test_wf_remote_capability_free_draft_lifecycle(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -833,7 +833,7 @@ def test_wf_remote_capability_free_draft_lifecycle(monkeypatch, tmp_path) -> Non def test_wf_draft_export_uses_remote_get_and_writes_only_draft( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) rpc_calls: list[tuple[str, dict[str, Any]]] = [] original_call = RpcClientTransport._call @@ -885,7 +885,7 @@ def test_wf_draft_export_uses_remote_get_and_writes_only_draft( def test_wf_draft_import_uses_exact_remote_replacement_payload( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) asyncio.run( server.api.create_empty_draft_workspace( workspace_id="source_ws", @@ -956,7 +956,7 @@ def test_wf_draft_import_uses_exact_remote_replacement_payload( def test_wf_draft_transfer_round_trip_preserves_document_and_destination_id( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) asyncio.run( server.api.create_empty_draft_workspace( workspace_id="source_ws", @@ -1023,7 +1023,7 @@ def test_wf_draft_transfer_round_trip_preserves_document_and_destination_id( def test_wf_remote_run_resume_interrupted_deployment(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) asyncio.run( server.api.create_artifact_from_plan( artifact_id="remote_approval", @@ -1083,7 +1083,7 @@ def test_wf_remote_run_resume_interrupted_deployment(monkeypatch, tmp_path) -> N def test_wf_status_uses_rpc_url_override(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) asyncio.run( server.api.create_artifact_from_plan( artifact_id="status_constant", @@ -1143,7 +1143,7 @@ def test_wf_status_uses_rpc_url_override(monkeypatch, tmp_path) -> None: def test_wf_status_reports_rpc_config_target(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text( @@ -1178,7 +1178,7 @@ def test_wf_status_reports_rpc_config_target(monkeypatch, tmp_path) -> None: def test_wf_draft_delete_requires_confirm(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -1192,7 +1192,7 @@ def test_wf_draft_delete_requires_confirm(monkeypatch, tmp_path) -> None: def test_wf_draft_delete_succeeds_with_confirm(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -1225,7 +1225,7 @@ def test_wf_draft_delete_succeeds_with_confirm(monkeypatch, tmp_path) -> None: def test_wf_source_diagnose_uses_rpc_url_override(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -1300,7 +1300,7 @@ def test_wf_draft_create_reports_optional_inputs_without_binding( def test_wf_draft_set_input_bindings_preserves_composite_expression_over_rpc( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) rpc_calls: list[tuple[str, dict[str, Any]]] = [] original_call = RpcClientTransport._call @@ -1381,7 +1381,7 @@ def test_wf_draft_set_input_bindings_preserves_composite_expression_over_rpc( def test_wf_draft_focused_edit_commands_use_rpc_target(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -1536,7 +1536,7 @@ def test_wf_draft_focused_edit_commands_use_rpc_target(monkeypatch, tmp_path) -> def test_wf_draft_set_workflow_output_replaces_canonical_bindings_over_rpc( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) rpc_calls: list[tuple[str, dict[str, Any]]] = [] original_call = RpcClientTransport._call @@ -1625,7 +1625,7 @@ def test_wf_draft_set_workflow_output_replaces_canonical_bindings_over_rpc( def test_wf_draft_set_workflow_output_merge_uses_compatibility_rpc_target( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) rpc_methods: list[str] = [] original_call = RpcClientTransport._call @@ -1685,7 +1685,7 @@ def test_wf_draft_set_workflow_output_merge_uses_compatibility_rpc_target( def test_wf_draft_set_workflow_output_merge_reports_canonical_replacement( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) rpc_methods: list[str] = [] original_call = RpcClientTransport._call @@ -1760,7 +1760,7 @@ def test_wf_draft_set_workflow_output_merge_reports_canonical_replacement( def test_wf_draft_remove_route_uses_rpc_target(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -1811,7 +1811,7 @@ def test_wf_draft_remove_route_uses_rpc_target(monkeypatch, tmp_path) -> None: def test_wf_draft_bind_uses_rpc_target(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -1868,7 +1868,7 @@ def test_wf_draft_bind_uses_rpc_target(monkeypatch, tmp_path) -> None: def test_wf_draft_set_input_preserves_nested_target_over_rpc( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -1921,7 +1921,7 @@ def test_wf_draft_set_input_preserves_nested_target_over_rpc( def test_wf_draft_set_input_replaces_canonical_bindings_over_rpc( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) rpc_methods: list[str] = [] original_call = RpcClientTransport._call @@ -2024,7 +2024,7 @@ def test_wf_draft_set_input_replaces_canonical_bindings_over_rpc( def test_wf_draft_set_output_replaces_canonical_bindings_over_rpc( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) rpc_calls: list[tuple[str, dict[str, Any]]] = [] original_call = RpcClientTransport._call @@ -2117,7 +2117,7 @@ def test_wf_draft_set_output_replaces_canonical_bindings_over_rpc( def test_wf_draft_add_capability_uses_rpc_target(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) rpc_methods: list[str] = [] original_call = RpcClientTransport._call @@ -2187,7 +2187,7 @@ def test_wf_draft_add_capability_uses_rpc_target(monkeypatch, tmp_path) -> None: def test_wf_draft_capability_add_and_update_preserve_rpc_payloads( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) rpc_calls: list[tuple[str, dict[str, Any]]] = [] original_call = RpcClientTransport._call @@ -2289,7 +2289,7 @@ def test_wf_draft_capability_add_and_update_preserve_rpc_payloads( def test_wf_draft_add_control_steps_use_generic_rpc_target( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) rpc_methods: list[str] = [] original_call = RpcClientTransport._call @@ -2544,7 +2544,7 @@ def test_wf_draft_add_control_steps_use_generic_rpc_target( def test_wf_draft_add_capability_reports_bare_output_target_without_traceback( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -2592,7 +2592,7 @@ def test_wf_draft_add_capability_reports_bare_output_target_without_traceback( def test_wf_draft_compile_prints_compiled_plan(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") @@ -2625,7 +2625,7 @@ def test_wf_draft_compile_prints_compiled_plan(monkeypatch, tmp_path) -> None: def test_wf_draft_compile_invalid_prints_diagnostics_to_stderr( monkeypatch, tmp_path ) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) asyncio.run( server.api.create_draft_workspace( workspace_id="invalid_compile_ws", @@ -2672,7 +2672,7 @@ def test_wf_draft_compile_invalid_prints_diagnostics_to_stderr( def test_wf_deploy_create_alias_saves_deployment(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) asyncio.run( server.api.create_artifact_from_plan( artifact_id="alias_artifact", @@ -2710,7 +2710,7 @@ def test_wf_deploy_create_alias_saves_deployment(monkeypatch, tmp_path) -> None: def test_wf_draft_forward_route_invalid_via_rpc(monkeypatch, tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) _patch_rpc_client_to_server(monkeypatch, server) config_path = tmp_path / "wf.json" config_path.write_text('{"version": 1}', encoding="utf-8") diff --git a/tests/wf_client/test_app.py b/tests/wf_client/test_app.py index 184a1b45..c88c4662 100644 --- a/tests/wf_client/test_app.py +++ b/tests/wf_client/test_app.py @@ -5,8 +5,14 @@ from typing import Any, cast import httpx import pytest -from wf_client import App, CapabilitySummary, Page -from wf_client.errors import InvalidResponse +import wf_client +from wf_client import App, CapabilitySummary, Page, WorkflowClientError +from wf_client.errors import ( + CapabilityNotFound, + InvalidResponse, + ProtocolError, + TransportError, +) from wf_client.protocols import WorkflowClientPort from wf_platform import CapabilityRef @@ -92,6 +98,152 @@ def test_from_http_jsonrpc_is_lazy(monkeypatch: pytest.MonkeyPatch) -> None: assert calls == [] +def test_package_does_not_export_internal_port_or_codecs() -> None: + assert not hasattr(wf_client, "WorkflowClientPort") + assert not hasattr(wf_client, "DecodedRunResult") + assert not hasattr(wf_client, "decode_run_result") + + +@pytest.mark.asyncio +async def test_http_app_translates_connection_failure_to_public_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def fail_post(*args: object, **kwargs: object) -> httpx.Response: + raise httpx.ConnectError("connection refused") + + monkeypatch.setattr(httpx.AsyncClient, "post", fail_post) + app = App.from_http_jsonrpc("http://unreachable.test/rpc") + + with pytest.raises(WorkflowClientError) as raised: + await app.capability("app.default.search") + + assert isinstance(raised.value, TransportError) + assert "connection refused" in str(raised.value) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", ["http", "json", "json-array"]) +async def test_http_app_translates_http_and_json_failures( + monkeypatch: pytest.MonkeyPatch, + failure: str, +) -> None: + async def fail_post(*args: object, **kwargs: object) -> httpx.Response: + request = httpx.Request("POST", "http://test/rpc") + if failure == "http": + return httpx.Response(503, request=request) + if failure == "json-array": + return httpx.Response(200, request=request, json=[]) + return httpx.Response(200, request=request, content=b"not-json") + + monkeypatch.setattr(httpx.AsyncClient, "post", fail_post) + app = App.from_http_jsonrpc("http://test/rpc") + + with pytest.raises(WorkflowClientError) as raised: + await app.capability("app.default.search") + + expected_type = ProtocolError if failure == "json-array" else TransportError + assert isinstance(raised.value, expected_type) + assert "workflow.capabilities.inspect" in str(raised.value) + + +@pytest.mark.asyncio +async def test_http_app_translates_known_workflow_protocol_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def error_post(*args: object, **kwargs: object) -> httpx.Response: + return httpx.Response( + 200, + request=httpx.Request("POST", "http://test/rpc"), + json={ + "jsonrpc": "2.0", + "id": "request", + "error": { + "code": 5000, + "message": "Workflow operation failed", + "data": { + "code": "capability_not_found", + "message": "unknown capability app.default.search", + }, + }, + }, + ) + + monkeypatch.setattr(httpx.AsyncClient, "post", error_post) + app = App.from_http_jsonrpc("http://test/rpc") + + with pytest.raises(WorkflowClientError) as raised: + await app.capability("app.default.search") + + assert isinstance(raised.value, CapabilityNotFound) + assert "unknown capability" in str(raised.value) + + +@pytest.mark.asyncio +async def test_http_app_preserves_unknown_protocol_error_details( + monkeypatch: pytest.MonkeyPatch, +) -> None: + data = {"code": "future_workflow_error", "message": "future detail", "retry": 3} + + async def error_post(*args: object, **kwargs: object) -> httpx.Response: + return httpx.Response( + 200, + request=httpx.Request("POST", "http://test/rpc"), + json={ + "jsonrpc": "2.0", + "id": "request", + "error": { + "code": 5999, + "message": "Future workflow error", + "data": data, + }, + }, + ) + + monkeypatch.setattr(httpx.AsyncClient, "post", error_post) + app = App.from_http_jsonrpc("http://test/rpc") + + with pytest.raises(ProtocolError) as raised: + await app.capability("app.default.search") + + assert raised.value.code == 5999 + assert raised.value.message == "Future workflow error" + assert raised.value.data == data + + +@pytest.mark.asyncio +async def test_workflow_rejects_mismatched_inspected_artifact_identity() -> None: + class ArtifactPort(_Port): + async def inspect_artifact(self, **params: Any) -> object: + return { + "id": "other", + "version": 2, + "title": "Other", + "kind": "workflow", + "description": None, + "input_schema": {"type": "object", "properties": {}}, + "output_schema": {"type": "object", "properties": {}}, + "outcomes": ["ok"], + "plan": { + "name": "other", + "input_schema": {"type": "object", "properties": {}}, + "state_schema": {"type": "object", "properties": {}}, + "output_schema": {"type": "object", "properties": {}}, + "outcomes": ["ok"], + "start": "done", + "nodes": [{"id": "done", "type": "end", "outcome": "ok"}], + "edges": [], + }, + "required_capabilities": [], + "workflow_dependencies": {}, + "created_from_catalog_version": None, + } + + app = App._from_port(cast(WorkflowClientPort, ArtifactPort())) + + with pytest.raises(InvalidResponse, match="workflow.artifacts.inspect"): + await app.workflow("report", version=1) + + @pytest.mark.asyncio async def test_capability_discovery_returns_rich_page() -> None: page = await _app().capabilities(query="search", limit=10) diff --git a/tests/wf_client/test_authoring.py b/tests/wf_client/test_authoring.py index 40296a0d..19980231 100644 --- a/tests/wf_client/test_authoring.py +++ b/tests/wf_client/test_authoring.py @@ -6,6 +6,7 @@ import pytest from wf_authoring import WorkflowBuilder from wf_client import App, ArtifactRef, EditableWorkflow, RemoteCapability +from wf_client.errors import InvalidResponse from wf_client.protocols import WorkflowClientPort from wf_platform import CapabilityRef @@ -20,6 +21,7 @@ class FakePort: "workflow_dependencies": {}, } self.inspect_artifact_result: dict[str, Any] | None = None + self.create_artifact_result: dict[str, Any] | None = None async def validate_artifact_plan(self, **params: Any) -> object: self.calls.append(("validate_artifact_plan", params)) @@ -27,7 +29,11 @@ class FakePort: async def create_artifact_from_plan(self, **params: Any) -> object: self.calls.append(("create_artifact_from_plan", params)) - return {"artifact_id": params["artifact_id"], "version": params["version"], "saved": True} + return self.create_artifact_result or { + "artifact_id": params["artifact_id"], + "version": params["version"], + "saved": True, + } async def inspect_artifact(self, **params: Any) -> object: self.calls.append(("inspect_artifact", params)) @@ -43,13 +49,22 @@ def valid_plan(version: int = 1) -> dict[str, Any]: "kind": "workflow", "description": None, "input_schema": {"type": "object", "properties": {}}, - "output_schema": {"type": "object", "properties": {"value": {"type": "string"}}}, + "output_schema": { + "type": "object", + "properties": {"value": {"type": "string"}}, + }, "outcomes": ["ok"], "plan": { "name": "report", "input_schema": {"type": "object", "properties": {}}, - "state_schema": {"type": "object", "properties": {"value": {"type": "string"}}}, - "output_schema": {"type": "object", "properties": {"value": {"type": "string"}}}, + "state_schema": { + "type": "object", + "properties": {"value": {"type": "string"}}, + }, + "output_schema": { + "type": "object", + "properties": {"value": {"type": "string"}}, + }, "outcomes": ["ok"], "output": [{"path": "state.value", "target": "value"}], "start": "done", @@ -75,9 +90,7 @@ def remote_plan_without_schema_snapshots(version: int = 1) -> dict[str, Any]: {"id": "done", "type": "end", "outcome": "ok"}, ] payload["plan"]["start"] = "remote" - payload["plan"]["edges"] = [ - {"from": "remote", "outcome": "ok", "to": "done"} - ] + payload["plan"]["edges"] = [{"from": "remote", "outcome": "ok", "to": "done"}] payload["required_capabilities"] = [ { "ref": {"source": "app.default", "capability_key": "remote"}, @@ -155,19 +168,73 @@ async def test_edit_and_save_inspects_exact_saved_version() -> None: graph = await app.edit_workflow("report", version=1) assert isinstance(graph, WorkflowBuilder) assert isinstance(graph, EditableWorkflow) - assert all(hasattr(graph, name) for name in ("when", "choose", "match", "foreach", "interrupt", "end", "connect", "set_entry_point")) + assert all( + hasattr(graph, name) + for name in ( + "when", + "choose", + "match", + "foreach", + "interrupt", + "end", + "connect", + "set_entry_point", + ) + ) port.inspect_artifact_result = valid_plan(version=2) saved = await graph.save(version=2) - create = next(params for operation, params in port.calls if operation == "create_artifact_from_plan") + create = next( + params + for operation, params in port.calls + if operation == "create_artifact_from_plan" + ) assert create["plan"] == valid_plan(version=1)["plan"] - inspect = [params for operation, params in port.calls if operation == "inspect_artifact"][-1] + inspect = [ + params for operation, params in port.calls if operation == "inspect_artifact" + ][-1] assert inspect == {"artifact_id": "report", "version": 2} assert saved.ref == ArtifactRef("report", 2) assert str(saved.workflow.output[0].target) == "value" +@pytest.mark.asyncio +async def test_save_rejects_mismatched_create_acknowledgement() -> None: + port = FakePort() + graph = App._from_port(cast(WorkflowClientPort, port)).new_workflow( + "report", + input_schema={"type": "object", "properties": {}}, + state_schema={"type": "object", "properties": {}}, + output_schema={"type": "object", "properties": {}}, + ) + graph.set_entry_point(graph.end("ok", id="done")) + port.create_artifact_result = { + "artifact_id": "other", + "version": 2, + "saved": True, + } + + with pytest.raises(InvalidResponse, match="workflow.artifacts.create_from_plan"): + await graph.save(version=2) + + +@pytest.mark.asyncio +async def test_save_rejects_mismatched_exact_inspection() -> None: + port = FakePort() + graph = App._from_port(cast(WorkflowClientPort, port)).new_workflow( + "report", + input_schema={"type": "object", "properties": {}}, + state_schema={"type": "object", "properties": {}}, + output_schema={"type": "object", "properties": {}}, + ) + graph.set_entry_point(graph.end("ok", id="done")) + port.inspect_artifact_result = valid_plan(version=3) + + with pytest.raises(InvalidResponse, match="workflow.artifacts.inspect"): + await graph.save(version=2) + + @pytest.mark.asyncio async def test_editable_artifact_without_schema_snapshots_remains_saveable() -> None: port = FakePort() diff --git a/tests/wf_client/test_capabilities.py b/tests/wf_client/test_capabilities.py index 2f471c47..4a178e90 100644 --- a/tests/wf_client/test_capabilities.py +++ b/tests/wf_client/test_capabilities.py @@ -36,14 +36,18 @@ def _inspect_payload() -> dict[str, Any]: class _Port: def __init__(self) -> None: self.calls: list[dict[str, Any]] = [] + self.result_qualified_name = "app.default.search" + self.result_source_id = "app.default" + self.result_kind = "node_spec" + self.result_deployment_id: str | None = None async def call_capability(self, **params: Any) -> object: self.calls.append(params) return { - "qualified_name": "app.default.search", - "source_id": "app.default", - "kind": "node_spec", - "deployment_id": None, + "qualified_name": self.result_qualified_name, + "source_id": self.result_source_id, + "kind": self.result_kind, + "deployment_id": self.result_deployment_id, "outcome": "ok", "output": {"results": ["one"]}, "diagnostics": [], @@ -118,6 +122,67 @@ async def test_remote_capability_rejects_mixed_payload_forms() -> None: await capability({"query": "workflow"}, query="again") +@pytest.mark.asyncio +async def test_remote_capability_rejects_mismatched_call_source() -> None: + port = _Port() + port.result_source_id = "other.source" + capability = RemoteCapability( + _port=cast(WorkflowClientPort, port), + ref=CapabilityRef.parse("app.default.search"), + qualified_name="app.default.search", + description=None, + input_schema={"type": "object"}, + output_schema={"type": "object"}, + outcomes=("ok",), + is_async=False, + ) + + with pytest.raises(InvalidResponse, match="workflow.capabilities.call"): + await capability({}) + + +@pytest.mark.asyncio +async def test_node_capability_rejects_unexpected_result_deployment() -> None: + port = _Port() + port.result_deployment_id = "unexpected" + capability = RemoteCapability( + _port=cast(WorkflowClientPort, port), + ref=CapabilityRef.parse("app.default.search"), + qualified_name="app.default.search", + description=None, + input_schema={"type": "object"}, + output_schema={"type": "object"}, + outcomes=("ok",), + is_async=False, + ) + + with pytest.raises(InvalidResponse, match="workflow.capabilities.call"): + await capability.call({}, deployment_id="ignored-by-node-spec") + + +@pytest.mark.asyncio +async def test_wrapper_capability_requires_exact_result_deployment() -> None: + port = _Port() + port.result_qualified_name = "workflow.report.v1" + port.result_source_id = "workflow" + port.result_kind = "wrapper_artifact" + port.result_deployment_id = "other.deployment" + capability = RemoteCapability( + _port=cast(WorkflowClientPort, port), + ref=CapabilityRef(source=SourceRef.parse("workflow"), name="report.v1"), + qualified_name="workflow.report.v1", + description=None, + input_schema={"type": "object"}, + output_schema={"type": "object"}, + outcomes=("ok",), + is_async=False, + _kind="wrapper_artifact", + ) + + with pytest.raises(InvalidResponse, match="workflow.capabilities.call"): + await capability.call({}, deployment_id="report.production") + + def test_remote_capability_rejects_invalid_inspected_schema() -> None: with pytest.raises(InvalidResponse, match="invalid JSON Schema"): RemoteCapability( diff --git a/tests/wf_client/test_deployments.py b/tests/wf_client/test_deployments.py index 6a868dc5..163fb268 100644 --- a/tests/wf_client/test_deployments.py +++ b/tests/wf_client/test_deployments.py @@ -5,8 +5,9 @@ from typing import Any, cast import pytest from wf_artifacts import WorkflowArtifact as ArtifactModel -from wf_client import DeploymentRequired, WorkflowClientPort +from wf_client import DeploymentRequired from wf_client.errors import DeploymentNotRunnable, InvalidResponse +from wf_client.protocols import WorkflowClientPort from wf_client.workflows import WorkflowArtifact from wf_core import Workflow @@ -193,6 +194,16 @@ async def test_artifact_deploy_rejects_wrong_created_deployment_id() -> None: await artifact.deploy("report.production") +@pytest.mark.asyncio +async def test_artifact_deploy_rejects_wrong_created_artifact_identity() -> None: + artifact = _artifact() + port = cast(_FakePort, artifact._port) + port.save_result["artifact_version"] = 2 + + with pytest.raises(InvalidResponse, match="workflow.deployments.save"): + await artifact.deploy("report.production") + + @pytest.mark.asyncio async def test_artifact_deploy_rejects_wrong_inspected_deployment_id() -> None: artifact = _artifact() @@ -310,3 +321,47 @@ async def test_deployment_run_preserves_server_error_and_diagnostics() -> None: assert captured.value.error == "dependency check failed" assert captured.value.outcome == "rejected" assert captured.value.diagnostics[0].code == "missing_source" + + +@pytest.mark.asyncio +async def test_artifact_snapshot_defensively_copies_nested_models() -> None: + artifact = _artifact() + + exposed_artifact = artifact.artifact + exposed_workflow = artifact.workflow + exposed_artifact.id = "mutated" + exposed_artifact.plan["name"] = "mutated" + exposed_workflow.name = "mutated" + + assert artifact.ref.artifact_id == "report" + assert artifact.inspect().name == "report" + assert artifact.edit().name == "report" + + +@pytest.mark.asyncio +async def test_deployment_snapshot_defensively_copies_model_and_diagnostics() -> None: + artifact = _artifact() + port = cast(_FakePort, artifact._port) + port.validation_result["diagnostics"] = [ + { + "severity": "warning", + "code": "drift", + "logical_ref": "app.default", + "bound_source": "company.production", + "message": "original", + "repair_hint": None, + } + ] + deployment = await artifact.deploy("report.production") + + exposed_model = deployment.model + exposed_diagnostics = deployment.diagnostics + exposed_model.id = "mutated" + exposed_model.bindings = [] + exposed_diagnostics[0].message = "mutated" + + assert deployment.deployment_id == "report.production" + assert deployment.bindings == {"app.default": "company.production"} + assert deployment.diagnostics[0].message == "original" + await deployment.run({}) + assert port.calls[-1][1]["deployment_id"] == "report.production" diff --git a/tests/wf_client/test_repr.py b/tests/wf_client/test_repr.py index df8f2587..cb7bbe24 100644 --- a/tests/wf_client/test_repr.py +++ b/tests/wf_client/test_repr.py @@ -126,7 +126,8 @@ def test_repr_does_not_materialize_an_unbounded_iterable() -> None: def test_all_rich_objects_render_without_port_access() -> None: - port = cast(WorkflowClientPort, _port()) + raw_port = _port() + port = cast(WorkflowClientPort, raw_port) diagnostic = WorkflowDiagnostic("error", "bad", "state.x", "broken") local = ValidationReport() objects = [ @@ -153,4 +154,4 @@ def test_all_rich_objects_render_without_port_access() -> None: }, ) ) - assert port.calls == [] + assert raw_port.calls == [] diff --git a/tests/wf_client/test_runs.py b/tests/wf_client/test_runs.py index 0b0b96dc..cbee10e5 100644 --- a/tests/wf_client/test_runs.py +++ b/tests/wf_client/test_runs.py @@ -4,8 +4,9 @@ from typing import Any, cast import pytest -from wf_client import App, Run, WorkflowClientPort +from wf_client import App, Run from wf_client.errors import DeploymentNotRunnable, InvalidResponse +from wf_client.protocols import WorkflowClientPort def _payload( @@ -135,6 +136,22 @@ async def test_resume_rejects_mismatched_result_id() -> None: await run.resume({"approved": True}) +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["refresh", "resume"]) +async def test_run_lifecycle_rejects_mismatched_deployment_identity( + operation: str, +) -> None: + port = _Port() + port.resume_payload["deployment_id"] = "other.deployment" + run = Run.from_payload(cast(WorkflowClientPort, port), _payload()) + + with pytest.raises(InvalidResponse, match="workflow.runs"): + if operation == "refresh": + await run.refresh() + else: + await run.resume({"approved": True}) + + @pytest.mark.asyncio async def test_malformed_interrupt_route_is_invalid_response() -> None: payload = _payload() @@ -180,3 +197,59 @@ async def test_trace_rejects_invalid_bounds_before_io() -> None: with pytest.raises(ValueError): await run.trace(limit=101) assert port.calls == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("field", "value"), + [ + ("run_id", "other-run"), + ("deployment_id", "other.deployment"), + ("trace_start", 1), + ("trace_limit", 26), + ], +) +async def test_trace_rejects_mismatched_identity_or_page( + field: str, + value: object, +) -> None: + port = _Port() + port.trace_payload[field] = value + run = Run.from_payload(cast(WorkflowClientPort, port), _payload()) + + with pytest.raises(InvalidResponse, match="workflow.runs.trace"): + await run.trace(start=0, limit=25) + + +@pytest.mark.asyncio +async def test_run_snapshot_defensively_copies_nested_public_values() -> None: + port = _Port() + payload = _payload() + payload["output"] = {"nested": {"value": "original"}} + payload["diagnostics"] = [ + { + "severity": "warning", + "code": "drift", + "logical_ref": "app.default", + "bound_source": "company.production", + "message": "original", + "repair_hint": None, + } + ] + run = Run.from_payload(cast(WorkflowClientPort, port), payload) + + exposed_output = run.output + exposed_interrupt = run.interrupt + exposed_diagnostics = run.diagnostics + assert exposed_output is not None + assert exposed_interrupt is not None + exposed_output["nested"]["value"] = "mutated" + exposed_interrupt.payload["question"] = "mutated" + exposed_diagnostics[0].message = "mutated" + + assert run.output == {"nested": {"value": "original"}} + assert run.interrupt is not None + assert run.interrupt.payload == {"question": "approve?"} + assert run.diagnostics[0].message == "original" + await run.resume({"approved": True}) + assert port.calls[-1][1]["run_id"] == "run-1" diff --git a/tests/wf_contract_manifest/test_cli.py b/tests/wf_contract_manifest/test_cli.py index 4f23cf72..1a2de778 100644 --- a/tests/wf_contract_manifest/test_cli.py +++ b/tests/wf_contract_manifest/test_cli.py @@ -18,7 +18,9 @@ def _manifest() -> ContractManifest: return manifest_from_openrpc(synthetic_openrpc_document()) -def test_write_generates_once_and_writes_requested_contract(monkeypatch, tmp_path: Path) -> None: +def test_write_generates_once_and_writes_requested_contract( + monkeypatch, tmp_path: Path +) -> None: manifest = _manifest() calls: list[tuple[object, Path]] = [] generate_calls = 0 @@ -35,7 +37,10 @@ def test_write_generates_once_and_writes_requested_contract(monkeypatch, tmp_pat "wf_contract_manifest.__main__.write_manifest", lambda value, path: calls.append((value, path)) or path, ) - monkeypatch.setattr("wf_contract_manifest.__main__.DEFAULT_MANIFEST_PATH", tmp_path / "manifest.json") + monkeypatch.setattr( + "wf_contract_manifest.__main__.DEFAULT_MANIFEST_PATH", + tmp_path / "manifest.json", + ) assert main(["write"]) == 0 assert calls == [(manifest, tmp_path / "manifest.json")] diff --git a/tests/wf_contract_manifest/test_generate.py b/tests/wf_contract_manifest/test_generate.py index 8463cb89..d8177cfd 100644 --- a/tests/wf_contract_manifest/test_generate.py +++ b/tests/wf_contract_manifest/test_generate.py @@ -151,9 +151,7 @@ def test_manifest_separates_recursive_step_inputs_from_workflow_outputs() -> Non input_binding_schema = schemas["InputExpressionBinding"] properties = input_binding_schema.get("properties") assert isinstance(properties, dict) - assert properties["expression"] == { - "$ref": "#/components/schemas/InputExpression" - } + assert properties["expression"] == {"$ref": "#/components/schemas/InputExpression"} expression_schema = schemas["InputExpression"] assert expression_schema["discriminator"] == { "mapping": { diff --git a/tests/wf_mcp/test_saved_subgraphs.py b/tests/wf_mcp/test_saved_subgraphs.py index 72312971..f204a97a 100644 --- a/tests/wf_mcp/test_saved_subgraphs.py +++ b/tests/wf_mcp/test_saved_subgraphs.py @@ -177,9 +177,7 @@ def test_interrupted_saved_child_blocks_resume_until_pinned_source_returns( assert blocked["status"] == "interrupted" assert blocked["resume_readiness"] == "blocked" assert blocked["diagnostics"][0]["code"] == "source_disabled" - assert ( - run_store.get_run(paused_run_id).resume_readiness is ResumeReadiness.BLOCKED - ) + assert run_store.get_run(paused_run_id).resume_readiness is ResumeReadiness.BLOCKED assert run_store.get_latest_checkpoint(paused_run_id).sequence == 1 handlers.service.capability_sources["demo.personal"].enabled = True diff --git a/tests/wf_transport_rpc_http/test_client.py b/tests/wf_transport_rpc_http/test_client.py index bd011d99..22175048 100644 --- a/tests/wf_transport_rpc_http/test_client.py +++ b/tests/wf_transport_rpc_http/test_client.py @@ -1,6 +1,5 @@ from __future__ import annotations -import sys from typing import Any import httpx @@ -32,25 +31,6 @@ from wf_transport_rpc_http.client.drafts import RpcDraftClientMixin from wf_transport_rpc_http.client.sources import RpcSourceAdminClientMixin -@pytest.fixture(autouse=True) -def _draft_enabled_composition(monkeypatch: pytest.MonkeyPatch) -> None: - """Opt draft-focused RPC client tests into the explicit draft surface.""" - build_local = build_local_static_workflow_server - create_app = create_rpc_app - - def draft_local(root, *args, **kwargs): - kwargs.setdefault("drafts", True) - return build_local(root, *args, **kwargs) - - def draft_app(server, *args, **kwargs): - kwargs.setdefault("drafts", True) - return create_app(server, *args, **kwargs) - - module = sys.modules[__name__] - monkeypatch.setattr(module, "build_local_static_workflow_server", draft_local) - monkeypatch.setattr(module, "create_rpc_app", draft_app) - - async def test_rpc_client_preserves_structured_jsonrpc_error() -> None: def handler(request: httpx.Request) -> httpx.Response: return httpx.Response( @@ -354,8 +334,8 @@ async def test_rpc_workflow_client_lists_inspects_validates_and_deletes_deployme async def test_rpc_workflow_client_draft_workspace_lifecycle(tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") - app = create_rpc_app(server) + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) + app = create_rpc_app(server, drafts=True) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient( transport=transport, base_url="http://test" @@ -550,8 +530,8 @@ async def test_rpc_client_sends_exact_replace_document_payload() -> None: async def test_rpc_client_builds_capability_free_draft_lifecycle(tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") - app = create_rpc_app(server) + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) + app = create_rpc_app(server, drafts=True) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient( transport=transport, @@ -625,8 +605,8 @@ def test_rpc_client_satisfies_draft_surface_static_shape() -> None: async def test_rpc_workflow_client_deletes_draft_workspace(tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") - app = create_rpc_app(server) + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) + app = create_rpc_app(server, drafts=True) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient( transport=transport, base_url="http://test" @@ -772,8 +752,8 @@ async def test_rpc_client_validates_artifact_plan_without_persisting(tmp_path) - async def test_rpc_client_set_workflow_output_map(tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") - app = create_rpc_app(server) + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) + app = create_rpc_app(server, drafts=True) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient( transport=transport, base_url="http://test" @@ -807,8 +787,8 @@ async def test_rpc_client_set_workflow_output_map(tmp_path) -> None: async def test_rpc_client_draft_workspace_focused_edit_methods(tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") - app = create_rpc_app(server) + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) + app = create_rpc_app(server, drafts=True) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient( transport=transport, base_url="http://test" @@ -1069,8 +1049,8 @@ async def test_rpc_client_draft_remove_methods(tmp_path) -> None: async def test_rpc_client_draft_workspace_add_step_from_capability(tmp_path) -> None: - server = build_local_static_workflow_server(tmp_path / "store") - app = create_rpc_app(server) + server = build_local_static_workflow_server(tmp_path / "store", drafts=True) + app = create_rpc_app(server, drafts=True) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient( transport=transport, base_url="http://test" diff --git a/tests/wf_transport_rpc_http/test_mcp_backed_server_rpc.py b/tests/wf_transport_rpc_http/test_mcp_backed_server_rpc.py index dab971fc..5248b6c7 100644 --- a/tests/wf_transport_rpc_http/test_mcp_backed_server_rpc.py +++ b/tests/wf_transport_rpc_http/test_mcp_backed_server_rpc.py @@ -200,7 +200,7 @@ def _runtime_reuse_server( ], ) store_roots = config.store_roots - workflow_stores = file_workflow_stores(store_roots.workflow_root) + workflow_stores = file_workflow_stores(store_roots.workflow_root, drafts=True) auth_store = FileAuthStore(store_roots.auth_root) catalog_store = FileCatalogStore(store_roots.catalog_cache_root) factory = _RecordingSessionFactory() @@ -232,7 +232,7 @@ async def test_mcp_backed_rpc_lists_and_mutates_source_registry(tmp_path) -> Non SourceRegistryFile(sources=[_registry_entry("demo.registry")]) ) server = build_workflow_server_from_config(config) - app = create_rpc_app(server) + app = create_rpc_app(server, drafts=True) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient( @@ -255,7 +255,7 @@ async def test_mcp_backed_rpc_lists_and_mutates_source_registry(tmp_path) -> Non async def test_mcp_backed_rpc_capability_list_filters_by_source(tmp_path) -> None: config = BrokerConfig(store_root=tmp_path / "store", connections=[]) server = build_workflow_server_from_config(config) - app = create_rpc_app(server) + app = create_rpc_app(server, drafts=True) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient( @@ -285,7 +285,7 @@ async def test_mcp_backed_rpc_reports_connections_and_events(tmp_path) -> None: ], ) server = build_workflow_server_from_config(config) - app = create_rpc_app(server) + app = create_rpc_app(server, drafts=True) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient( @@ -299,7 +299,7 @@ async def test_mcp_backed_rpc_reports_connections_and_events(tmp_path) -> None: async def test_mcp_backed_rpc_applies_source_registry_changes(tmp_path) -> None: config = BrokerConfig(store_root=tmp_path / "store", connections=[]) server = build_workflow_server_from_config(config) - app = create_rpc_app(server) + app = create_rpc_app(server, drafts=True) async with httpx.AsyncClient( transport=httpx.ASGITransport(app=app), @@ -368,7 +368,7 @@ async def test_mcp_backed_rpc_can_be_built_from_neutral_workflow_config( } ) server = build_workflow_server_from_workflow_config(workflow_config) - app = create_rpc_app(server) + app = create_rpc_app(server, drafts=True) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient( @@ -408,7 +408,7 @@ async def test_mcp_backed_rpc_resumes_interrupted_run_after_server_rebuild( } ) async with httpx.AsyncClient( - transport=httpx.ASGITransport(app=create_rpc_app(first_server)), + transport=httpx.ASGITransport(app=create_rpc_app(first_server, drafts=True)), base_url="http://test", ) as http_client: first_client = RpcWorkflowApiClient( @@ -433,7 +433,7 @@ async def test_mcp_backed_rpc_resumes_interrupted_run_after_server_rebuild( rebuilt_server = build_workflow_server_from_workflow_config(workflow_config) async with httpx.AsyncClient( - transport=httpx.ASGITransport(app=create_rpc_app(rebuilt_server)), + transport=httpx.ASGITransport(app=create_rpc_app(rebuilt_server, drafts=True)), base_url="http://test", ) as http_client: rebuilt_client = RpcWorkflowApiClient( @@ -463,7 +463,7 @@ async def test_mcp_backed_rpc_workflow_reuses_runtime_session_across_runs( assert factory.created_connections[0].id == "fixture.default" async with httpx.AsyncClient( - transport=httpx.ASGITransport(app=create_rpc_app(server)), + transport=httpx.ASGITransport(app=create_rpc_app(server, drafts=True)), base_url="http://test", ) as http_client: client = RpcWorkflowApiClient(url="http://test/rpc", http_client=http_client) @@ -597,7 +597,7 @@ async def test_mcp_backed_rpc_workflow_reuses_runtime_session_direct_setup( ) async with httpx.AsyncClient( - transport=httpx.ASGITransport(app=create_rpc_app(server)), + transport=httpx.ASGITransport(app=create_rpc_app(server, drafts=True)), base_url="http://test", ) as http_client: client = RpcWorkflowApiClient(url="http://test/rpc", http_client=http_client) @@ -696,7 +696,7 @@ async def test_mcp_backed_rpc_deployment_becomes_unrunnable_after_source_removed ) async with httpx.AsyncClient( - transport=httpx.ASGITransport(app=create_rpc_app(server)), + transport=httpx.ASGITransport(app=create_rpc_app(server, drafts=True)), base_url="http://test", ) as http_client: client = RpcWorkflowApiClient(url="http://test/rpc", http_client=http_client) @@ -748,7 +748,7 @@ async def test_mcp_backed_rpc_workflow_reuses_real_stdio_fixture_session( ) async with httpx.AsyncClient( - transport=httpx.ASGITransport(app=create_rpc_app(server)), + transport=httpx.ASGITransport(app=create_rpc_app(server, drafts=True)), base_url="http://test", ) as http_client: client = RpcWorkflowApiClient(url="http://test/rpc", http_client=http_client)