From d4ca8ec6a611baa2a452f7c3982be0e4699bd8a6 Mon Sep 17 00:00:00 2001 From: lda Date: Mon, 31 Aug 2026 01:42:48 +0700 Subject: [PATCH] fix: upgrade placeholder node contracts --- src/wf_client/authoring.py | 21 +++++++++++++-- tests/wf_client/test_authoring.py | 43 ++++++++++++++++++++++++++++++- 2 files changed, 61 insertions(+), 3 deletions(-) diff --git a/src/wf_client/authoring.py b/src/wf_client/authoring.py index 4c33f232..2be054e3 100644 --- a/src/wf_client/authoring.py +++ b/src/wf_client/authoring.py @@ -42,12 +42,15 @@ class EditableWorkflow(WorkflowBuilder): _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 + ) @classmethod def from_artifact(cls, artifact: WorkflowArtifact) -> EditableWorkflow: """Copy every canonical graph field from an immutable artifact snapshot.""" builder = WorkflowBuilder.from_workflow(artifact.workflow) - _seed_remote_node_defs(builder, artifact) + permissive_node_defs = _seed_remote_node_defs(builder, artifact) return cls( _port=artifact._port, based_on=artifact.ref, @@ -68,6 +71,7 @@ class EditableWorkflow(WorkflowBuilder): prepared_subgraphs=builder.prepared_subgraphs, _source_plan=deepcopy(artifact.artifact.plan), _source_workflow=builder._build_workflow(start=builder.start or ""), + _permissive_node_defs=permissive_node_defs, ) @overload @@ -101,6 +105,13 @@ class EditableWorkflow(WorkflowBuilder): return self.use_contract(spec.node_def(), **kwargs) return super().use(spec, **kwargs) + def use_contract(self, node_def: NodeDef, **kwargs: Any) -> NodeUse: + """Upgrade an artifact placeholder before normal duplicate checks.""" + if node_def.name in self._permissive_node_defs: + self.seeded_node_defs.pop(node_def.name, None) + self._permissive_node_defs.remove(node_def.name) + return super().use_contract(node_def, **kwargs) + def subgraph( self, workflow: WorkflowArtifact, @@ -206,7 +217,7 @@ class EditableWorkflow(WorkflowBuilder): def _seed_remote_node_defs( builder: WorkflowBuilder, artifact: WorkflowArtifact, -) -> None: +) -> set[str]: """Restore remote node contracts retained as artifact dependency snapshots.""" node_name_by_step_id = { node.id: node.node @@ -218,6 +229,7 @@ def _seed_remote_node_defs( node_name = node_name_by_step_id.get(edge.from_) if node_name is not None: outcomes_by_node.setdefault(node_name, []).append(edge.outcome) + permissive_names: set[str] = set() for requirement in artifact.required_capabilities: if requirement.kind != "node_spec": continue @@ -247,6 +259,10 @@ def _seed_remote_node_defs( requirement.output_schema_snapshot, output_fields, ) + if not isinstance(requirement.input_schema_snapshot, dict) or not isinstance( + requirement.output_schema_snapshot, dict + ): + permissive_names.add(name) builder.seeded_node_defs.setdefault( name, NodeDef( @@ -256,6 +272,7 @@ def _seed_remote_node_defs( outcomes=outcomes_by_node.get(name, ["ok"]), ), ) + return permissive_names def _binding_root_field(path: object) -> str | None: diff --git a/tests/wf_client/test_authoring.py b/tests/wf_client/test_authoring.py index ce1750ca..6487e30e 100644 --- a/tests/wf_client/test_authoring.py +++ b/tests/wf_client/test_authoring.py @@ -5,8 +5,9 @@ from typing import Any, cast import pytest from wf_authoring import WorkflowBuilder -from wf_client import App, ArtifactRef, EditableWorkflow +from wf_client import App, ArtifactRef, EditableWorkflow, RemoteCapability from wf_client.protocols import WorkflowClientPort +from wf_platform import CapabilityRef class FakePort: @@ -165,3 +166,43 @@ async def test_editable_artifact_without_schema_snapshots_remains_saveable() -> port.inspect_artifact_result = remote_plan_without_schema_snapshots(version=2) saved = await graph.save(version=2) assert saved.ref == ArtifactRef("report", 2) + + +@pytest.mark.asyncio +async def test_snapshotless_remote_node_upgrades_to_real_capability_contract() -> None: + port = FakePort() + port.inspect_artifact_result = remote_plan_without_schema_snapshots(version=1) + graph = await App._from_port(cast(WorkflowClientPort, port)).edit_workflow( + "report", version=1 + ) + capability = RemoteCapability( + _port=cast(WorkflowClientPort, port), + ref=CapabilityRef.parse("app.default.remote"), + qualified_name="app.default.remote", + description=None, + input_schema={"type": "object", "properties": {"query": {"type": "string"}}}, + output_schema={"type": "object", "properties": {}}, + outcomes=("ok",), + is_async=False, + ) + + replacement = graph.use(capability, id="replacement", input=[], output=[]) + graph.connect(replacement, "ok", "done") + + assert graph.validate_local().ok is True + assert graph.seeded_node_defs["app.default.remote"].input_schema.properties == { + "query": {"type": "string"} + } + + incompatible = RemoteCapability( + _port=cast(WorkflowClientPort, port), + ref=CapabilityRef.parse("app.default.remote"), + qualified_name="app.default.remote", + description=None, + input_schema={"type": "object", "properties": {"other": {"type": "string"}}}, + output_schema={"type": "object", "properties": {}}, + outcomes=("ok",), + is_async=False, + ) + with pytest.raises(ValueError, match="incompatible duplicate"): + graph.use(incompatible, id="incompatible", input=[], output=[])