example concurrent foreach

This commit is contained in:
lda
2026-05-23 21:57:50 +07:00 Verified
parent e09d714c72
commit 88b17f6806
4 changed files with 204 additions and 133 deletions
+6
View File
@@ -228,6 +228,12 @@ step with item-local child lineages:
In async execution, admitted async item node handlers may run at the same time.
Run-state mutation, tracing, and barrier commits remain deterministic.
See `examples/authoring_concurrent_foreach.py` for a runnable example covering:
- sync concurrent foreach with `item_error={"action": "collect", ...}`;
- async item-node batching with deterministic output order;
- `item_error` as a string, mapping, or `ForeachItemErrorPolicy` object.
## Deprecated `route`
`route()` is a compatibility shim:
+167
View File
@@ -0,0 +1,167 @@
from __future__ import annotations
import asyncio
from typing import Annotated, Any
from pydantic import BaseModel, Field
from wf_authoring import (
NodeSpec,
WorkflowBuilder,
build_async_registry,
context_path,
input_from,
node,
output_to,
state_field,
state_path,
)
from wf_core import (
END,
ForeachConcurrentPolicy,
ForeachItemErrorPolicy,
execute_workflow_async,
)
from wf_core.paths import StatePath
from wf_core.run_state import RunState
class ItemsInput(BaseModel):
"""Workflow input containing items to process."""
items: list[str]
class ConcurrentForeachState(BaseModel):
"""State shape used by the concurrent foreach authoring example."""
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):
"""Workflow output showing successful items and collected item failures."""
seen: list[str]
errors: list[dict[str, object]]
class RecordInput(BaseModel):
"""Input for one foreach item node call."""
value: str
seen: str
class RecordOutput(BaseModel):
"""Output appended to workflow state at the foreach barrier."""
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 build_concurrent_foreach_workflow(
spec: NodeSpec[Any, RecordOutput] = record_item,
*,
item_error: ForeachItemErrorPolicy | dict[str, object] | str | None = None,
concurrent: ForeachConcurrentPolicy | dict[str, object] | None = None,
) -> WorkflowBuilder:
"""Build a public authoring workflow that uses concurrent foreach.
`item_error` accepts the same canonical forms as `WorkflowBuilder.foreach`:
a bare action string, a mapping, or a `ForeachItemErrorPolicy` object.
"""
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=concurrent or {"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 run_collected_errors_example() -> RunState:
"""Run the sync example with one failing item collected into state.errors."""
builder = build_concurrent_foreach_workflow(
record_item,
item_error=ForeachItemErrorPolicy(
action="collect",
collect_to=StatePath.of("errors"),
),
)
return builder.execute({"items": ["a", "bad", "c"]})
async def run_async_ordered_example() -> RunState:
"""Run the async example; barrier commits still preserve item order."""
builder = build_concurrent_foreach_workflow(record_item_async)
return await execute_workflow_async(
builder.compile(),
{"items": ["a", "b", "c"]},
build_async_registry(record_item_async),
)
def _item_error_action(
item_error: ForeachItemErrorPolicy | dict[str, object] | str | None,
) -> object:
"""Return the policy action without forcing callers into one input shape."""
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
def main() -> None:
"""Run the example directly from the command line."""
sync_run = run_collected_errors_example()
async_run = asyncio.run(run_async_ordered_example())
print("collected_errors", sync_run.output)
print("async_ordered", async_run.output)
if __name__ == "__main__":
main()
+4 -3
View File
@@ -10,6 +10,7 @@ from wf_authoring.ops.values import runtime_error
from wf_core import (
ConditionNode,
Edge,
ForeachConcurrentPolicy,
ForeachItemErrorPolicy,
ForeachNode,
InterruptNode,
@@ -420,7 +421,7 @@ class WorkflowBuilder:
as_: str,
mode: Literal["serial", "concurrent"] = "serial",
item_error: ForeachItemErrorPolicy | Mapping[str, object] | str | None = None,
concurrent: Mapping[str, object] | None = None,
concurrent: ForeachConcurrentPolicy | Mapping[str, object] | None = None,
) -> ForeachNode: ...
@overload
@@ -433,7 +434,7 @@ class WorkflowBuilder:
as_: str,
mode: Literal["serial", "concurrent"] = "serial",
on_item_error: Literal["fail", "collect", "skip"] = "fail",
concurrent: Mapping[str, object] | None = None,
concurrent: ForeachConcurrentPolicy | Mapping[str, object] | None = None,
) -> ForeachNode: ...
def foreach(
@@ -445,7 +446,7 @@ class WorkflowBuilder:
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,
concurrent: ForeachConcurrentPolicy | Mapping[str, object] | None = None,
) -> ForeachNode:
"""Add a foreach step.
@@ -1,82 +1,26 @@
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 examples.authoring_concurrent_foreach import (
ConcurrentForeachOutput,
ConcurrentForeachState,
ItemsInput,
build_concurrent_foreach_workflow,
record_item,
run_async_ordered_example,
run_collected_errors_example,
)
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)
from wf_authoring import WorkflowBuilder, state_path
from wf_core import ForeachConcurrentPolicy
from wf_core import RunStatus
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"]})
run = run_collected_errors_example()
assert run.status == RunStatus.COMPLETED
assert run.output["seen"] == ["a", "c"]
@@ -90,23 +34,14 @@ def test_authoring_concurrent_foreach_collects_item_errors() -> None:
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,
)
)
run = asyncio.run(run_async_ordered_example())
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(
builder = build_concurrent_foreach_workflow(
record_item,
item_error={"action": "collect", "collect_to": state_path("errors")},
)
@@ -120,13 +55,25 @@ def test_authoring_foreach_accepts_item_error_mapping_with_authoring_path() -> N
def test_authoring_foreach_accepts_item_error_action_string() -> None:
builder = _concurrent_foreach_builder(record_item, item_error="skip")
builder = build_concurrent_foreach_workflow(record_item, item_error="skip")
foreach = builder.compile().nodes[0]
assert foreach.model_dump(mode="json")["item_error"]["action"] == "skip"
def test_authoring_foreach_accepts_concurrent_policy_object() -> None:
builder = build_concurrent_foreach_workflow(
record_item,
concurrent=ForeachConcurrentPolicy(max_active=1, max_outstanding=3),
)
foreach = builder.compile().nodes[0]
assert foreach.model_dump(mode="json")["concurrent"]["max_active"] == 1
assert foreach.model_dump(mode="json")["concurrent"]["max_outstanding"] == 3
def test_authoring_foreach_deprecated_on_item_error_warns() -> None:
builder = WorkflowBuilder(
name="deprecated_item_error",
@@ -162,53 +109,3 @@ def test_authoring_foreach_rejects_mixed_item_error_styles() -> None:
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