added some good stuff
This commit is contained in:
@@ -15,9 +15,12 @@ from wf_core import (
|
|||||||
StateSchema,
|
StateSchema,
|
||||||
Workflow,
|
Workflow,
|
||||||
)
|
)
|
||||||
|
from wf_core.errors import WorkflowExecutionError
|
||||||
from wf_core.model import Condition as CoreCondition
|
from wf_core.model import Condition as CoreCondition
|
||||||
|
|
||||||
from .dsl import Expr, GraphPath, PathArg, compile_condition
|
from .dsl import Expr, GraphPath, PathArg, compile_condition
|
||||||
|
from .nodes.callables import SyncRegistryHandler
|
||||||
|
from .nodes.registry import build_registry
|
||||||
from .schemas import SchemaLike, StateSchemaLike, schema_ref_from, state_schema_from
|
from .schemas import SchemaLike, StateSchemaLike, schema_ref_from, state_schema_from
|
||||||
from .spec import NodeSpec
|
from .spec import NodeSpec
|
||||||
|
|
||||||
@@ -98,7 +101,7 @@ class WorkflowBuilder:
|
|||||||
input_schema: SchemaLike
|
input_schema: SchemaLike
|
||||||
state_schema: StateSchemaLike
|
state_schema: StateSchemaLike
|
||||||
output_schema: SchemaLike
|
output_schema: SchemaLike
|
||||||
start: str
|
start: str | None = None
|
||||||
node_specs: dict[str, NodeSpec[Any, Any]] = field(default_factory=dict)
|
node_specs: dict[str, NodeSpec[Any, Any]] = field(default_factory=dict)
|
||||||
nodes: list[Any] = field(default_factory=list)
|
nodes: list[Any] = field(default_factory=list)
|
||||||
edges: list[Edge] = field(default_factory=list)
|
edges: list[Edge] = field(default_factory=list)
|
||||||
@@ -154,6 +157,14 @@ class WorkflowBuilder:
|
|||||||
suffix += 1
|
suffix += 1
|
||||||
return f"{base}_{suffix}"
|
return f"{base}_{suffix}"
|
||||||
|
|
||||||
|
def set_entry_point(self, step: StepRef) -> None:
|
||||||
|
"""Set the workflow start node explicitly."""
|
||||||
|
self.start = _step_id(step)
|
||||||
|
|
||||||
|
def registry(self) -> dict[str, SyncRegistryHandler]:
|
||||||
|
"""Export handlers for all node specs used by this builder."""
|
||||||
|
return build_registry(*self.node_specs.values())
|
||||||
|
|
||||||
def condition(self, *, id: str, check: CoreCondition | Expr) -> ConditionNode:
|
def condition(self, *, id: str, check: CoreCondition | Expr) -> ConditionNode:
|
||||||
node = ConditionNode(
|
node = ConditionNode(
|
||||||
id=id,
|
id=id,
|
||||||
@@ -213,6 +224,11 @@ class WorkflowBuilder:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def compile(self) -> Workflow:
|
def compile(self) -> Workflow:
|
||||||
|
if self.start is None:
|
||||||
|
raise WorkflowExecutionError(
|
||||||
|
"workflow builder requires an explicit start; call set_entry_point(...) "
|
||||||
|
"or pass start=..."
|
||||||
|
)
|
||||||
node_defs = [spec.to_node_def() for spec in self.node_specs.values()]
|
node_defs = [spec.to_node_def() for spec in self.node_specs.values()]
|
||||||
return Workflow(
|
return Workflow(
|
||||||
name=self.name,
|
name=self.name,
|
||||||
|
|||||||
@@ -4,8 +4,10 @@ from typing import Annotated, TypedDict
|
|||||||
|
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from wf_authoring import WorkflowBuilder, build_registry, node, state_field
|
import pytest
|
||||||
from wf_core import RunStatus, execute_workflow
|
|
||||||
|
from wf_authoring import WorkflowBuilder, node, state_field
|
||||||
|
from wf_core import RunStatus, WorkflowExecutionError, execute_workflow
|
||||||
|
|
||||||
|
|
||||||
class WorkflowInput(BaseModel):
|
class WorkflowInput(BaseModel):
|
||||||
@@ -138,7 +140,7 @@ def test_builder_auto_binds_matching_node_inputs_and_outputs_to_state() -> None:
|
|||||||
run = execute_workflow(
|
run = execute_workflow(
|
||||||
workflow,
|
workflow,
|
||||||
{"text": "hello", "count": 1},
|
{"text": "hello", "count": 1},
|
||||||
build_registry(auto_bind_node),
|
builder.registry(),
|
||||||
)
|
)
|
||||||
|
|
||||||
assert step.in_map == {
|
assert step.in_map == {
|
||||||
@@ -168,3 +170,43 @@ def test_builder_can_auto_id_node_uses_from_spec_name() -> None:
|
|||||||
|
|
||||||
assert first.id == "test_auto_bind"
|
assert first.id == "test_auto_bind"
|
||||||
assert second.id == "test_auto_bind_2"
|
assert second.id == "test_auto_bind_2"
|
||||||
|
|
||||||
|
|
||||||
|
def test_builder_can_compile_with_explicit_start_set_later() -> None:
|
||||||
|
builder = WorkflowBuilder(
|
||||||
|
name="optional_start_demo",
|
||||||
|
input_schema=AutoBindInput,
|
||||||
|
state_schema=AutoBindState,
|
||||||
|
output_schema=AutoBindOutput,
|
||||||
|
)
|
||||||
|
step = builder.use(auto_bind_node)
|
||||||
|
builder.set_entry_point(step)
|
||||||
|
|
||||||
|
workflow = builder.compile()
|
||||||
|
|
||||||
|
assert workflow.start == "test_auto_bind"
|
||||||
|
|
||||||
|
|
||||||
|
def test_builder_requires_explicit_start_before_compile() -> None:
|
||||||
|
builder = WorkflowBuilder(
|
||||||
|
name="missing_start_demo",
|
||||||
|
input_schema=AutoBindInput,
|
||||||
|
state_schema=AutoBindState,
|
||||||
|
output_schema=AutoBindOutput,
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(WorkflowExecutionError, match="start"):
|
||||||
|
builder.compile()
|
||||||
|
|
||||||
|
|
||||||
|
def test_builder_registry_exports_used_node_specs() -> None:
|
||||||
|
builder = WorkflowBuilder(
|
||||||
|
name="registry_demo",
|
||||||
|
input_schema=AutoBindInput,
|
||||||
|
state_schema=AutoBindState,
|
||||||
|
output_schema=AutoBindOutput,
|
||||||
|
start="test_auto_bind",
|
||||||
|
)
|
||||||
|
builder.use(auto_bind_node)
|
||||||
|
|
||||||
|
assert set(builder.registry()) == {"test.auto_bind"}
|
||||||
|
|||||||
Reference in New Issue
Block a user