This commit is contained in:
Quanzheng Long
2026-03-18 11:41:05 -07:00
parent 12ea4c24d5
commit c0bd2dbc07
10 changed files with 505 additions and 61 deletions
@@ -2,10 +2,12 @@ from .state import (
AdvancedStateGraph,
AnyOfCondition,
ChannelCondition,
ConditionResult,
CompiledGraphEngine,
Context,
GraphRunHandler,
TimerCondition,
WaitForResult,
any_of,
channel_condition,
timer_condition,
@@ -19,6 +21,8 @@ __all__ = [
"ChannelCondition",
"TimerCondition",
"AnyOfCondition",
"ConditionResult",
"WaitForResult",
"channel_condition",
"timer_condition",
"any_of",
@@ -40,6 +40,18 @@ class AnyOfCondition:
conditions: tuple[WaitCondition, ...]
@dataclass(frozen=True)
class ConditionResult:
met: bool
channel_name: str | None = None
values: list[Any] | None = None
@dataclass(frozen=True)
class WaitForResult:
conditions: list[ConditionResult]
WaitCondition = ChannelCondition | TimerCondition
_EXECUTOR_LOCK = threading.Lock()
@@ -183,7 +195,7 @@ class Context:
def __init__(self, run: _GraphEngineRun) -> None:
self._run = run
async def wait_for(self, target: WaitCondition | AnyOfCondition) -> Any:
async def wait_for(self, target: WaitCondition | AnyOfCondition) -> WaitForResult:
resumed = self._run._consume_resume_event(target)
if resumed is not None:
return resumed
@@ -287,25 +299,31 @@ class _GraphEngineRun:
def publish_nowait(self, channel: str, value: Any) -> None:
self._publish_sync(channel, value)
async def wait_for(self, target: WaitCondition | AnyOfCondition) -> Any:
async def wait_for(self, target: WaitCondition | AnyOfCondition) -> WaitForResult:
if isinstance(target, ChannelCondition):
value = await self._wait_for_channel_values(
target.channel, min=target.min, max=target.max
)
return {
"condition": "channel",
"channel": target.channel,
"value": value,
}
return WaitForResult(
conditions=[
ConditionResult(
met=True,
channel_name=target.channel,
values=_normalize_channel_values(value),
)
]
)
if isinstance(target, TimerCondition):
loop = asyncio.get_running_loop()
return await loop.run_in_executor(
await loop.run_in_executor(
_advanced_graph_executor(),
self._rust_engine.wait_timer,
target.seconds,
)
return WaitForResult(conditions=[ConditionResult(met=True)])
if isinstance(target, AnyOfCondition):
return await self._wait_for_any_of(target)
raw_event = await self._wait_for_any_of(target)
return _wait_for_result_from_any_of_event(target, raw_event)
raise ValueError(f"Unsupported wait condition type: {type(target)!r}")
async def _wait_for_channel_values(
@@ -327,7 +345,7 @@ class _GraphEngineRun:
)
return event["value"]
async def _wait_for_any_of(self, condition: AnyOfCondition) -> Any:
async def _wait_for_any_of(self, condition: AnyOfCondition) -> dict[str, Any]:
if not condition.conditions:
raise ValueError("any_of() requires at least one condition")
payload = {
@@ -390,12 +408,14 @@ class _GraphEngineRun:
def _set_resume_event(self, event: dict[str, Any] | None) -> None:
self._local.resume_event = event
def _consume_resume_event(self, target: WaitCondition | AnyOfCondition) -> Any | None:
def _consume_resume_event(
self, target: WaitCondition | AnyOfCondition
) -> WaitForResult | None:
event = cast(dict[str, Any] | None, getattr(self._local, "resume_event", None))
if event is None:
return None
self._local.resume_event = None
return event
return _wait_for_result_from_resume_event(target, event)
def _run_awaitable_in_worker(self, awaitable: Coroutine[Any, Any, Any]) -> Any:
# Create and close a dedicated loop per execution to avoid
@@ -490,6 +510,82 @@ def any_of(*conditions: WaitCondition) -> AnyOfCondition:
return AnyOfCondition(conditions=tuple(conditions))
def _normalize_channel_values(value: Any) -> list[Any]:
if isinstance(value, list):
return value
return [value]
def _wait_for_result_from_resume_event(
target: WaitCondition | AnyOfCondition, event: dict[str, Any]
) -> WaitForResult:
if isinstance(target, ChannelCondition):
return WaitForResult(
conditions=[
ConditionResult(
met=True,
channel_name=target.channel,
values=_normalize_channel_values(event.get("value")),
)
]
)
if isinstance(target, TimerCondition):
return WaitForResult(conditions=[ConditionResult(met=True)])
return _wait_for_result_from_any_of_event(target, event)
def _wait_for_result_from_any_of_event(
target: AnyOfCondition, event: dict[str, Any]
) -> WaitForResult:
results = [ConditionResult(met=False) for _ in target.conditions]
condition = event.get("condition")
if condition == "timer":
for idx, cond in enumerate(target.conditions):
if isinstance(cond, TimerCondition):
results[idx] = ConditionResult(met=True)
break
return WaitForResult(conditions=results)
if condition != "channel":
return WaitForResult(conditions=results)
channel = cast(str | None, event.get("channel"))
value = event.get("value")
if channel == "__any_of__" and isinstance(value, list):
matched = list(value)
cursor = 0
for idx, cond in enumerate(target.conditions):
if not isinstance(cond, ChannelCondition):
continue
if cursor >= len(matched):
continue
item = matched[cursor]
if (
isinstance(item, dict)
and item.get("channel") == cond.channel
and "value" in item
):
results[idx] = ConditionResult(
met=True,
channel_name=cond.channel,
values=_normalize_channel_values(item.get("value")),
)
cursor += 1
return WaitForResult(conditions=results)
for idx, cond in enumerate(target.conditions):
if isinstance(cond, ChannelCondition) and cond.channel == channel:
results[idx] = ConditionResult(
met=True,
channel_name=cond.channel,
values=_normalize_channel_values(value),
)
break
return WaitForResult(conditions=results)
def _condition_to_rust(condition: WaitCondition) -> dict[str, Any]:
if isinstance(condition, ChannelCondition):
return {
@@ -1,7 +1,12 @@
import pytest
from typing_extensions import TypedDict
from saf_python_sdk.advanced_graph import AdvancedStateGraph, Context, channel_condition
from saf_python_sdk.advanced_graph import (
AdvancedStateGraph,
Context,
any_of,
channel_condition,
)
from saf_python_sdk.types import Command, Send
pytestmark = pytest.mark.anyio
@@ -77,8 +82,8 @@ async def test_channel_wait_respects_max_m() -> None:
return Command(update=state, goto=Send("wait_node", None))
async def wait_node(ctx: Context, _input: None, state: PrimitiveState) -> dict[str, object]:
event = await ctx.wait_for(channel_condition("events", min=2, max=4))
values = event["value"]
result = await ctx.wait_for(channel_condition("events", min=2, max=4))
values = result.conditions[0].values or []
assert isinstance(values, list)
return {"counter": len(values), "logs": values, "done": "ok"}
@@ -90,3 +95,54 @@ async def test_channel_wait_respects_max_m() -> None:
assert result["logs"] == ["a", "b", "c"]
assert result["done"] == "ok"
async def test_any_of_consumes_all_ready_channels() -> None:
graph = AdvancedStateGraph(PrimitiveState)
graph.add_async_channel("alpha", str)
graph.add_async_channel("beta", str)
async def start_node(ctx: Context, state: PrimitiveState) -> Command:
ctx.publish_to_channel("alpha", "a1")
ctx.publish_to_channel("beta", "b1")
return Command(update=state, goto=Send("wait_node", None))
async def wait_node(ctx: Context, _input: None, state: PrimitiveState) -> Command:
first = await ctx.wait_for(
any_of(channel_condition("alpha"), channel_condition("beta"))
)
assert len(first.conditions) == 2
assert first.conditions[0].met is True
assert first.conditions[0].channel_name == "alpha"
assert first.conditions[0].values == ["a1"]
assert first.conditions[1].met is True
assert first.conditions[1].channel_name == "beta"
assert first.conditions[1].values == ["b1"]
ctx.publish_to_channel("beta", "b2")
return Command(
update={"counter": 1, "logs": ["matched=2"], "done": None},
goto=Send("verify_node", None),
)
async def verify_node(
ctx: Context, _input: None, state: PrimitiveState
) -> dict[str, object]:
second = await ctx.wait_for(channel_condition("beta"))
values = second.conditions[0].values or []
return {
"counter": 2,
"logs": [*state["logs"], f"beta={values[0]}"],
"done": "ok",
}
graph.add_entry_node(start_node)
graph.add_node(wait_node)
graph.add_finish_node(verify_node)
result = await graph.compile().ainvoke({"counter": 0, "logs": [], "done": None})
assert result["counter"] == 2
assert result["logs"] == [
"matched=2",
"beta=b2",
]
assert result["done"] == "ok"
@@ -92,7 +92,7 @@ def build_main_agent(planner: MockLLM, sub_agent: Any) -> Any:
async def wait_node(ctx: Context, state: MainAgentState) -> Command:
# Lightweight interrupt: only this node blocks for the next relevant signal.
event = await ctx.wait_for(
result = await ctx.wait_for(
any_of(
channel_condition("tool_completion_channel"),
channel_condition("subagent_completion_channel"),
@@ -100,21 +100,33 @@ def build_main_agent(planner: MockLLM, sub_agent: Any) -> Any:
timer_condition(seconds=1),
)
)
if event["condition"] == "channel":
channel = event["channel"]
payload = event["value"]
if channel == "tool_completion_channel":
state["output"].append(f"tool: {payload}")
elif channel == "subagent_completion_channel":
state["output"].append(f"sub_agent: {payload}")
elif channel == "user_input_channel":
state["output"].append(f"user_input: {payload}")
had_channel_update = False
for item in result.conditions:
if not item.met:
continue
if item.channel_name == "tool_completion_channel":
payloads = item.values or []
for payload in payloads:
state["output"].append(f"tool: {payload}")
had_channel_update = True
elif item.channel_name == "subagent_completion_channel":
payloads = item.values or []
for payload in payloads:
state["output"].append(f"sub_agent: {payload}")
had_channel_update = True
elif item.channel_name == "user_input_channel":
payloads = item.values or []
for payload in payloads:
state["output"].append(f"user_input: {payload}")
had_channel_update = True
if had_channel_update:
# State changed -> ask planner what to do next.
return Command(update=state, goto=Send("llm_node", None))
else:
state["output"].append("timer: no updates yet")
# No meaningful state change -> keep waiting without calling planner.
return Command(update=state, goto=Send("wait_node", None))
state["output"].append("timer: no updates yet")
# No meaningful state change -> keep waiting without calling planner.
return Command(update=state, goto=Send("wait_node", None))
async def tool_node(ctx: Context, tool_input: str) -> None:
await asyncio.sleep(0.1)