feat: reserve async workflow step attempts
This commit is contained in:
@@ -16,7 +16,7 @@ from wf_core.models.steps import (
|
|||||||
from wf_core.models.workflow import Workflow
|
from wf_core.models.workflow import Workflow
|
||||||
from wf_core.run_state import ExecutionFrame, FrameStatus, RunState, StepExecutionResult
|
from wf_core.run_state import ExecutionFrame, FrameStatus, RunState, StepExecutionResult
|
||||||
from wf_core.runtime.foreach_state import item_frame_owner, load_foreach_activation
|
from wf_core.runtime.foreach_state import item_frame_owner, load_foreach_activation
|
||||||
from wf_core.runtime.limits import admit_step_attempt
|
from wf_core.runtime.limits import admit_step_attempt, remaining_step_attempts
|
||||||
from wf_core.runtime.ops.flow import advance_frame, append_step_result_trace
|
from wf_core.runtime.ops.flow import advance_frame, append_step_result_trace
|
||||||
from wf_core.runtime.ops.foreach import step_foreach
|
from wf_core.runtime.ops.foreach import step_foreach
|
||||||
from wf_core.runtime.ops.handlers import (
|
from wf_core.runtime.ops.handlers import (
|
||||||
@@ -322,12 +322,25 @@ async def _step_async_foreach_item_batch(
|
|||||||
advancement happen afterward in ready-queue order so `RunState` is mutated
|
advancement happen afterward in ready-queue order so `RunState` is mutated
|
||||||
deterministically.
|
deterministically.
|
||||||
|
|
||||||
Task 3 will bound the claimed siblings by the remaining budget and pin the
|
Batch reservation is bounded by the remaining step budget: the batch claims
|
||||||
reservation/failure semantics. Until then every frame in the batch is
|
at most ``remaining`` frames (``first_frame`` plus up to ``remaining - 1``
|
||||||
admitted in ready-queue order before any handler starts, so each trace has
|
siblings in ready-queue order) and admits every selected frame before
|
||||||
a number; a denied frame raises before any handler in the batch runs.
|
creating any handler coroutine. A denied admission raises before any
|
||||||
|
handler runs, so unclaimed siblings stay PENDING in their original queue
|
||||||
|
order and admitted attempts stay consumed even if a handler later fails.
|
||||||
|
All handlers are awaited via ``gather(return_exceptions=True)`` and then
|
||||||
|
finalized in reservation order; the first unhandled failure keeps preceding
|
||||||
|
commits and discards later sibling state/trace commits.
|
||||||
"""
|
"""
|
||||||
frames = [first_frame, *_claim_matching_async_item_frames(run, index, first_frame)]
|
# Snapshot the remainder once so the claim bound and the admissions agree.
|
||||||
|
# `first_frame` already occupies the ready-queue head, so siblings are
|
||||||
|
# limited to `remaining - 1`; when remaining is 0 the first admission below
|
||||||
|
# raises before any handler is created.
|
||||||
|
remaining = remaining_step_attempts(run)
|
||||||
|
siblings = _claim_matching_async_item_frames(
|
||||||
|
run, index, first_frame, limit=max(remaining - 1, 0)
|
||||||
|
)
|
||||||
|
frames = [first_frame, *siblings]
|
||||||
for frame in frames:
|
for frame in frames:
|
||||||
admit_step_attempt(run, frame, frame.node_id)
|
admit_step_attempt(run, frame, frame.node_id)
|
||||||
tasks = []
|
tasks = []
|
||||||
@@ -375,12 +388,18 @@ def _claim_matching_async_item_frames(
|
|||||||
run: RunState,
|
run: RunState,
|
||||||
index: WorkflowIndex,
|
index: WorkflowIndex,
|
||||||
first_frame: ExecutionFrame,
|
first_frame: ExecutionFrame,
|
||||||
|
limit: int,
|
||||||
) -> list[ExecutionFrame]:
|
) -> list[ExecutionFrame]:
|
||||||
"""Claim sibling item frames from the same activation for async batching.
|
"""Claim sibling item frames from the same activation for async batching.
|
||||||
|
|
||||||
Batching never mixes activations: only frames naming the same parent,
|
Batching never mixes activations: only frames naming the same parent,
|
||||||
foreach, and activation id run together, preserving deterministic barrier
|
foreach, and activation id run together, preserving deterministic barrier
|
||||||
commits across revisits.
|
commits across revisits.
|
||||||
|
|
||||||
|
At most ``limit`` matching frames are claimed in ready-queue order. Only
|
||||||
|
claimed frames flip to RUNNING; the rest stay PENDING and keep their
|
||||||
|
original relative order in ``ready_frame_ids`` so a later dispatch can
|
||||||
|
admit them once budget allows.
|
||||||
"""
|
"""
|
||||||
owner = item_frame_owner(first_frame)
|
owner = item_frame_owner(first_frame)
|
||||||
if owner is None:
|
if owner is None:
|
||||||
@@ -397,6 +416,7 @@ def _claim_matching_async_item_frames(
|
|||||||
and frame_owner.foreach_node_id == owner.foreach_node_id
|
and frame_owner.foreach_node_id == owner.foreach_node_id
|
||||||
and frame_owner.activation_id == owner.activation_id
|
and frame_owner.activation_id == owner.activation_id
|
||||||
and isinstance(index.nodes_by_id.get(frame.node_id), NodeUse)
|
and isinstance(index.nodes_by_id.get(frame.node_id), NodeUse)
|
||||||
|
and len(claimed) < limit
|
||||||
):
|
):
|
||||||
frame.status = FrameStatus.RUNNING
|
frame.status = FrameStatus.RUNNING
|
||||||
claimed.append(frame)
|
claimed.append(frame)
|
||||||
|
|||||||
@@ -0,0 +1,456 @@
|
|||||||
|
"""Async step-budget reservation tests (Task 3).
|
||||||
|
|
||||||
|
Covers deterministic async batch reservation: bound claims by remaining
|
||||||
|
budget, admit in ready-queue order before launching handlers, keep
|
||||||
|
reservations after failure, settle siblings before raising, and discard
|
||||||
|
later sibling commits after the first unhandled result in reservation order.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from wf_core import (
|
||||||
|
END,
|
||||||
|
Edge,
|
||||||
|
ForeachNode,
|
||||||
|
FrameStatus,
|
||||||
|
NodeDef,
|
||||||
|
NodeUse,
|
||||||
|
ReducerRef,
|
||||||
|
RunStatus,
|
||||||
|
SchemaRef,
|
||||||
|
StateField,
|
||||||
|
StateSchema,
|
||||||
|
Workflow,
|
||||||
|
execute_workflow,
|
||||||
|
execute_workflow_async,
|
||||||
|
step_workflow_async,
|
||||||
|
)
|
||||||
|
from wf_core.errors import WorkflowStepLimitExceeded
|
||||||
|
from wf_core.runtime.limits import RunLimits, remaining_step_attempts
|
||||||
|
from wf_core.runtime.ops.runs import create_run_state
|
||||||
|
from wf_core.runtime.preparation import prepare_new_run, prepare_resume
|
||||||
|
|
||||||
|
|
||||||
|
def _concurrent_workflow(*, max_active: int, name: str = "async_budget") -> Workflow:
|
||||||
|
return Workflow(
|
||||||
|
name=name,
|
||||||
|
input_schema=SchemaRef(type="object", properties={"items": {"type": "array"}}),
|
||||||
|
state_schema=StateSchema.from_field_map(
|
||||||
|
{
|
||||||
|
"items": StateField(type="array"),
|
||||||
|
"seen": StateField(
|
||||||
|
type="array",
|
||||||
|
reducer=ReducerRef(name="wf.std.append"),
|
||||||
|
),
|
||||||
|
}
|
||||||
|
),
|
||||||
|
output_schema=SchemaRef(type="object", properties={"seen": {"type": "array"}}),
|
||||||
|
node_defs=[
|
||||||
|
NodeDef(
|
||||||
|
name="record",
|
||||||
|
input_schema=SchemaRef(
|
||||||
|
type="object",
|
||||||
|
properties={"value": {}, "seen": {}},
|
||||||
|
required=["value", "seen"],
|
||||||
|
),
|
||||||
|
output_schema=SchemaRef(
|
||||||
|
type="object",
|
||||||
|
properties={"value": {}, "seen": {}},
|
||||||
|
required=["seen"],
|
||||||
|
),
|
||||||
|
outcomes=["ok"],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
start="each",
|
||||||
|
nodes=[
|
||||||
|
ForeachNode.model_validate(
|
||||||
|
{
|
||||||
|
"id": "each",
|
||||||
|
"type": "foreach",
|
||||||
|
"over": "state.items",
|
||||||
|
"as": "item",
|
||||||
|
"mode": "concurrent",
|
||||||
|
"concurrent": {
|
||||||
|
"max_active": max_active,
|
||||||
|
"max_outstanding": max_active,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
),
|
||||||
|
NodeUse.model_validate(
|
||||||
|
{
|
||||||
|
"id": "record",
|
||||||
|
"type": "node",
|
||||||
|
"node": "record",
|
||||||
|
"input": [
|
||||||
|
{"target": "value", "path": "context.item"},
|
||||||
|
{"target": "seen", "path": "context.item"},
|
||||||
|
],
|
||||||
|
"output": [{"source": "seen", "target": "state.seen"}],
|
||||||
|
}
|
||||||
|
),
|
||||||
|
],
|
||||||
|
edges=[
|
||||||
|
Edge.model_validate({"from": "each", "outcome": "loop", "to": "record"}),
|
||||||
|
Edge.model_validate({"from": "record", "outcome": "ok", "to": "each"}),
|
||||||
|
Edge.model_validate({"from": "each", "outcome": "done", "to": END}),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _prepare_limited_run(
|
||||||
|
workflow: Workflow, items: list[Any], *, max_steps: int
|
||||||
|
):
|
||||||
|
run = create_run_state(workflow, {"items": items}, limits=RunLimits(max_steps=max_steps))
|
||||||
|
prepare_new_run(workflow, {"items": items}, run)
|
||||||
|
index = prepare_resume(workflow, run, resume_payload=None, resume_outcome="submitted")
|
||||||
|
assert index is not None
|
||||||
|
return run, index
|
||||||
|
|
||||||
|
|
||||||
|
async def _wait_for(predicate, *, timeout: float = 2.0) -> None: # type: ignore[no-untyped-def]
|
||||||
|
async def _poll() -> None:
|
||||||
|
while not predicate():
|
||||||
|
await asyncio.sleep(0.005)
|
||||||
|
|
||||||
|
await asyncio.wait_for(_poll(), timeout)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_async_batch_bounded_to_remaining_budget() -> None:
|
||||||
|
"""A 3-unit remainder starts only the first 3 of 5 eligible frames."""
|
||||||
|
workflow = _concurrent_workflow(max_active=5, name="async_bounded")
|
||||||
|
items = ["a", "b", "c", "d", "e"]
|
||||||
|
run, index = _prepare_limited_run(workflow, items, max_steps=4)
|
||||||
|
|
||||||
|
async def _noop(payload: dict[str, Any], _ctx: object) -> dict[str, Any]:
|
||||||
|
return {"outcome": "ok", "output": payload}
|
||||||
|
|
||||||
|
await step_workflow_async(workflow, run, {"record": _noop}, index=index)
|
||||||
|
assert run.steps_executed == 1
|
||||||
|
assert remaining_step_attempts(run) == 3
|
||||||
|
assert run.ready_frame_ids == [f"root:each#0:{i}" for i in range(5)]
|
||||||
|
|
||||||
|
release = asyncio.Event()
|
||||||
|
started: list[str] = []
|
||||||
|
|
||||||
|
async def record(payload: dict[str, Any], _ctx: object) -> dict[str, Any]:
|
||||||
|
started.append(payload["value"])
|
||||||
|
await release.wait()
|
||||||
|
return {"outcome": "ok", "output": payload}
|
||||||
|
|
||||||
|
batch = asyncio.create_task(
|
||||||
|
step_workflow_async(workflow, run, {"record": record}, index=index)
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
# The first three in ready-queue order start while gated; the other
|
||||||
|
# two must never start because only three units remain.
|
||||||
|
await _wait_for(lambda: len(started) == 3)
|
||||||
|
assert started == ["a", "b", "c"]
|
||||||
|
await asyncio.sleep(0.05)
|
||||||
|
assert started == ["a", "b", "c"]
|
||||||
|
|
||||||
|
# Unclaimed siblings stay PENDING in their original ready-queue order
|
||||||
|
# while the admitted batch is still in flight.
|
||||||
|
assert run.frames["root:each#0:3"].status == FrameStatus.PENDING
|
||||||
|
assert run.frames["root:each#0:4"].status == FrameStatus.PENDING
|
||||||
|
assert run.frames["root:each#0:3"].step_number is None
|
||||||
|
assert run.frames["root:each#0:4"].step_number is None
|
||||||
|
unclaimed = [fid for fid in run.ready_frame_ids if fid.endswith((":3", ":4"))]
|
||||||
|
assert unclaimed == ["root:each#0:3", "root:each#0:4"]
|
||||||
|
finally:
|
||||||
|
release.set()
|
||||||
|
await batch
|
||||||
|
|
||||||
|
# Numbers follow ready-queue order, reservations are consumed.
|
||||||
|
assert run.steps_executed == 4
|
||||||
|
assert remaining_step_attempts(run) == 0
|
||||||
|
assert run.frames["root:each#0:0"].step_number == 2
|
||||||
|
assert run.frames["root:each#0:1"].step_number == 3
|
||||||
|
assert run.frames["root:each#0:2"].step_number == 4
|
||||||
|
assert run.frames["root:each#0:3"].step_number is None
|
||||||
|
assert run.frames["root:each#0:4"].step_number is None
|
||||||
|
assert run.frames["root:each#0:3"].status == FrameStatus.PENDING
|
||||||
|
assert run.frames["root:each#0:4"].status == FrameStatus.PENDING
|
||||||
|
batch_numbers = [
|
||||||
|
entry.step_number
|
||||||
|
for entry in run.trace
|
||||||
|
if entry.frame_id.startswith("root:each#0:")
|
||||||
|
]
|
||||||
|
assert batch_numbers == [2, 3, 4]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_async_batch_denies_first_when_remaining_zero() -> None:
|
||||||
|
"""With no remainder the first admission raises before any handler runs."""
|
||||||
|
workflow = _concurrent_workflow(max_active=5, name="async_zero_remainder")
|
||||||
|
items = ["a", "b", "c"]
|
||||||
|
run, index = _prepare_limited_run(workflow, items, max_steps=1)
|
||||||
|
|
||||||
|
async def _noop(payload: dict[str, Any], _ctx: object) -> dict[str, Any]:
|
||||||
|
return {"outcome": "ok", "output": payload}
|
||||||
|
|
||||||
|
await step_workflow_async(workflow, run, {"record": _noop}, index=index)
|
||||||
|
assert run.steps_executed == 1
|
||||||
|
assert remaining_step_attempts(run) == 0
|
||||||
|
|
||||||
|
started: list[str] = []
|
||||||
|
|
||||||
|
async def record(payload: dict[str, Any], _ctx: object) -> dict[str, Any]:
|
||||||
|
started.append(payload["value"])
|
||||||
|
return {"outcome": "ok", "output": payload}
|
||||||
|
|
||||||
|
with pytest.raises(WorkflowStepLimitExceeded):
|
||||||
|
await step_workflow_async(workflow, run, {"record": record}, index=index)
|
||||||
|
|
||||||
|
assert started == []
|
||||||
|
assert run.steps_executed == 1
|
||||||
|
# Nothing was claimed: the popped first frame stays RUNNING without a
|
||||||
|
# number, every eligible sibling stays PENDING in queue order.
|
||||||
|
assert run.frames["root:each#0:0"].status == FrameStatus.RUNNING
|
||||||
|
assert run.frames["root:each#0:0"].step_number is None
|
||||||
|
assert run.ready_frame_ids == [f"root:each#0:{i}" for i in range(1, 3)]
|
||||||
|
for i in range(1, 3):
|
||||||
|
frame = run.frames[f"root:each#0:{i}"]
|
||||||
|
assert frame.status == FrameStatus.PENDING
|
||||||
|
assert frame.step_number is None
|
||||||
|
|
||||||
|
|
||||||
|
async def test_async_batch_numbers_follow_queue_order_not_completion() -> None:
|
||||||
|
"""Step numbers and trace order follow reservation, not completion order."""
|
||||||
|
workflow = _concurrent_workflow(max_active=3, name="async_queue_order")
|
||||||
|
items = ["a", "b", "c"]
|
||||||
|
run, index = _prepare_limited_run(workflow, items, max_steps=20)
|
||||||
|
|
||||||
|
async def _noop(payload: dict[str, Any], _ctx: object) -> dict[str, Any]:
|
||||||
|
return {"outcome": "ok", "output": payload}
|
||||||
|
|
||||||
|
await step_workflow_async(workflow, run, {"record": _noop}, index=index)
|
||||||
|
assert run.steps_executed == 1
|
||||||
|
|
||||||
|
allow_a = asyncio.Event()
|
||||||
|
started: list[str] = []
|
||||||
|
finished: list[str] = []
|
||||||
|
|
||||||
|
async def record(payload: dict[str, Any], _ctx: object) -> dict[str, Any]:
|
||||||
|
value = payload["value"]
|
||||||
|
started.append(value)
|
||||||
|
if value == "a":
|
||||||
|
# Gate the queue-first item so it finishes last even though it
|
||||||
|
# was admitted first.
|
||||||
|
await allow_a.wait()
|
||||||
|
else:
|
||||||
|
await asyncio.sleep(0.01)
|
||||||
|
finished.append(value)
|
||||||
|
return {"outcome": "ok", "output": payload}
|
||||||
|
|
||||||
|
batch = asyncio.create_task(
|
||||||
|
step_workflow_async(workflow, run, {"record": record}, index=index)
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
await _wait_for(lambda: len(started) == 3)
|
||||||
|
assert started == ["a", "b", "c"]
|
||||||
|
await _wait_for(lambda: len(finished) == 2)
|
||||||
|
assert sorted(finished) == ["b", "c"]
|
||||||
|
assert "a" not in finished
|
||||||
|
finally:
|
||||||
|
allow_a.set()
|
||||||
|
await batch
|
||||||
|
|
||||||
|
assert finished[-1] == "a"
|
||||||
|
assert finished != ["a", "b", "c"]
|
||||||
|
# Reservation order still drives numbering and finalization order.
|
||||||
|
assert run.frames["root:each#0:0"].step_number == 2
|
||||||
|
assert run.frames["root:each#0:1"].step_number == 3
|
||||||
|
assert run.frames["root:each#0:2"].step_number == 4
|
||||||
|
batch_entries = [
|
||||||
|
entry for entry in run.trace if entry.frame_id.startswith("root:each#0:")
|
||||||
|
]
|
||||||
|
assert [entry.frame_id for entry in batch_entries] == [
|
||||||
|
"root:each#0:0",
|
||||||
|
"root:each#0:1",
|
||||||
|
"root:each#0:2",
|
||||||
|
]
|
||||||
|
assert [entry.step_number for entry in batch_entries] == [2, 3, 4]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_async_batch_reservations_kept_after_failure() -> None:
|
||||||
|
"""Admitted attempts stay consumed even when a handler raises."""
|
||||||
|
workflow = _concurrent_workflow(max_active=3, name="async_reserved_failure")
|
||||||
|
items = ["a", "b", "c"]
|
||||||
|
run, index = _prepare_limited_run(workflow, items, max_steps=20)
|
||||||
|
|
||||||
|
async def _noop(payload: dict[str, Any], _ctx: object) -> dict[str, Any]:
|
||||||
|
return {"outcome": "ok", "output": payload}
|
||||||
|
|
||||||
|
await step_workflow_async(workflow, run, {"record": _noop}, index=index)
|
||||||
|
base_steps = run.steps_executed
|
||||||
|
assert base_steps == 1
|
||||||
|
|
||||||
|
release = asyncio.Event()
|
||||||
|
started: list[str] = []
|
||||||
|
|
||||||
|
async def record(payload: dict[str, Any], _ctx: object) -> dict[str, Any]:
|
||||||
|
started.append(payload["value"])
|
||||||
|
await release.wait()
|
||||||
|
if payload["value"] == "b":
|
||||||
|
raise ValueError("bad item")
|
||||||
|
return {"outcome": "ok", "output": payload}
|
||||||
|
|
||||||
|
batch = asyncio.create_task(
|
||||||
|
step_workflow_async(workflow, run, {"record": record}, index=index)
|
||||||
|
)
|
||||||
|
await _wait_for(lambda: len(started) == 3)
|
||||||
|
release.set()
|
||||||
|
with pytest.raises(ValueError, match="bad item"):
|
||||||
|
await batch
|
||||||
|
|
||||||
|
# All three admissions remain consumed; the failure did not refund them.
|
||||||
|
assert run.steps_executed == base_steps + 3
|
||||||
|
assert remaining_step_attempts(run) == 20 - 4
|
||||||
|
assert run.frames["root:each#0:0"].step_number == 2
|
||||||
|
assert run.frames["root:each#0:1"].step_number == 3
|
||||||
|
assert run.frames["root:each#0:2"].step_number == 4
|
||||||
|
|
||||||
|
|
||||||
|
async def test_async_batch_settles_siblings_and_discards_later_commits() -> None:
|
||||||
|
"""Siblings settle before raising; later commits are discarded in order."""
|
||||||
|
workflow = _concurrent_workflow(max_active=3, name="async_settle_discard")
|
||||||
|
items = ["a", "b", "c"]
|
||||||
|
run, index = _prepare_limited_run(workflow, items, max_steps=20)
|
||||||
|
|
||||||
|
async def _noop(payload: dict[str, Any], _ctx: object) -> dict[str, Any]:
|
||||||
|
return {"outcome": "ok", "output": payload}
|
||||||
|
|
||||||
|
await step_workflow_async(workflow, run, {"record": _noop}, index=index)
|
||||||
|
|
||||||
|
release = asyncio.Event()
|
||||||
|
started: list[str] = []
|
||||||
|
finished: list[str] = []
|
||||||
|
|
||||||
|
async def record(payload: dict[str, Any], _ctx: object) -> dict[str, Any]:
|
||||||
|
value = payload["value"]
|
||||||
|
started.append(value)
|
||||||
|
await release.wait()
|
||||||
|
if value == "b":
|
||||||
|
finished.append(value)
|
||||||
|
raise ValueError("middle fails")
|
||||||
|
if value == "c":
|
||||||
|
# The reservation-later sibling is slow: it must still settle
|
||||||
|
# before the middle failure is raised.
|
||||||
|
await asyncio.sleep(0.05)
|
||||||
|
finished.append(value)
|
||||||
|
return {"outcome": "ok", "output": payload}
|
||||||
|
|
||||||
|
batch = asyncio.create_task(
|
||||||
|
step_workflow_async(workflow, run, {"record": record}, index=index)
|
||||||
|
)
|
||||||
|
await _wait_for(lambda: len(started) == 3)
|
||||||
|
assert started == ["a", "b", "c"]
|
||||||
|
release.set()
|
||||||
|
with pytest.raises(ValueError, match="middle fails"):
|
||||||
|
await batch
|
||||||
|
|
||||||
|
# Every handler settled, including the slow reservation-later sibling.
|
||||||
|
assert sorted(finished) == ["a", "b", "c"]
|
||||||
|
# Reservations are kept even for the discarded sibling.
|
||||||
|
assert run.steps_executed == 4
|
||||||
|
# Preceding success committed in reservation order...
|
||||||
|
committed = [entry.frame_id for entry in run.trace if entry.step_type == "node"]
|
||||||
|
assert committed == ["root:each#0:0"]
|
||||||
|
assert run.frames["root:each#0:0"].status == FrameStatus.COMPLETED
|
||||||
|
# ...while the later sibling result was discarded without state/trace.
|
||||||
|
assert "root:each#0:2" not in committed
|
||||||
|
assert run.frames["root:each#0:2"].status == FrameStatus.RUNNING
|
||||||
|
assert run.frames["root:each#0:2"].step_number == 4
|
||||||
|
|
||||||
|
|
||||||
|
async def test_sync_async_parity_for_serial_execution() -> None:
|
||||||
|
"""Equivalent serial runs count the same steps sync and async."""
|
||||||
|
|
||||||
|
def _serial_workflow(name: str) -> Workflow:
|
||||||
|
return Workflow(
|
||||||
|
name=name,
|
||||||
|
input_schema=SchemaRef(type="object", properties={"items": {"type": "array"}}),
|
||||||
|
state_schema=StateSchema.from_field_map(
|
||||||
|
{
|
||||||
|
"items": StateField(type="array"),
|
||||||
|
"seen": StateField(
|
||||||
|
type="array",
|
||||||
|
reducer=ReducerRef(name="wf.std.append"),
|
||||||
|
),
|
||||||
|
}
|
||||||
|
),
|
||||||
|
output_schema=SchemaRef(type="object", properties={"seen": {"type": "array"}}),
|
||||||
|
node_defs=[
|
||||||
|
NodeDef(
|
||||||
|
name="record",
|
||||||
|
input_schema=SchemaRef(
|
||||||
|
type="object",
|
||||||
|
properties={"value": {}, "seen": {}},
|
||||||
|
required=["value", "seen"],
|
||||||
|
),
|
||||||
|
output_schema=SchemaRef(
|
||||||
|
type="object",
|
||||||
|
properties={"seen": {}},
|
||||||
|
required=["seen"],
|
||||||
|
),
|
||||||
|
outcomes=["ok"],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
start="each",
|
||||||
|
nodes=[
|
||||||
|
ForeachNode.model_validate(
|
||||||
|
{
|
||||||
|
"id": "each",
|
||||||
|
"type": "foreach",
|
||||||
|
"over": "state.items",
|
||||||
|
"as": "item",
|
||||||
|
"mode": "concurrent",
|
||||||
|
"concurrent": {"max_active": 1, "max_outstanding": 1},
|
||||||
|
}
|
||||||
|
),
|
||||||
|
NodeUse.model_validate(
|
||||||
|
{
|
||||||
|
"id": "record",
|
||||||
|
"type": "node",
|
||||||
|
"node": "record",
|
||||||
|
"input": [
|
||||||
|
{"target": "value", "path": "context.item"},
|
||||||
|
{"target": "seen", "path": "context.item"},
|
||||||
|
],
|
||||||
|
"output": [{"source": "seen", "target": "state.seen"}],
|
||||||
|
}
|
||||||
|
),
|
||||||
|
],
|
||||||
|
edges=[
|
||||||
|
Edge.model_validate({"from": "each", "outcome": "loop", "to": "record"}),
|
||||||
|
Edge.model_validate({"from": "record", "outcome": "ok", "to": "each"}),
|
||||||
|
Edge.model_validate({"from": "each", "outcome": "done", "to": END}),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
sync_workflow = _serial_workflow("parity_sync")
|
||||||
|
async_workflow = _serial_workflow("parity_sync")
|
||||||
|
|
||||||
|
sync_run = execute_workflow(
|
||||||
|
sync_workflow,
|
||||||
|
{"items": ["a", "b"]},
|
||||||
|
{"record": lambda payload, _ctx: {"outcome": "ok", "output": payload}},
|
||||||
|
)
|
||||||
|
|
||||||
|
async def record(payload: dict[str, Any], _ctx: object) -> dict[str, Any]:
|
||||||
|
return {"outcome": "ok", "output": payload}
|
||||||
|
|
||||||
|
async_run = await execute_workflow_async(
|
||||||
|
async_workflow, {"items": ["a", "b"]}, {"record": record}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert sync_run.status == RunStatus.COMPLETED
|
||||||
|
assert async_run.status == RunStatus.COMPLETED
|
||||||
|
assert async_run.steps_executed == sync_run.steps_executed
|
||||||
|
assert [entry.step_number for entry in async_run.trace] == [
|
||||||
|
entry.step_number for entry in sync_run.trace
|
||||||
|
]
|
||||||
Reference in New Issue
Block a user