From 4b2167dd2590bd668720ff5195c114b2819f22eb Mon Sep 17 00:00:00 2001 From: Quanzheng Long Date: Thu, 12 Mar 2026 13:13:08 -0700 Subject: [PATCH] 1stpass --- .../langgraph/graph_engine/__init__.py | 13 + .../langgraph/langgraph/graph_engine/state.py | 352 ++++++++++++++++++ .../tests/advanced-graph/test_sub_agents.py | 256 ++++++++----- 3 files changed, 518 insertions(+), 103 deletions(-) create mode 100644 libs/langgraph/langgraph/graph_engine/__init__.py create mode 100644 libs/langgraph/langgraph/graph_engine/state.py diff --git a/libs/langgraph/langgraph/graph_engine/__init__.py b/libs/langgraph/langgraph/graph_engine/__init__.py new file mode 100644 index 000000000..5e85971d7 --- /dev/null +++ b/libs/langgraph/langgraph/graph_engine/__init__.py @@ -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", +) diff --git a/libs/langgraph/langgraph/graph_engine/state.py b/libs/langgraph/langgraph/graph_engine/state.py new file mode 100644 index 000000000..f6fc7636b --- /dev/null +++ b/libs/langgraph/langgraph/graph_engine/state.py @@ -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 == "": + 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}") diff --git a/libs/langgraph/tests/advanced-graph/test_sub_agents.py b/libs/langgraph/tests/advanced-graph/test_sub_agents.py index 583f648af..a5704d9d0 100644 --- a/libs/langgraph/tests/advanced-graph/test_sub_agents.py +++ b/libs/langgraph/tests/advanced-graph/test_sub_agents.py @@ -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":[ - - ] - } \ No newline at end of file + "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", + }