This commit is contained in:
Quanzheng Long
2026-03-12 13:13:08 -07:00
parent be46e91180
commit 4b2167dd25
3 changed files with 518 additions and 103 deletions
@@ -0,0 +1,13 @@
from langgraph.graph_engine.state import (
AdvancedStateGraph,
CompiledGraphEngine,
publish_to_channel,
wait_for,
)
__all__ = (
"AdvancedStateGraph",
"CompiledGraphEngine",
"publish_to_channel",
"wait_for",
)
@@ -0,0 +1,352 @@
from __future__ import annotations
import asyncio
import contextvars
import inspect
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import Any, Generic, TypeVar, cast
from langgraph.types import Command, Send
StateT = TypeVar("StateT")
_CURRENT_RUN: contextvars.ContextVar[_GraphEngineRun | None] = contextvars.ContextVar(
"langgraph_graph_engine_run", default=None
)
@dataclass(frozen=True)
class _ChannelSpec:
typ: Any
maxsize: int
class AdvancedStateGraph(Generic[StateT]):
"""Experimental in-memory graph engine with async channels."""
def __init__(self, state_schema: type[StateT]) -> None:
self.state_schema = state_schema
self._nodes: dict[str, Callable[..., Any]] = {}
self._channels: dict[str, _ChannelSpec] = {}
self._entry_point: str | None = None
self._finish_point: str | None = None
def add_node(
self,
name_or_node: str | Callable[..., Any],
node: Callable[..., Any] | None = None,
) -> str:
if node is None:
if not callable(name_or_node):
raise TypeError("add_node() expects a callable when name is omitted")
node_name = _infer_node_name(name_or_node)
node_fn = name_or_node
else:
if not isinstance(name_or_node, str):
raise TypeError("add_node() expects a string node name")
node_name = name_or_node
node_fn = node
if node_name in self._nodes:
raise ValueError(f"Node `{node_name}` already exists")
self._nodes[node_name] = node_fn
return node_name
def node(self, node: Callable[..., Any]) -> str:
"""Register a node using the function name as node id."""
return self.add_node(node)
def add_async_channel(
self, name: str, typ: Any, maxsize: int | None = None
) -> None:
if name in self._channels:
raise ValueError(f"Channel `{name}` already exists")
self._channels[name] = _ChannelSpec(typ=typ, maxsize=maxsize or 0)
def set_entry_point(self, name_or_node: str | Callable[..., Any]) -> None:
self._entry_point = self._resolve_node_name(name_or_node)
def set_finish_point(self, name_or_node: str | Callable[..., Any]) -> None:
self._finish_point = self._resolve_node_name(name_or_node)
def add_entry_node(self, node: Callable[..., Any]) -> str:
node_name = self.add_node(node)
self.set_entry_point(node_name)
return node_name
def add_finish_node(self, node: Callable[..., Any]) -> str:
node_name = self.add_node(node)
self.set_finish_point(node_name)
return node_name
def _resolve_node_name(self, name_or_node: str | Callable[..., Any]) -> str:
if isinstance(name_or_node, str):
return name_or_node
node_name = _infer_node_name(name_or_node)
if node_name not in self._nodes:
self._nodes[node_name] = name_or_node
return node_name
def compile(self) -> CompiledGraphEngine[StateT]:
if self._entry_point is None:
raise ValueError("Entry point is not set")
if self._finish_point is None:
raise ValueError("Finish point is not set")
if self._entry_point not in self._nodes:
raise ValueError(f"Entry point node `{self._entry_point}` does not exist")
if self._finish_point not in self._nodes:
raise ValueError(f"Finish point node `{self._finish_point}` does not exist")
return CompiledGraphEngine(
nodes=dict(self._nodes),
channels=dict(self._channels),
entry_point=self._entry_point,
finish_point=self._finish_point,
)
class CompiledGraphEngine(Generic[StateT]):
"""Executable runtime for `AdvancedStateGraph`."""
def __init__(
self,
*,
nodes: dict[str, Callable[..., Any]],
channels: dict[str, _ChannelSpec],
entry_point: str,
finish_point: str,
) -> None:
self._nodes = nodes
self._channels = channels
self._entry_point = entry_point
self._finish_point = finish_point
self._active_run: _GraphEngineRun | None = None
self._run_lock = asyncio.Lock()
async def ainvoke(self, initial_state: StateT) -> StateT:
async with self._run_lock:
if self._active_run is not None:
raise RuntimeError("Graph engine already has an active run")
run = _GraphEngineRun(
nodes=self._nodes,
channel_specs=self._channels,
entry_point=self._entry_point,
finish_point=self._finish_point,
)
self._active_run = run
try:
return await run.run(initial_state)
finally:
async with self._run_lock:
if self._active_run is run:
self._active_run = None
async def apublish_to_channel(self, channel: str, value: Any) -> None:
run = self._active_run
if run is None:
raise RuntimeError("No active graph run to publish to")
await run.publish(channel, value)
class _GraphEngineRun:
def __init__(
self,
*,
nodes: dict[str, Callable[..., Any]],
channel_specs: dict[str, _ChannelSpec],
entry_point: str,
finish_point: str,
) -> None:
self._nodes = nodes
self._entry_point = entry_point
self._finish_point = finish_point
self._channels: dict[str, asyncio.Queue[Any]] = {
name: asyncio.Queue(maxsize=spec.maxsize)
for name, spec in channel_specs.items()
}
self._tasks: set[asyncio.Task[list[Send]]] = set()
self._finished = False
self._state: Any = None
async def run(self, initial_state: StateT) -> StateT:
self._state = initial_state
self._schedule(Send(self._entry_point, initial_state))
try:
while self._tasks and not self._finished:
done, _ = await asyncio.wait(
self._tasks, return_when=asyncio.FIRST_COMPLETED
)
for task in done:
self._tasks.remove(task)
exc = task.exception()
if exc is not None:
await self._cancel_all_tasks()
raise exc
sends = task.result()
for send in sends:
self._schedule(send)
if self._finished:
await self._cancel_all_tasks()
return cast(StateT, self._state)
finally:
await self._cancel_all_tasks()
async def publish(self, channel: str, value: Any) -> None:
queue = self._get_channel(channel)
await queue.put(value)
def publish_nowait(self, channel: str, value: Any) -> None:
queue = self._get_channel(channel)
queue.put_nowait(value)
async def wait_for(self, channel: str, n: int = 1) -> Any:
if n < 1:
raise ValueError("wait_for count `n` must be >= 1")
queue = self._get_channel(channel)
if n == 1:
return await queue.get()
values: list[Any] = []
for _ in range(n):
values.append(await queue.get())
return values
def _get_channel(self, channel: str) -> asyncio.Queue[Any]:
if channel not in self._channels:
raise ValueError(f"Unknown channel `{channel}`")
return self._channels[channel]
def _schedule(self, send: Send) -> None:
if self._finished:
return
task: asyncio.Task[list[Send]] = asyncio.create_task(self._execute_send(send))
self._tasks.add(task)
async def _cancel_all_tasks(self) -> None:
if not self._tasks:
return
to_cancel = list(self._tasks)
for task in to_cancel:
task.cancel()
await asyncio.gather(*to_cancel, return_exceptions=True)
self._tasks.clear()
async def _execute_send(self, send: Send) -> list[Send]:
node_name = _resolve_target_name(send.node)
if node_name not in self._nodes:
raise ValueError(f"Unknown node `{node_name}`")
node = self._nodes[node_name]
token = _CURRENT_RUN.set(self)
try:
result = node(send.arg)
if inspect.isawaitable(result):
result = await result
finally:
_CURRENT_RUN.reset(token)
if isinstance(result, Command):
self._apply_update(result.update)
next_sends = _normalize_goto(result.goto, default_arg=self._state)
else:
self._apply_update(result)
next_sends = _normalize_result_to_sends(result, default_arg=self._state)
if node_name == self._finish_point:
self._finished = True
return []
return next_sends
def _apply_update(self, update: Any) -> None:
if update is None:
return
if isinstance(update, Mapping):
if isinstance(self._state, Mapping):
# Keep semantics simple: in-place update for mapping-like state.
cast(dict[str, Any], self._state).update(update)
return
if isinstance(update, Sequence) and not isinstance(update, (str, bytes)):
pairs = list(update)
if all(
isinstance(item, tuple) and len(item) == 2 and isinstance(item[0], str)
for item in pairs
):
if isinstance(self._state, Mapping):
cast(dict[str, Any], self._state).update(
cast(dict[str, Any], pairs)
)
return
def _normalize_result_to_sends(result: Any, *, default_arg: Any) -> list[Send]:
if result is None:
return []
if isinstance(result, Send):
return [result]
if callable(result):
return [Send(_infer_node_name(result), default_arg)]
if isinstance(result, str):
return [Send(result, default_arg)]
if isinstance(result, Sequence) and not isinstance(result, (str, bytes)):
sends: list[Send] = []
for item in result:
if isinstance(item, Send):
sends.append(item)
elif callable(item):
sends.append(Send(_infer_node_name(item), default_arg))
elif isinstance(item, str):
sends.append(Send(item, default_arg))
return sends
return []
def _normalize_goto(goto: Any, *, default_arg: Any) -> list[Send]:
if not goto:
return []
if isinstance(goto, Send):
return [goto]
if callable(goto):
return [Send(_infer_node_name(goto), default_arg)]
if isinstance(goto, str):
return [Send(goto, default_arg)]
if isinstance(goto, Sequence):
sends: list[Send] = []
for item in goto:
if isinstance(item, Send):
sends.append(item)
elif callable(item):
sends.append(Send(_infer_node_name(item), default_arg))
elif isinstance(item, str):
sends.append(Send(item, default_arg))
return sends
return []
async def wait_for(channel: str, n: int = 1) -> Any:
run = _CURRENT_RUN.get()
if run is None:
raise RuntimeError("wait_for() can only be used inside graph_engine nodes")
return await run.wait_for(channel, n=n)
def publish_to_channel(channel: str, value: Any) -> None:
run = _CURRENT_RUN.get()
if run is None:
raise RuntimeError(
"publish_to_channel() can only be used inside graph_engine nodes"
)
run.publish_nowait(channel, value)
def _infer_node_name(node: Callable[..., Any]) -> str:
node_name = getattr(node, "__name__", "")
if not node_name or node_name == "<lambda>":
raise ValueError("Cannot infer node name from anonymous callable")
return node_name
def _resolve_target_name(target: Any) -> str:
if isinstance(target, str):
return target
if callable(target):
return _infer_node_name(target)
raise ValueError(f"Unsupported node target type: {type(target)!r}")
@@ -1,124 +1,174 @@
from typing import Annotated
import asyncio
from dataclasses import dataclass
from typing import Any, Literal
import pytest
from typing_extensions import TypedDict
from langgraph.constants import END, START
from langgraph.graph import StateGraph
from langgraph.graph_engine import AdvancedStateGraph, publish_to_channel, wait_for
from langgraph.types import Command, Send
pytestmark = pytest.mark.anyio
class MainAgentState(TypedDict):
input: str
output: list[str]
advanced_flow = AdvancedStateGraph(MainAgentState)
advanced_flow.add_async_channel("inbox", str) # default to infinite buffer like Rust channel
prompt = "based on current state, decide to use tool, or kick off subagent, or complete"
main_llm = new_mock_llm(prompt, [email_tool, slack_tool])
async def llm_node(state: MainAgentState):
decisions = await main_llm.ainvoke(state)
sends = []
for decision in decisions:
if decision.type == "end":
return Command(goto=Send("end", decision.complete)) # NOTE: we can simplify to just Complete(decision.complete)
elif decision.type == "sub_agent":
sends.append(Send("sub_agent", decision.sub_agent)) # NOTE: we can simplify to just sends.apppend(subagent, decision.sub_agent)
elif decision.type == "tool":
sends.append(Send("tool", decision.tool))
sends.append(Send("wait_node"))
return Command(goto=sends)
async def wait_node(state: MainAgentState):
# wait for at least one message on the inbox channel
# this is a "lightweight" interrupt, that does not block the entire graph
msgs = wait_for("inbox")
output = state.output
if msg.type == "tool":
output.append("tool: " + msg.payload)
elif msg.type == "sub_agent":
output.append("sub_agent: " + msg.payload)
elif msg.type == "user_input":
output.append("user_input: " + msg.payload)
# loop back to llm node with new output
return Command(goto=Send("llm_node", output))
done: str | None
async def tool_node(tool_input: str):
await asyncio.sleep(5)
output = "tool completed for: " + tool_input
publish_to_channel("inbox", {"type": "tool", "payload": output})
# just complete without going to next node
class SubAgentState(TypedDict):
input: str
output: str
async def order_food_node(state: str):
return "order_food_node completed for: " + state
# sub agent uses regular/simple state graph
sub_agent = StateGraph(str)
async def research_node (state: str):
await asyncio.sleep(10)
return "research sub agent completed for: " + state
sub_agent.add_node("research_node", research_node)
sub_agent.add_edge(START, "research_node")
sub_agent.add_edge("research_node", END)
@dataclass(frozen=True)
class Decision:
type: Literal["end", "sub_agent", "tool"]
sub_agent: str | None = None
tool: str | None = None
complete: str | None = None
async def sub_agent_node(sub_agent_input: str):
sub_agent_output = sub_agent.invoke({"input": sub_agent_input})
publish_to_channel("inbox", {"type": "sub_agent", "payload": sub_agent_output})
# just complete without going to next node
advanced_flow.add_node("llm_node", llm_node)
advanced_flow.add_node("wait_node", wait_node)
advanced_flow.add_node("tool_node", tool_node)
advanced_flow.add_node("sub_agent_node", sub_agent_node)
advanced_flow.add_node("order_food_node", order_food_node)
advanced_flow.set_entry_point("llm_node")
advanced_flow.set_finish_point("order_food_node")
class MockPlanner:
def __init__(self) -> None:
self.responses: list[list[Decision]] = []
self._idx = 0
## NOTE: above can be simplified to:
# advanced_flow.add_entry_node(llm_node)
# advanced_flow.node(wait_node)
# advanced_flow.node(tool_node)
# advanced_flow.node(sub_agent_node)
# advanced_flow.add_finish_node(order_food_node)
async def ainvoke(self, _: MainAgentState) -> list[Decision]:
if self._idx >= len(self.responses):
return []
response = self.responses[self._idx]
self._idx += 1
return response
main_agent = advanced_flow.compile()
async def test_async_sub_graph():
main_llm.mock_response = [
# first llm invoke
def build_sub_agent() -> Any:
# Sub-agent uses the regular/simple StateGraph API.
sub_agent = StateGraph(SubAgentState)
async def research_node(state: SubAgentState) -> dict[str, str]:
# Make timing deterministic for the prototype flow assertions.
if state["input"] == "research lunch options":
await asyncio.sleep(0.05)
else:
await asyncio.sleep(0.09)
return {"output": f"research sub agent completed for: {state['input']}"}
sub_agent.add_node("research_node", research_node)
sub_agent.add_edge(START, "research_node")
sub_agent.add_edge("research_node", END)
return sub_agent.compile()
async def test_async_sub_graph() -> None:
planner = MockPlanner()
sub_agent = build_sub_agent()
advanced_flow = AdvancedStateGraph(MainAgentState)
# Default behavior is an unbounded async channel (maxsize=None).
advanced_flow.add_async_channel("inbox", dict)
async def llm_node(state: MainAgentState) -> Command:
# Planner decides whether to call a tool, spawn a sub-agent, or finish.
decisions = await planner.ainvoke(state)
sends: list[Send] = []
for decision in decisions:
if decision.type == "end":
# NOTE: this can be simplified further in the future with a dedicated
# complete primitive, instead of routing to a finish node manually.
return Command(
goto=Send(
order_food_node,
decision.complete or "order flow completed",
)
)
if decision.type == "sub_agent" and decision.sub_agent:
sends.append(Send("sub_agent_node", decision.sub_agent))
if decision.type == "tool" and decision.tool:
sends.append(Send("tool_node", decision.tool))
# Keep the main loop responsive: wait for one inbound message and continue.
sends.append(Send("wait_node", state))
return Command(goto=sends)
async def wait_node(state: MainAgentState) -> Command:
# Lightweight interrupt: only this node blocks on inbox.
msg = await wait_for("inbox")
state["output"].append(f"{msg['type']}: {msg['payload']}")
# Loop back to planner with updated output.
return Command(goto=Send("llm_node", state))
async def tool_node(tool_input: str) -> None:
await asyncio.sleep(0.03)
# Fire-and-forget style completion: publish result to inbox and exit.
# (i.e., just complete without explicitly going to a next node)
publish_to_channel(
"inbox",
{"type": "tool", "payload": f"tool completed for: {tool_input}"},
)
async def sub_agent_node(sub_agent_input: str) -> None:
# Sub-agent remains a regular StateGraph, compiled independently.
sub_agent_output = await sub_agent.ainvoke(
{"input": sub_agent_input, "output": ""}
)
# Same pattern as tool node: publish result and complete current node.
publish_to_channel(
"inbox",
{"type": "sub_agent", "payload": sub_agent_output["output"]},
)
async def order_food_node(complete_message: str) -> dict[str, str]:
return {"done": complete_message}
# NOTE: This is the simplified API shape:
# - add_entry_node(fn)
# - node(fn)
# - add_finish_node(fn)
# where node names are inferred from function names via reflection.
advanced_flow.add_entry_node(llm_node)
advanced_flow.node(wait_node)
advanced_flow.node(tool_node)
advanced_flow.node(sub_agent_node)
advanced_flow.add_finish_node(order_food_node)
main_agent = advanced_flow.compile()
planner.responses = [
[
{
"type": "sub_agent",
"sub_agent": "research_node"
},
{
"type": "tool",
"tool": "slack_tool"
}
# First planner pass triggers one sub-agent + one tool.
Decision(type="sub_agent", sub_agent="research lunch options"),
Decision(type="tool", tool="slack_tool"),
],
# 2nd llm invoke, after additional user input
[
{
"type": "sub_agent",
}
],
# 3rd llm invoke, after tool node completes
[],
# 4th llm invoke, after 1st sub agent node completes
# Second planner pass triggers another sub-agent.
[Decision(type="sub_agent", sub_agent="find vegetarian fallback")],
[],
# 5th llm invoke, after 2nd sub agent node completes
[
{
"type": "end"
}
],
[],
# Final pass decides to end.
[Decision(type="end", complete="order submitted")],
]
started = main_agent.ainvoke({"input": "help me get something for lunch"})
# provide more info after 2 seconds
await asyncio.sleep(2)
main_agent.apublish_to_channel("inbox", {"type": "user_input", "payload": "No spicy food please"})
started = asyncio.create_task(
main_agent.ainvoke(
{"input": "help me get something for lunch", "output": [], "done": None}
)
)
# External input can be injected while graph execution is in progress.
await asyncio.sleep(0.01)
await main_agent.apublish_to_channel(
"inbox", {"type": "user_input", "payload": "No spicy food please"}
)
result = await started
assert result == {
"input": "help me get something for lunch",
"output":[
]
}
"output": [
"user_input: No spicy food please",
"tool: tool completed for: slack_tool",
"sub_agent: research sub agent completed for: research lunch options",
"sub_agent: research sub agent completed for: find vegetarian fallback",
],
"done": "order submitted",
}