async concurrent foreach

This commit is contained in:
lda
2026-05-22 21:24:21 +07:00 Verified
parent f8426d2c9b
commit 6802593d2a
5 changed files with 399 additions and 13 deletions
+119 -4
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable, Mapping
from dataclasses import dataclass
from typing import Any, cast
from wf_core.conditions import safe_resolve_path
@@ -10,7 +11,12 @@ from wf_core.models.results import NodeResult
from wf_core.models.schemas import NodeDef
from wf_core.models.steps import InputPathBinding, InputValueBinding, NodeUse
from wf_core.models.workflow import Workflow
from wf_core.run_state import RunState, RuntimeContext, StepExecutionResult
from wf_core.run_state import (
ExecutionFrame,
RunState,
RuntimeContext,
StepExecutionResult,
)
from wf_core.runtime.foreach_state import ForeachBarrierState, item_frame_owner
from wf_core.runtime.ops.frames import frame_context_values
from wf_core.runtime.ops.merges import ReducerDefinition
@@ -25,14 +31,26 @@ AsyncNodeHandler = Callable[
]
@dataclass(slots=True)
class PendingAsyncNodeResult:
"""Async handler result captured before sequential state finalization."""
frame: ExecutionFrame
node: NodeUse
node_def: NodeDef
resolved_input: dict[str, Any]
raw_result: NodeResult | dict[str, Any]
state_view: dict[str, Any]
def _resolve_node_execution(
*,
workflow: Workflow,
run: RunState,
frame: ExecutionFrame,
node: NodeUse,
node_def: NodeDef,
) -> tuple[dict[str, Any], RuntimeContext, dict[str, Any]]:
frame = run.current_frame()
context_values = frame_context_values(frame)
state_view = state_view_for_frame(run, frame)
resolved_input: dict[str, Any] = {}
@@ -72,6 +90,7 @@ def _finalize_node_execution(
*,
workflow: Workflow,
run: RunState,
frame: ExecutionFrame,
node: NodeUse,
node_def: NodeDef,
resolved_input: dict[str, Any],
@@ -96,7 +115,7 @@ def _finalize_node_execution(
state_view,
reducers=reducers,
)
owner = item_frame_owner(run.current_frame())
owner = item_frame_owner(frame)
if owner is None:
state_changes = commit_state_patch(run.state, patch)
else:
@@ -106,7 +125,7 @@ def _finalize_node_execution(
if barrier is not None and barrier.mode == "concurrent":
barrier.add_success_patch(
index=item_index,
frame_id=run.current_frame().id,
frame_id=frame.id,
patch=patch,
)
barrier.save_to_frame(parent_frame, foreach_node_id)
@@ -138,6 +157,7 @@ def execute_node_use(
resolved_input, context, state_view = _resolve_node_execution(
workflow=workflow,
run=run,
frame=run.current_frame(),
node=node,
node_def=node_def,
)
@@ -145,6 +165,7 @@ def execute_node_use(
return _finalize_node_execution(
workflow=workflow,
run=run,
frame=run.current_frame(),
node=node,
node_def=node_def,
resolved_input=resolved_input,
@@ -171,6 +192,7 @@ async def execute_node_use_async(
resolved_input, context, state_view = _resolve_node_execution(
workflow=workflow,
run=run,
frame=run.current_frame(),
node=node,
node_def=node_def,
)
@@ -182,6 +204,7 @@ async def execute_node_use_async(
return _finalize_node_execution(
workflow=workflow,
run=run,
frame=run.current_frame(),
node=node,
node_def=node_def,
resolved_input=resolved_input,
@@ -191,6 +214,98 @@ async def execute_node_use_async(
)
async def invoke_node_use_async_for_frame(
workflow: Workflow,
run: RunState,
frame: ExecutionFrame,
node: NodeUse,
node_def: NodeDef,
registry: Mapping[str, AsyncNodeHandler],
) -> PendingAsyncNodeResult:
"""Resolve input, await the async handler, and defer state finalization.
Async concurrent foreach can run handler awaits concurrently, but state
patches and traces must still be finalized sequentially against `RunState`.
"""
handler = registry.get(node.node)
if handler is None:
raise WorkflowExecutionError(
f"no handler registered for node def {node.node!r}"
)
resolved_input, context, state_view = _resolve_node_execution(
workflow=workflow,
run=run,
frame=frame,
node=node,
node_def=node_def,
)
raw_or_awaitable = handler(resolved_input, context)
if isinstance(raw_or_awaitable, Awaitable):
raw_result = await raw_or_awaitable
else:
raw_result = raw_or_awaitable
return PendingAsyncNodeResult(
frame=frame,
node=node,
node_def=node_def,
resolved_input=resolved_input,
raw_result=cast(NodeResult | dict[str, Any], raw_result),
state_view=state_view,
)
def finalize_pending_async_node_result(
workflow: Workflow,
run: RunState,
pending: PendingAsyncNodeResult,
reducers: Mapping[str, ReducerDefinition] | None = None,
) -> StepExecutionResult:
"""Finalize a previously awaited async node result sequentially."""
return _finalize_node_execution(
workflow=workflow,
run=run,
frame=pending.frame,
node=pending.node,
node_def=pending.node_def,
resolved_input=pending.resolved_input,
raw_result=pending.raw_result,
state_view=pending.state_view,
reducers=reducers,
)
async def execute_node_use_async_for_frame(
workflow: Workflow,
run: RunState,
frame: ExecutionFrame,
node: NodeUse,
node_def: NodeDef,
registry: Mapping[str, AsyncNodeHandler],
reducers: Mapping[str, ReducerDefinition] | None = None,
) -> StepExecutionResult:
"""Execute one async node against an explicit frame.
Async concurrent foreach cannot rely on `run.current_frame()` while several
handlers are in flight. This explicit-frame helper keeps input resolution
and item-local overlay lookup tied to the frame that launched the handler.
"""
pending = await invoke_node_use_async_for_frame(
workflow,
run,
frame=frame,
node=node,
node_def=node_def,
registry=registry,
)
return finalize_pending_async_node_result(
workflow=workflow,
run=run,
pending=pending,
reducers=reducers,
)
def coerce_node_result(raw_result: NodeResult | dict[str, Any]) -> NodeResult:
if isinstance(raw_result, NodeResult):
return raw_result
+129 -3
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
import asyncio
from collections.abc import Mapping
from typing import Any
@@ -19,14 +20,17 @@ from wf_core.runtime.ops.handlers import (
handle_interrupt_step,
handle_join_step,
)
from wf_core.runtime.ops.index import WorkflowIndex
from wf_core.runtime.ops.index import WorkflowIndex, build_workflow_index
from wf_core.runtime.ops.merges import ReducerDefinition
from wf_core.runtime.ops.nodes import (
AsyncNodeHandler,
NodeHandler,
execute_node_use,
execute_node_use_async,
finalize_pending_async_node_result,
invoke_node_use_async_for_frame,
)
from wf_core.runtime.foreach_state import ForeachBarrierState, item_frame_owner
from wf_core.runtime.scheduler import (
ForeachIterationMetadata,
select_next_frame,
@@ -129,7 +133,7 @@ def _mark_handled_item_failure(
run: RunState,
index: WorkflowIndex,
frame: ExecutionFrame,
exc: Exception,
exc: BaseException,
) -> bool:
"""Record skip/collect item failures without failing the whole run here."""
metadata = ForeachIterationMetadata.from_frame(frame)
@@ -163,7 +167,18 @@ async def step_workflow_async(
if frame is None or frame.status != FrameStatus.RUNNING:
if select_next_frame(run) is None:
return run
prepared = prepare_step(workflow, run, index)
frame = run.current_frame()
resolved_index = index or build_workflow_index(workflow)
if _can_batch_async_foreach_item(run, resolved_index, frame):
return await _step_async_foreach_item_batch(
workflow,
run,
registry,
index=resolved_index,
reducers=reducers,
first_frame=frame,
)
prepared = prepare_step(workflow, run, resolved_index)
if prepared is None:
return run
index, step = prepared
@@ -206,3 +221,114 @@ async def step_workflow_async(
step_type=step.type,
step_result=step_result,
)
async def _step_async_foreach_item_batch(
workflow: Workflow,
run: RunState,
registry: Mapping[str, AsyncNodeHandler],
*,
index: WorkflowIndex,
first_frame: ExecutionFrame,
reducers: Mapping[str, ReducerDefinition] | None,
) -> RunState:
"""Run one batch of ready concurrent-foreach item node handlers.
Only handler awaits run concurrently. Finalization, tracing, and frame
advancement happen afterward in frame-id order so `RunState` is mutated
deterministically.
"""
frames = [first_frame, *_claim_matching_async_item_frames(run, index, first_frame)]
tasks = [
invoke_node_use_async_for_frame(
workflow,
run,
frame,
_node_use_for_frame(index, frame),
index.node_defs[_node_use_for_frame(index, frame).node],
registry,
)
for frame in frames
]
results = await asyncio.gather(*tasks, return_exceptions=True)
for frame, result in zip(frames, results, strict=True):
run.current_frame_id = frame.id
run.sync_from_current_frame()
if isinstance(result, BaseException):
if _mark_handled_item_failure(run, index, frame, result):
continue
raise result
step_result = finalize_pending_async_node_result(
workflow,
run,
result,
reducers=reducers,
)
node = _node_use_for_frame(index, frame)
complete_step(
run=run,
index=index,
outcome=step_result.outcome,
frame_id=frame.id,
node_id=frame.node_id,
step_type=node.type,
step_result=step_result,
)
return run
def _claim_matching_async_item_frames(
run: RunState,
index: WorkflowIndex,
first_frame: ExecutionFrame,
) -> list[ExecutionFrame]:
owner = item_frame_owner(first_frame)
if owner is None:
return []
parent_frame_id, foreach_node_id, _item_index = owner
claimed: list[ExecutionFrame] = []
remaining_ready: list[str] = []
for frame_id in run.ready_frame_ids:
frame = run.frames[frame_id]
frame_owner = item_frame_owner(frame)
if (
frame.status == FrameStatus.PENDING
and frame_owner is not None
and frame_owner[:2] == (parent_frame_id, foreach_node_id)
and isinstance(index.nodes_by_id.get(frame.node_id), NodeUse)
):
frame.status = FrameStatus.RUNNING
claimed.append(frame)
else:
remaining_ready.append(frame_id)
run.ready_frame_ids = remaining_ready
return claimed
def _can_batch_async_foreach_item(
run: RunState,
index: WorkflowIndex,
frame: ExecutionFrame,
) -> bool:
owner = item_frame_owner(frame)
if owner is None:
return False
parent_frame_id, foreach_node_id, _item_index = owner
parent_frame = run.frames.get(parent_frame_id)
if parent_frame is None:
return False
barrier = ForeachBarrierState.from_frame(parent_frame, foreach_node_id)
return (
barrier is not None
and barrier.mode == "concurrent"
and isinstance(index.nodes_by_id.get(frame.node_id), NodeUse)
)
def _node_use_for_frame(index: WorkflowIndex, frame: ExecutionFrame) -> NodeUse:
step = index.nodes_by_id[frame.node_id]
if not isinstance(step, NodeUse):
raise WorkflowExecutionError(
f"async foreach batch requires node frame, got {frame.node_id!r}"
)
return step