get end node support in builder and drafts

This commit is contained in:
lda
2026-05-26 00:59:29 +07:00 Verified
parent 6846c649b6
commit 95532b1cf9
7 changed files with 115 additions and 2 deletions
+2
View File
@@ -8,6 +8,7 @@ from .api import (
from .models import ( from .models import (
DraftChooseClause, DraftChooseClause,
DraftChooseStep, DraftChooseStep,
DraftEndStep,
DraftForeachStep, DraftForeachStep,
DraftInterruptStep, DraftInterruptStep,
DraftJoinStep, DraftJoinStep,
@@ -22,6 +23,7 @@ __all__ = [
"DraftDiagnostic", "DraftDiagnostic",
"DraftChooseClause", "DraftChooseClause",
"DraftChooseStep", "DraftChooseStep",
"DraftEndStep",
"DraftForeachStep", "DraftForeachStep",
"DraftInterruptStep", "DraftInterruptStep",
"DraftJoinStep", "DraftJoinStep",
+4
View File
@@ -7,6 +7,7 @@ from wf_core.paths import GraphSourcePath
from .models import ( from .models import (
DraftChooseStep, DraftChooseStep,
DraftEndStep,
DraftForeachStep, DraftForeachStep,
DraftInterruptStep, DraftInterruptStep,
DraftJoinStep, DraftJoinStep,
@@ -29,6 +30,7 @@ def build_workflow_from_draft(draft: WorkflowDraft) -> Workflow:
input_schema=draft.input_schema, input_schema=draft.input_schema,
state_schema=draft.state_schema, state_schema=draft.state_schema,
output_schema=draft.output_schema, output_schema=draft.output_schema,
outcomes=draft.outcomes,
) )
step_refs = { step_refs = {
step_id: _add_step(builder, step_id, step) step_id: _add_step(builder, step_id, step)
@@ -72,6 +74,8 @@ def _add_step(builder: WorkflowBuilder, step_id: str, step: DraftStep):
node = JoinNode(id=step_id, type="join") node = JoinNode(id=step_id, type="join")
builder.nodes.append(node) builder.nodes.append(node)
return node return node
if isinstance(step, DraftEndStep):
return builder.end(step.end.outcome, id=step_id)
if isinstance(step, DraftWhenStep): if isinstance(step, DraftWhenStep):
return builder.when( return builder.when(
step.when.if_, step.when.if_,
+23
View File
@@ -20,6 +20,7 @@ STEP_KIND_KEYS = frozenset(
"foreach", "foreach",
"interrupt", "interrupt",
"join", "join",
"end",
"when", "when",
"choose", "choose",
"match", "match",
@@ -191,6 +192,26 @@ class DraftJoinStep(BaseModel):
join: JsonObject = Field(default_factory=dict) join: JsonObject = Field(default_factory=dict)
class DraftEndPayload(BaseModel):
"""Payload for one explicit workflow terminal outcome."""
model_config = ConfigDict(extra="forbid")
outcome: str = Field(default="ok", min_length=1)
class DraftEndStep(BaseModel):
"""Draft step that lowers to core `EndNode`.
Route to this step when the workflow should finish with a non-`ok` public
outcome. Routing directly to `__end__` remains the shorthand for `ok`.
"""
model_config = ConfigDict(extra="forbid")
end: DraftEndPayload = Field(default_factory=DraftEndPayload)
class DraftWhenPayload(BaseModel): class DraftWhenPayload(BaseModel):
"""Payload for one boolean draft decision.""" """Payload for one boolean draft decision."""
@@ -267,6 +288,7 @@ DraftStep = (
| DraftForeachStep | DraftForeachStep
| DraftInterruptStep | DraftInterruptStep
| DraftJoinStep | DraftJoinStep
| DraftEndStep
| DraftWhenStep | DraftWhenStep
| DraftChooseStep | DraftChooseStep
| DraftMatchStep | DraftMatchStep
@@ -285,6 +307,7 @@ class WorkflowDraft(BaseModel):
input_schema: JsonObject input_schema: JsonObject
state_schema: JsonObject state_schema: JsonObject
output_schema: JsonObject output_schema: JsonObject
outcomes: list[str] = Field(default_factory=lambda: ["ok"], min_length=1)
output: list[InputBinding] = Field(default_factory=list) output: list[InputBinding] = Field(default_factory=list)
start: str start: str
steps: dict[str, DraftStep] steps: dict[str, DraftStep]
+18
View File
@@ -10,6 +10,7 @@ from wf_authoring.ops.values import runtime_error
from wf_core import ( from wf_core import (
ConditionNode, ConditionNode,
Edge, Edge,
EndNode,
ForeachConcurrentPolicy, ForeachConcurrentPolicy,
ForeachItemErrorPolicy, ForeachItemErrorPolicy,
ForeachNode, ForeachNode,
@@ -206,6 +207,7 @@ class WorkflowBuilder:
input_schema: SchemaLike input_schema: SchemaLike
state_schema: StateSchemaLike state_schema: StateSchemaLike
output_schema: SchemaLike output_schema: SchemaLike
outcomes: Sequence[str] | None = None
start: str | None = None start: str | None = None
reducers: ReducerCatalog | Mapping[str, ReducerDefinition] | None = None reducers: ReducerCatalog | Mapping[str, ReducerDefinition] | None = None
node_specs: dict[str, NodeSpec[Any, Any]] = field(default_factory=dict) node_specs: dict[str, NodeSpec[Any, Any]] = field(default_factory=dict)
@@ -220,6 +222,7 @@ class WorkflowBuilder:
self.input_schema = schema_ref_from(self.input_schema) self.input_schema = schema_ref_from(self.input_schema)
self.state_schema = state_schema_from(self.state_schema) self.state_schema = state_schema_from(self.state_schema)
self.output_schema = schema_ref_from(self.output_schema) self.output_schema = schema_ref_from(self.output_schema)
self.outcomes = list(self.outcomes or ["ok"])
@overload @overload
def use( def use(
@@ -418,6 +421,20 @@ class WorkflowBuilder:
self.nodes.append(node) self.nodes.append(node)
return node return node
def end(self, outcome: str = "ok", *, id: str | None = None) -> EndNode:
"""Add an explicit workflow terminal for a declared public outcome.
Routing to `__end__` is still the compact `ok` shorthand. Use this when
the graph should expose a non-`ok` terminal such as `error`.
"""
node = EndNode(
id=id or self._next_step_id(f"end_{slug_id(outcome)}"),
type="end",
outcome=outcome,
)
self.nodes.append(node)
return node
def prepare_subgraph(self, child: WorkflowBuilder) -> Workflow: def prepare_subgraph(self, child: WorkflowBuilder) -> Workflow:
"""Register a local child builder for native execution and return its graph. """Register a local child builder for native execution and return its graph.
@@ -854,6 +871,7 @@ class WorkflowBuilder:
input_schema=cast(SchemaRef, self.input_schema), input_schema=cast(SchemaRef, self.input_schema),
state_schema=cast(StateSchema, self.state_schema), state_schema=cast(StateSchema, self.state_schema),
output_schema=cast(SchemaRef, self.output_schema), output_schema=cast(SchemaRef, self.output_schema),
outcomes=list(self.outcomes or ["ok"]),
node_defs=node_defs, node_defs=node_defs,
start=self.start, start=self.start,
nodes=self.nodes, nodes=self.nodes,
+28 -1
View File
@@ -5,7 +5,7 @@ from pydantic import ValidationError
from wf_artifacts.drafts import WorkflowDraft from wf_artifacts.drafts import WorkflowDraft
from wf_artifacts.drafts.api import compile_workflow_draft, validate_workflow_draft from wf_artifacts.drafts.api import compile_workflow_draft, validate_workflow_draft
from wf_artifacts.drafts.adapter import build_workflow_from_draft from wf_artifacts.drafts.adapter import build_workflow_from_draft
from wf_core import ConditionNode, ForeachNode, NodeUse from wf_core import ConditionNode, EndNode, ForeachNode, NodeUse
from wf_core.models.steps import InputValueBinding from wf_core.models.steps import InputValueBinding
@@ -213,6 +213,33 @@ def test_adapter_lowers_when_step_through_builder() -> None:
] ]
def test_adapter_lowers_explicit_end_step() -> None:
draft = WorkflowDraft.model_validate(
{
"name": "end_example",
"input_schema": {},
"state_schema": {"fields": {}},
"output_schema": {},
"outcomes": ["ok", "error"],
"start": "echo",
"steps": {
"echo": {"use": "demo.echo"},
"end_error": {"end": {"outcome": "error"}},
},
"routes": {"echo": {"error": "end_error"}},
}
)
workflow = build_workflow_from_draft(draft)
terminal = workflow.nodes[1]
assert isinstance(terminal, EndNode)
assert terminal.id == "end_error"
assert terminal.outcome == "error"
assert workflow.outcomes == ["ok", "error"]
assert workflow.edges[0].to == "end_error"
def test_adapter_lowers_choose_step_through_builder() -> None: def test_adapter_lowers_choose_step_through_builder() -> None:
draft = WorkflowDraft.model_validate( draft = WorkflowDraft.model_validate(
{ {
+20
View File
@@ -7,6 +7,7 @@ from pydantic import ValidationError
from wf_artifacts.drafts import ( from wf_artifacts.drafts import (
DraftChooseStep, DraftChooseStep,
DraftEndStep,
DraftForeachStep, DraftForeachStep,
DraftMatchStep, DraftMatchStep,
DraftUseStep, DraftUseStep,
@@ -135,6 +136,25 @@ def test_workflow_draft_accepts_when_step() -> None:
assert isinstance(draft.steps["decide"], DraftWhenStep) assert isinstance(draft.steps["decide"], DraftWhenStep)
def test_workflow_draft_accepts_explicit_end_step() -> None:
draft = WorkflowDraft.model_validate(
{
**_keyed_echo_draft(),
"outcomes": ["ok", "error"],
"steps": {
**_keyed_echo_draft()["steps"],
"end_error": {"end": {"outcome": "error"}},
},
"routes": {"echo": {"error": "end_error"}},
}
)
terminal = draft.steps["end_error"]
assert isinstance(terminal, DraftEndStep)
assert terminal.end.outcome == "error"
def test_workflow_draft_accepts_choose_step() -> None: def test_workflow_draft_accepts_choose_step() -> None:
draft = WorkflowDraft.model_validate( draft = WorkflowDraft.model_validate(
{ {
+20 -1
View File
@@ -14,7 +14,7 @@ from wf_authoring import (
state, state,
state_path, state_path,
) )
from wf_core import END, RunStatus, WorkflowExecutionError from wf_core import END, EndNode, RunStatus, WorkflowExecutionError
from wf_core.models.steps import InputPathBinding, InputValueBinding from wf_core.models.steps import InputPathBinding, InputValueBinding
from wf_core.paths import GraphSourcePath, LocalPath, StatePath from wf_core.paths import GraphSourcePath, LocalPath, StatePath
from wf_platform import CapabilityRef from wf_platform import CapabilityRef
@@ -455,6 +455,25 @@ def test_builder_rejects_mixed_canonical_and_deprecated_output_styles() -> None:
) )
def test_builder_adds_explicit_end_node() -> None:
builder = WorkflowBuilder(
name="explicit_end",
input_schema={},
state_schema={"fields": {}},
output_schema={},
outcomes=["ok", "error"],
)
terminal = builder.end("error", id="end_error")
assert isinstance(terminal, EndNode)
assert terminal.id == "end_error"
assert terminal.outcome == "error"
assert builder.nodes[-1] is terminal
builder.set_entry_point(terminal)
assert builder.compile().outcomes == ["ok", "error"]
class _StructuralKeyMap(Mapping[object, object]): class _StructuralKeyMap(Mapping[object, object]):
def __getitem__(self, key: object) -> object: def __getitem__(self, key: object) -> object:
raise KeyError(key) raise KeyError(key)