expose itempolicy to the wf_authoring
This commit is contained in:
@@ -9,7 +9,7 @@ more later
|
|||||||
|
|
||||||
```bash
|
```bash
|
||||||
uv run /* --env-file .env */ pytest -q
|
uv run /* --env-file .env */ pytest -q
|
||||||
(uv run / uvx) ruff check / format
|
(uvx / uv run) ruff check / format
|
||||||
uv run basedpyright --level error # error to cut spam
|
uv run basedpyright --level error # error to cut spam
|
||||||
# maybe uvx ty
|
# maybe uvx ty
|
||||||
```
|
```
|
||||||
|
|||||||
@@ -185,6 +185,49 @@ creates a fresh `use()` step with auto-mapping and an auto id.
|
|||||||
Use existing step refs when the same node use should be shared. Pass a
|
Use existing step refs when the same node use should be shared. Pass a
|
||||||
`NodeSpec` when you want a new use at that point in the graph.
|
`NodeSpec` when you want a new use at that point in the graph.
|
||||||
|
|
||||||
|
## Concurrent `foreach`
|
||||||
|
|
||||||
|
Use `foreach(mode="concurrent")` when item lineages may make progress
|
||||||
|
independently but should still commit their state writes at one deterministic
|
||||||
|
barrier.
|
||||||
|
|
||||||
|
```python
|
||||||
|
each = g.foreach(
|
||||||
|
id="each",
|
||||||
|
over=state_path("items"),
|
||||||
|
as_="item",
|
||||||
|
mode="concurrent",
|
||||||
|
concurrent={"max_active": 2, "max_outstanding": 2},
|
||||||
|
item_error={
|
||||||
|
"action": "collect",
|
||||||
|
"collect_to": state_path("errors"),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
`item_error` is the canonical policy field. It accepts:
|
||||||
|
|
||||||
|
- `"fail"` or `"skip"` when no extra policy fields are needed;
|
||||||
|
- a mapping when fields such as `collect_to` are needed;
|
||||||
|
- the core `ForeachItemErrorPolicy` object.
|
||||||
|
|
||||||
|
`item_error="collect"` is intentionally incomplete and fails validation because
|
||||||
|
`collect` must say where error records should be written. `on_item_error` is
|
||||||
|
deprecated compatibility shorthand and should not appear in new examples.
|
||||||
|
|
||||||
|
Concurrent foreach is not a general fork/gather node. It is still one foreach
|
||||||
|
step with item-local child lineages:
|
||||||
|
|
||||||
|
- each item sees its own buffered writes while it runs;
|
||||||
|
- sibling item writes do not leak into each other before the barrier;
|
||||||
|
- final barrier commits happen in item-index order;
|
||||||
|
- same-path sibling writes require a mergeable reducer on that exact state path;
|
||||||
|
- `item_error.action="collect"` requires `collect_to` to point at a declared
|
||||||
|
array state field.
|
||||||
|
|
||||||
|
In async execution, admitted async item node handlers may run at the same time.
|
||||||
|
Run-state mutation, tracing, and barrier commits remain deterministic.
|
||||||
|
|
||||||
## Deprecated `route`
|
## Deprecated `route`
|
||||||
|
|
||||||
`route()` is a compatibility shim:
|
`route()` is a compatibility shim:
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ from wf_authoring.ops.values import runtime_error
|
|||||||
from wf_core import (
|
from wf_core import (
|
||||||
ConditionNode,
|
ConditionNode,
|
||||||
Edge,
|
Edge,
|
||||||
|
ForeachItemErrorPolicy,
|
||||||
ForeachNode,
|
ForeachNode,
|
||||||
InterruptNode,
|
InterruptNode,
|
||||||
NodeUse,
|
NodeUse,
|
||||||
@@ -32,7 +33,7 @@ from wf_core.models.steps import (
|
|||||||
from wf_core.paths import GraphSourcePath, LocalPath, StatePath
|
from wf_core.paths import GraphSourcePath, LocalPath, StatePath
|
||||||
from wf_core.runtime.ops.merges import ReducerDefinition
|
from wf_core.runtime.ops.merges import ReducerDefinition
|
||||||
|
|
||||||
from ..dsl import Expr, PathArg, PathExpr, compile_condition
|
from ..dsl import Expr, GraphPath, PathArg, PathExpr, compile_condition
|
||||||
from ..nodes.callables import SyncRegistryHandler
|
from ..nodes.callables import SyncRegistryHandler
|
||||||
from ..nodes.registry import build_registry
|
from ..nodes.registry import build_registry
|
||||||
from ..reducers import ReducerCatalog
|
from ..reducers import ReducerCatalog
|
||||||
@@ -145,6 +146,28 @@ def _warn_deprecated_binding_sugar(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_foreach_item_error(
|
||||||
|
item_error: ForeachItemErrorPolicy | Mapping[str, object] | str | None,
|
||||||
|
) -> ForeachItemErrorPolicy | dict[str, object] | str | None:
|
||||||
|
"""Coerce authoring path helpers inside the canonical item-error policy."""
|
||||||
|
if item_error is None or isinstance(item_error, ForeachItemErrorPolicy | str):
|
||||||
|
return item_error
|
||||||
|
|
||||||
|
normalized = dict(item_error)
|
||||||
|
collect_to = normalized.get("collect_to")
|
||||||
|
if isinstance(collect_to, PathExpr):
|
||||||
|
collect_to = collect_to.path
|
||||||
|
if isinstance(collect_to, GraphPath):
|
||||||
|
collect_to = collect_to.path
|
||||||
|
if isinstance(collect_to, GraphSourcePath):
|
||||||
|
if collect_to.root != "state":
|
||||||
|
raise ValueError("foreach item_error.collect_to must be a state path")
|
||||||
|
collect_to = StatePath(collect_to.parts)
|
||||||
|
if collect_to is not None:
|
||||||
|
normalized["collect_to"] = collect_to
|
||||||
|
return normalized
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
@dataclass(slots=True)
|
||||||
class WorkflowBuilder:
|
class WorkflowBuilder:
|
||||||
name: str
|
name: str
|
||||||
@@ -388,6 +411,20 @@ class WorkflowBuilder:
|
|||||||
self.nodes.append(node)
|
self.nodes.append(node)
|
||||||
return node
|
return node
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def foreach(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
id: str | None = None,
|
||||||
|
over: PathArg,
|
||||||
|
as_: str,
|
||||||
|
mode: Literal["serial", "concurrent"] = "serial",
|
||||||
|
item_error: ForeachItemErrorPolicy | Mapping[str, object] | str | None = None,
|
||||||
|
concurrent: Mapping[str, object] | None = None,
|
||||||
|
) -> ForeachNode: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
@deprecated("use item_error canonical policy instead")
|
||||||
def foreach(
|
def foreach(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -397,13 +434,36 @@ class WorkflowBuilder:
|
|||||||
mode: Literal["serial", "concurrent"] = "serial",
|
mode: Literal["serial", "concurrent"] = "serial",
|
||||||
on_item_error: Literal["fail", "collect", "skip"] = "fail",
|
on_item_error: Literal["fail", "collect", "skip"] = "fail",
|
||||||
concurrent: Mapping[str, object] | None = None,
|
concurrent: Mapping[str, object] | None = None,
|
||||||
|
) -> ForeachNode: ...
|
||||||
|
|
||||||
|
def foreach(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
id: str | None = None,
|
||||||
|
over: PathArg,
|
||||||
|
as_: str,
|
||||||
|
mode: Literal["serial", "concurrent"] = "serial",
|
||||||
|
item_error: ForeachItemErrorPolicy | Mapping[str, object] | str | None = None,
|
||||||
|
on_item_error: Literal["fail", "collect", "skip"] = "fail",
|
||||||
|
concurrent: Mapping[str, object] | None = None,
|
||||||
) -> ForeachNode:
|
) -> ForeachNode:
|
||||||
"""Add a foreach step.
|
"""Add a foreach step.
|
||||||
|
|
||||||
Concurrent mode is supported by the runtime with deterministic barrier
|
Concurrent mode is supported by the runtime with deterministic barrier
|
||||||
commits, item error policies, and async item-node batching. See ADR 0002
|
commits, item error policies, and async item-node batching. See ADR 0002
|
||||||
for the exact merge and interrupt semantics.
|
for the exact merge and interrupt semantics. Prefer the canonical
|
||||||
|
`item_error` policy object/mapping; `on_item_error` is compatibility
|
||||||
|
shorthand for older callers.
|
||||||
"""
|
"""
|
||||||
|
if item_error is not None and on_item_error != "fail":
|
||||||
|
raise TypeError("cannot mix item_error with deprecated on_item_error")
|
||||||
|
if item_error is None and on_item_error != "fail":
|
||||||
|
warnings.warn(
|
||||||
|
"on_item_error is deprecated WorkflowBuilder foreach sugar; "
|
||||||
|
"use item_error={'action': ...} instead",
|
||||||
|
DeprecationWarning,
|
||||||
|
stacklevel=2,
|
||||||
|
)
|
||||||
node = ForeachNode.model_validate(
|
node = ForeachNode.model_validate(
|
||||||
{
|
{
|
||||||
"id": id or self._next_step_id(f"foreach_{slug_id(as_)}"),
|
"id": id or self._next_step_id(f"foreach_{slug_id(as_)}"),
|
||||||
@@ -411,7 +471,9 @@ class WorkflowBuilder:
|
|||||||
"over": coerce_path(over),
|
"over": coerce_path(over),
|
||||||
"as": as_,
|
"as": as_,
|
||||||
"mode": mode,
|
"mode": mode,
|
||||||
"on_item_error": on_item_error,
|
"item_error": _normalize_foreach_item_error(item_error)
|
||||||
|
if item_error is not None
|
||||||
|
else {"action": on_item_error},
|
||||||
"concurrent": concurrent,
|
"concurrent": concurrent,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -171,6 +171,14 @@ class ForeachItemErrorPolicy(BaseModel):
|
|||||||
action: Literal["fail", "skip", "collect"] = "fail"
|
action: Literal["fail", "skip", "collect"] = "fail"
|
||||||
collect_to: StatePath | None = None
|
collect_to: StatePath | None = None
|
||||||
|
|
||||||
|
@model_validator(mode="before")
|
||||||
|
@classmethod
|
||||||
|
def _coerce_action_string(cls, data: object) -> object:
|
||||||
|
"""Accept bare action strings for policies with no extra fields."""
|
||||||
|
if isinstance(data, str):
|
||||||
|
return {"action": data}
|
||||||
|
return data
|
||||||
|
|
||||||
@model_validator(mode="after")
|
@model_validator(mode="after")
|
||||||
def _validate_collect_to(self) -> Self:
|
def _validate_collect_to(self) -> Self:
|
||||||
if self.action == "collect" and self.collect_to is None:
|
if self.action == "collect" and self.collect_to is None:
|
||||||
|
|||||||
@@ -0,0 +1,214 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from typing import Annotated
|
||||||
|
from typing import Any, cast
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
from wf_authoring import (
|
||||||
|
WorkflowBuilder,
|
||||||
|
build_async_registry,
|
||||||
|
context_path,
|
||||||
|
input_from,
|
||||||
|
node,
|
||||||
|
output_to,
|
||||||
|
state_field,
|
||||||
|
state_path,
|
||||||
|
)
|
||||||
|
from wf_core import (
|
||||||
|
END,
|
||||||
|
ForeachItemErrorPolicy,
|
||||||
|
RunStatus,
|
||||||
|
execute_workflow_async,
|
||||||
|
)
|
||||||
|
from wf_core.paths import StatePath
|
||||||
|
|
||||||
|
|
||||||
|
class ItemsInput(BaseModel):
|
||||||
|
items: list[str]
|
||||||
|
|
||||||
|
|
||||||
|
class ConcurrentForeachState(BaseModel):
|
||||||
|
items: list[str]
|
||||||
|
seen: Annotated[list[str], state_field(reducer="wf.std.append")] = Field(
|
||||||
|
default_factory=list
|
||||||
|
)
|
||||||
|
errors: list[dict[str, object]] = Field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
|
class ConcurrentForeachOutput(BaseModel):
|
||||||
|
seen: list[str]
|
||||||
|
errors: list[dict[str, object]]
|
||||||
|
|
||||||
|
|
||||||
|
class RecordInput(BaseModel):
|
||||||
|
value: str
|
||||||
|
seen: str
|
||||||
|
|
||||||
|
|
||||||
|
class RecordOutput(BaseModel):
|
||||||
|
seen: str
|
||||||
|
|
||||||
|
|
||||||
|
@node(name="example.record_item")
|
||||||
|
def record_item(payload: RecordInput) -> RecordOutput:
|
||||||
|
"""Record one foreach item, failing on a sentinel item for examples."""
|
||||||
|
if payload.value == "bad":
|
||||||
|
raise ValueError("bad item")
|
||||||
|
return RecordOutput(seen=payload.seen)
|
||||||
|
|
||||||
|
|
||||||
|
@node(name="example.record_item_async")
|
||||||
|
async def record_item_async(payload: RecordInput) -> RecordOutput:
|
||||||
|
"""Async variant used to prove authoring workflows can use async batching."""
|
||||||
|
await asyncio.sleep({"a": 0.03, "b": 0.01, "c": 0.02}[payload.value])
|
||||||
|
return RecordOutput(seen=payload.seen)
|
||||||
|
|
||||||
|
|
||||||
|
def test_authoring_concurrent_foreach_collects_item_errors() -> None:
|
||||||
|
builder = _concurrent_foreach_builder(
|
||||||
|
record_item,
|
||||||
|
item_error=ForeachItemErrorPolicy(
|
||||||
|
action="collect",
|
||||||
|
collect_to=StatePath.of("errors"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
run = builder.execute({"items": ["a", "bad", "c"]})
|
||||||
|
|
||||||
|
assert run.status == RunStatus.COMPLETED
|
||||||
|
assert run.output["seen"] == ["a", "c"]
|
||||||
|
assert len(run.output["errors"]) == 1
|
||||||
|
error = run.output["errors"][0]
|
||||||
|
assert error["index"] == 1
|
||||||
|
assert error["node_id"] == "record"
|
||||||
|
assert error["error_type"] == "ValueError"
|
||||||
|
assert error["message"] == "bad item"
|
||||||
|
assert error["item"] == "bad"
|
||||||
|
|
||||||
|
|
||||||
|
def test_authoring_async_concurrent_foreach_commits_in_item_order() -> None:
|
||||||
|
builder = _concurrent_foreach_builder(record_item_async)
|
||||||
|
registry = build_async_registry(record_item_async)
|
||||||
|
|
||||||
|
run = asyncio.run(
|
||||||
|
execute_workflow_async(
|
||||||
|
builder.compile(),
|
||||||
|
{"items": ["a", "b", "c"]},
|
||||||
|
registry,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert run.status == RunStatus.COMPLETED
|
||||||
|
assert run.output["seen"] == ["a", "b", "c"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_authoring_foreach_accepts_item_error_mapping_with_authoring_path() -> None:
|
||||||
|
builder = _concurrent_foreach_builder(
|
||||||
|
record_item,
|
||||||
|
item_error={"action": "collect", "collect_to": state_path("errors")},
|
||||||
|
)
|
||||||
|
|
||||||
|
foreach = builder.compile().nodes[0]
|
||||||
|
|
||||||
|
assert foreach.model_dump(mode="json")["item_error"]["collect_to"] == {
|
||||||
|
"root": "state",
|
||||||
|
"parts": ["errors"],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_authoring_foreach_accepts_item_error_action_string() -> None:
|
||||||
|
builder = _concurrent_foreach_builder(record_item, item_error="skip")
|
||||||
|
|
||||||
|
foreach = builder.compile().nodes[0]
|
||||||
|
|
||||||
|
assert foreach.model_dump(mode="json")["item_error"]["action"] == "skip"
|
||||||
|
|
||||||
|
|
||||||
|
def test_authoring_foreach_deprecated_on_item_error_warns() -> None:
|
||||||
|
builder = WorkflowBuilder(
|
||||||
|
name="deprecated_item_error",
|
||||||
|
input_schema=ItemsInput,
|
||||||
|
state_schema=ConcurrentForeachState,
|
||||||
|
output_schema=ConcurrentForeachOutput,
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.warns(DeprecationWarning, match="on_item_error"):
|
||||||
|
foreach = builder.foreach(
|
||||||
|
id="each",
|
||||||
|
over=state_path("items"),
|
||||||
|
as_="item",
|
||||||
|
on_item_error="skip",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert foreach.item_error.action == "skip"
|
||||||
|
|
||||||
|
|
||||||
|
def test_authoring_foreach_rejects_mixed_item_error_styles() -> None:
|
||||||
|
builder = WorkflowBuilder(
|
||||||
|
name="mixed_item_error",
|
||||||
|
input_schema=ItemsInput,
|
||||||
|
state_schema=ConcurrentForeachState,
|
||||||
|
output_schema=ConcurrentForeachOutput,
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(TypeError, match="cannot mix item_error"):
|
||||||
|
cast(Any, builder.foreach)(
|
||||||
|
id="each",
|
||||||
|
over=state_path("items"),
|
||||||
|
as_="item",
|
||||||
|
item_error="skip",
|
||||||
|
on_item_error="collect",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _concurrent_foreach_builder(
|
||||||
|
spec,
|
||||||
|
*,
|
||||||
|
item_error: ForeachItemErrorPolicy | dict[str, object] | str | None = None,
|
||||||
|
) -> WorkflowBuilder:
|
||||||
|
"""Build the public authoring shape for concurrent foreach examples."""
|
||||||
|
builder = WorkflowBuilder(
|
||||||
|
name="authoring_concurrent_foreach",
|
||||||
|
input_schema=ItemsInput,
|
||||||
|
state_schema=ConcurrentForeachState,
|
||||||
|
output_schema=ConcurrentForeachOutput,
|
||||||
|
)
|
||||||
|
each = builder.foreach(
|
||||||
|
id="each",
|
||||||
|
over=state_path("items"),
|
||||||
|
as_="item",
|
||||||
|
mode="concurrent",
|
||||||
|
item_error=item_error,
|
||||||
|
concurrent={"max_active": 2, "max_outstanding": 2},
|
||||||
|
)
|
||||||
|
record = builder.use(
|
||||||
|
spec,
|
||||||
|
id="record",
|
||||||
|
input=[
|
||||||
|
input_from(context_path("item"), "value"),
|
||||||
|
input_from(context_path("item"), "seen"),
|
||||||
|
],
|
||||||
|
output=[output_to("seen", state_path("seen"))],
|
||||||
|
)
|
||||||
|
builder.set_entry_point(each)
|
||||||
|
builder.connect(each, "loop", record)
|
||||||
|
builder.connect(record, "ok", END)
|
||||||
|
builder.connect(each, "done", END)
|
||||||
|
if _item_error_action(item_error) in {"collect", "skip"}:
|
||||||
|
builder.connect(each, "completed_with_errors", END)
|
||||||
|
return builder
|
||||||
|
|
||||||
|
|
||||||
|
def _item_error_action(
|
||||||
|
item_error: ForeachItemErrorPolicy | dict[str, object] | str | None,
|
||||||
|
) -> object:
|
||||||
|
if isinstance(item_error, ForeachItemErrorPolicy):
|
||||||
|
return item_error.action
|
||||||
|
if isinstance(item_error, dict):
|
||||||
|
return item_error.get("action")
|
||||||
|
if isinstance(item_error, str):
|
||||||
|
return item_error
|
||||||
|
return None
|
||||||
@@ -41,6 +41,36 @@ def test_deprecated_on_item_error_parses_to_nested_policy() -> None:
|
|||||||
assert dumped["item_error"]["action"] == "skip"
|
assert dumped["item_error"]["action"] == "skip"
|
||||||
|
|
||||||
|
|
||||||
|
def test_item_error_string_parses_to_policy_action() -> None:
|
||||||
|
node = ForeachNode.model_validate(
|
||||||
|
{
|
||||||
|
"id": "each",
|
||||||
|
"type": "foreach",
|
||||||
|
"over": "state.items",
|
||||||
|
"as": "item",
|
||||||
|
"item_error": "skip",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
dumped = node.model_dump(mode="json", by_alias=True)
|
||||||
|
|
||||||
|
assert node.item_error.action == "skip"
|
||||||
|
assert dumped["item_error"]["action"] == "skip"
|
||||||
|
|
||||||
|
|
||||||
|
def test_collect_item_error_string_explains_required_shape() -> None:
|
||||||
|
with pytest.raises(ValidationError, match="collect_to"):
|
||||||
|
ForeachNode.model_validate(
|
||||||
|
{
|
||||||
|
"id": "each",
|
||||||
|
"type": "foreach",
|
||||||
|
"over": "state.items",
|
||||||
|
"as": "item",
|
||||||
|
"item_error": "collect",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_collect_item_policy_requires_collect_to() -> None:
|
def test_collect_item_policy_requires_collect_to() -> None:
|
||||||
with pytest.raises(ValidationError, match="collect_to"):
|
with pytest.raises(ValidationError, match="collect_to"):
|
||||||
ForeachNode.model_validate(
|
ForeachNode.model_validate(
|
||||||
|
|||||||
Reference in New Issue
Block a user