feat: reserve async workflow step attempts

This commit is contained in:
lda
2026-09-05 19:07:08 +07:00 Verified
parent ae27006abb
commit 57b53eb3f6
2 changed files with 482 additions and 6 deletions
+26 -6
View File
@@ -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)
+456
View File
@@ -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
]