fix: upgrade placeholder node contracts
This commit is contained in:
@@ -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:
|
||||||
|
|||||||
@@ -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=[])
|
||||||
|
|||||||
Reference in New Issue
Block a user