fix: upgrade placeholder node contracts

This commit is contained in:
lda
2026-08-31 01:42:48 +07:00 Verified
parent 870fa65f0a
commit d4ca8ec6a6
2 changed files with 61 additions and 3 deletions
+19 -2
View File
@@ -42,12 +42,15 @@ class EditableWorkflow(WorkflowBuilder):
_source_workflow: Workflow | None = field( _source_workflow: Workflow | None = field(
default=None, repr=False, kw_only=True default=None, repr=False, kw_only=True
) )
_permissive_node_defs: set[str] = field(
default_factory=set, repr=False, kw_only=True
)
@classmethod @classmethod
def from_artifact(cls, artifact: WorkflowArtifact) -> EditableWorkflow: def from_artifact(cls, artifact: WorkflowArtifact) -> EditableWorkflow:
"""Copy every canonical graph field from an immutable artifact snapshot.""" """Copy every canonical graph field from an immutable artifact snapshot."""
builder = WorkflowBuilder.from_workflow(artifact.workflow) builder = WorkflowBuilder.from_workflow(artifact.workflow)
_seed_remote_node_defs(builder, artifact) permissive_node_defs = _seed_remote_node_defs(builder, artifact)
return cls( return cls(
_port=artifact._port, _port=artifact._port,
based_on=artifact.ref, based_on=artifact.ref,
@@ -68,6 +71,7 @@ class EditableWorkflow(WorkflowBuilder):
prepared_subgraphs=builder.prepared_subgraphs, prepared_subgraphs=builder.prepared_subgraphs,
_source_plan=deepcopy(artifact.artifact.plan), _source_plan=deepcopy(artifact.artifact.plan),
_source_workflow=builder._build_workflow(start=builder.start or ""), _source_workflow=builder._build_workflow(start=builder.start or ""),
_permissive_node_defs=permissive_node_defs,
) )
@overload @overload
@@ -101,6 +105,13 @@ class EditableWorkflow(WorkflowBuilder):
return self.use_contract(spec.node_def(), **kwargs) return self.use_contract(spec.node_def(), **kwargs)
return super().use(spec, **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( def subgraph(
self, self,
workflow: WorkflowArtifact, workflow: WorkflowArtifact,
@@ -206,7 +217,7 @@ class EditableWorkflow(WorkflowBuilder):
def _seed_remote_node_defs( def _seed_remote_node_defs(
builder: WorkflowBuilder, builder: WorkflowBuilder,
artifact: WorkflowArtifact, artifact: WorkflowArtifact,
) -> None: ) -> set[str]:
"""Restore remote node contracts retained as artifact dependency snapshots.""" """Restore remote node contracts retained as artifact dependency snapshots."""
node_name_by_step_id = { node_name_by_step_id = {
node.id: node.node node.id: node.node
@@ -218,6 +229,7 @@ def _seed_remote_node_defs(
node_name = node_name_by_step_id.get(edge.from_) node_name = node_name_by_step_id.get(edge.from_)
if node_name is not None: if node_name is not None:
outcomes_by_node.setdefault(node_name, []).append(edge.outcome) outcomes_by_node.setdefault(node_name, []).append(edge.outcome)
permissive_names: set[str] = set()
for requirement in artifact.required_capabilities: for requirement in artifact.required_capabilities:
if requirement.kind != "node_spec": if requirement.kind != "node_spec":
continue continue
@@ -247,6 +259,10 @@ def _seed_remote_node_defs(
requirement.output_schema_snapshot, requirement.output_schema_snapshot,
output_fields, 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( builder.seeded_node_defs.setdefault(
name, name,
NodeDef( NodeDef(
@@ -256,6 +272,7 @@ def _seed_remote_node_defs(
outcomes=outcomes_by_node.get(name, ["ok"]), outcomes=outcomes_by_node.get(name, ["ok"]),
), ),
) )
return permissive_names
def _binding_root_field(path: object) -> str | None: def _binding_root_field(path: object) -> str | None:
+42 -1
View File
@@ -5,8 +5,9 @@ from typing import Any, cast
import pytest import pytest
from wf_authoring import WorkflowBuilder 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_client.protocols import WorkflowClientPort
from wf_platform import CapabilityRef
class FakePort: 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) port.inspect_artifact_result = remote_plan_without_schema_snapshots(version=2)
saved = await graph.save(version=2) saved = await graph.save(version=2)
assert saved.ref == ArtifactRef("report", 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=[])