From b3474d2db1adee6db1cac872c707035c1e552b9e Mon Sep 17 00:00:00 2001 From: Quanzheng Long Date: Sat, 14 Mar 2026 10:29:42 -0700 Subject: [PATCH] more --- .../langgraph/advanced_graph/state.py | 76 +++++++- .../benchmark_stategraph_vs_advancedgraph.py | 174 ++++++++++++++++++ 2 files changed, 244 insertions(+), 6 deletions(-) create mode 100644 libs/langgraph/tests/advanced-graph/benchmark_stategraph_vs_advancedgraph.py diff --git a/libs/langgraph/langgraph/advanced_graph/state.py b/libs/langgraph/langgraph/advanced_graph/state.py index efb6f4701..4d5a2198e 100644 --- a/libs/langgraph/langgraph/advanced_graph/state.py +++ b/libs/langgraph/langgraph/advanced_graph/state.py @@ -1,8 +1,11 @@ from __future__ import annotations +import atexit import asyncio +import os import inspect import threading +from concurrent.futures import ThreadPoolExecutor from collections.abc import Callable, Coroutine, Sequence from dataclasses import dataclass from datetime import timedelta @@ -38,6 +41,31 @@ class AnyOfCondition: WaitCondition = ChannelCondition | TimerCondition +_EXECUTOR_LOCK = threading.Lock() +_EXECUTOR: ThreadPoolExecutor | None = None + + +def _advanced_graph_executor() -> ThreadPoolExecutor: + global _EXECUTOR + with _EXECUTOR_LOCK: + if _EXECUTOR is None: + worker_count = int(os.getenv("LANGGRAPH_ADVANCED_GRAPH_PY_THREADS", "256")) + worker_count = max(worker_count, 1) + _EXECUTOR = ThreadPoolExecutor( + max_workers=worker_count, + thread_name_prefix="langgraph-advanced-py", + ) + atexit.register(_shutdown_advanced_graph_executor) + return _EXECUTOR + + +def _shutdown_advanced_graph_executor() -> None: + global _EXECUTOR + with _EXECUTOR_LOCK: + if _EXECUTOR is not None: + _EXECUTOR.shutdown(wait=False, cancel_futures=False) + _EXECUTOR = None + class WaitRequested(Exception): def __init__(self, payload: dict[str, Any]) -> None: @@ -207,7 +235,9 @@ class _GraphEngineRun: self.context = Context(self) async def run(self, initial_state: StateT) -> StateT: - result_obj = await asyncio.to_thread( + loop = asyncio.get_running_loop() + result_obj = await loop.run_in_executor( + _advanced_graph_executor(), self._rust_engine.run_graph_py, self._entry_point, self._finish_point, @@ -218,7 +248,13 @@ class _GraphEngineRun: return cast(StateT, self._state) async def publish(self, channel: str, value: Any) -> None: - await asyncio.to_thread(self._publish_sync, channel, value) + loop = asyncio.get_running_loop() + await loop.run_in_executor( + _advanced_graph_executor(), + self._publish_sync, + channel, + value, + ) def publish_nowait(self, channel: str, value: Any) -> None: self._publish_sync(channel, value) @@ -232,7 +268,12 @@ class _GraphEngineRun: "value": value, } if isinstance(target, TimerCondition): - return await asyncio.to_thread(self._rust_engine.wait_timer, target.seconds) + loop = asyncio.get_running_loop() + return await loop.run_in_executor( + _advanced_graph_executor(), + self._rust_engine.wait_timer, + target.seconds, + ) if isinstance(target, AnyOfCondition): return await self._wait_for_any_of(target) raise ValueError(f"Unsupported wait condition type: {type(target)!r}") @@ -240,7 +281,13 @@ class _GraphEngineRun: async def _wait_for_channel_values(self, channel: str, n: int) -> Any: if n < 1: raise ValueError("wait_for count `n` must be >= 1") - event = await asyncio.to_thread(self._rust_engine.wait_channel, channel, n) + loop = asyncio.get_running_loop() + event = await loop.run_in_executor( + _advanced_graph_executor(), + self._rust_engine.wait_channel, + channel, + n, + ) return event["value"] async def _wait_for_any_of(self, condition: AnyOfCondition) -> Any: @@ -249,7 +296,12 @@ class _GraphEngineRun: payload = { "conditions": [_condition_to_rust(cond) for cond in condition.conditions] } - return await asyncio.to_thread(self._rust_engine.wait_any_of_obj, payload) + loop = asyncio.get_running_loop() + return await loop.run_in_executor( + _advanced_graph_executor(), + self._rust_engine.wait_any_of_obj, + payload, + ) def _publish_sync(self, channel: str, value: Any) -> None: self._rust_engine.publish_obj(channel, value) @@ -265,7 +317,9 @@ class _GraphEngineRun: try: result = _invoke_node(node, self.context, node_input, state) if inspect.isawaitable(result): - result = asyncio.run(cast(Coroutine[Any, Any, Any], result)) + result = self._run_awaitable_in_worker( + cast(Coroutine[Any, Any, Any], result) + ) except WaitRequested as suspend: return {"suspend": suspend.payload} finally: @@ -296,6 +350,16 @@ class _GraphEngineRun: self._local.resume_event = None return event + def _run_awaitable_in_worker(self, awaitable: Coroutine[Any, Any, Any]) -> Any: + loop = cast( + asyncio.AbstractEventLoop | None, + getattr(self._local, "worker_loop", None), + ) + if loop is None or loop.is_closed(): + loop = asyncio.new_event_loop() + self._local.worker_loop = loop + return loop.run_until_complete(awaitable) + def _normalize_result_to_sends(result: Any, *, default_input: Any) -> list[Send]: if result is None: return [] diff --git a/libs/langgraph/tests/advanced-graph/benchmark_stategraph_vs_advancedgraph.py b/libs/langgraph/tests/advanced-graph/benchmark_stategraph_vs_advancedgraph.py new file mode 100644 index 000000000..58b2d7259 --- /dev/null +++ b/libs/langgraph/tests/advanced-graph/benchmark_stategraph_vs_advancedgraph.py @@ -0,0 +1,174 @@ +from __future__ import annotations + +import asyncio +import os +import time +from typing import Any + +# Configure advanced-graph runtime pools for this benchmark run. +os.environ["LANGGRAPH_RUN_POOL_SIZE"] = "10" +os.environ["LANGGRAPH_NODE_POOL_SIZE"] = "100" + +from langgraph.advanced_graph import ( + AdvancedStateGraph, + channel_condition, + timer_condition, +) +from langgraph.graph import END, START, StateGraph +from langgraph.types import Command, Send + +RUNS = 100 +MIDDLE_COUNT = 10 +SLEEP_SECONDS = 2.0 +STATE_BYTES = 10 * 1024 + + +def make_initial_state() -> dict[str, Any]: + return {"payload": "x" * STATE_BYTES, "done": False} + + +def build_advanced_parallel() -> Any: + graph: AdvancedStateGraph[dict[str, Any]] = AdvancedStateGraph(dict) + done_channel = "__bench_done_channel" + graph.add_async_channel(done_channel, str) + + async def start_node(state: dict[str, Any]) -> Command: + _ = state + sends = [Send(f"middle_{i}", None) for i in range(MIDDLE_COUNT)] + sends.append(Send("end_node", None)) + return Command(goto=sends) + + async def end_node(ctx: Any, state: dict[str, Any]) -> dict[str, Any]: + await ctx.wait_for(channel_condition(done_channel, n=MIDDLE_COUNT)) + out = dict(state) + out["done"] = True + return out + + graph.add_entry_node(start_node) + for i in range(MIDDLE_COUNT): + async def middle_node(ctx: Any, state: dict[str, Any], idx: int = i) -> None: + _ = idx + _ = state + await ctx.wait_for(timer_condition(seconds=SLEEP_SECONDS)) + ctx.publish_to_channel(done_channel, "done") + + graph.add_node(f"middle_{i}", middle_node) + graph.add_finish_node(end_node) + return graph.compile() + + +def build_advanced_sequential() -> Any: + graph: AdvancedStateGraph[dict[str, Any]] = AdvancedStateGraph(dict) + + async def start_node(state: dict[str, Any]) -> Command: + _ = state + return Command(goto=Send("middle_0", None)) + + async def end_node(state: dict[str, Any]) -> dict[str, Any]: + out = dict(state) + out["done"] = True + return out + + graph.add_entry_node(start_node) + def make_middle(target: str): + async def middle_node(ctx: Any, state: dict[str, Any]) -> Command: + _ = state + await ctx.wait_for(timer_condition(seconds=SLEEP_SECONDS)) + return Command(goto=Send(target, None)) + + return middle_node + + for i in range(MIDDLE_COUNT): + next_name = "end_node" if i == MIDDLE_COUNT - 1 else f"middle_{i+1}" + graph.add_node(f"middle_{i}", make_middle(next_name)) + graph.add_finish_node(end_node) + return graph.compile() + + +def build_stategraph_parallel() -> Any: + graph = StateGraph(dict) + + async def start_node(state: dict[str, Any]) -> None: + _ = state + + async def end_node(state: dict[str, Any]) -> dict[str, Any]: + out = dict(state) + out["done"] = True + return out + + graph.add_node("start_node", start_node) + for i in range(MIDDLE_COUNT): + async def middle_node(state: dict[str, Any], idx: int = i) -> None: + _ = idx + _ = state + await asyncio.sleep(SLEEP_SECONDS) + + graph.add_node(f"middle_{i}", middle_node) + graph.add_node("end_node", end_node) + + graph.add_edge(START, "start_node") + for i in range(MIDDLE_COUNT): + graph.add_edge("start_node", f"middle_{i}") + graph.add_edge(f"middle_{i}", "end_node") + graph.add_edge("end_node", END) + return graph.compile() + + +def build_stategraph_sequential() -> Any: + graph = StateGraph(dict) + + async def start_node(state: dict[str, Any]) -> None: + _ = state + + async def end_node(state: dict[str, Any]) -> dict[str, Any]: + out = dict(state) + out["done"] = True + return out + + graph.add_node("start_node", start_node) + for i in range(MIDDLE_COUNT): + async def middle_node(state: dict[str, Any], idx: int = i) -> None: + _ = idx + _ = state + await asyncio.sleep(SLEEP_SECONDS) + + graph.add_node(f"middle_{i}", middle_node) + graph.add_node("end_node", end_node) + + graph.add_edge(START, "start_node") + graph.add_edge("start_node", "middle_0") + for i in range(MIDDLE_COUNT - 1): + graph.add_edge(f"middle_{i}", f"middle_{i+1}") + graph.add_edge(f"middle_{MIDDLE_COUNT - 1}", "end_node") + graph.add_edge("end_node", END) + return graph.compile() + + +async def run_benchmark(name: str, compiled: Any) -> float: + started = time.perf_counter() + tasks = [asyncio.create_task(compiled.ainvoke(make_initial_state())) for _ in range(RUNS)] + results = await asyncio.gather(*tasks) + elapsed = time.perf_counter() - started + if not all(item.get("done") is True for item in results): + raise RuntimeError(f"{name} produced unfinished runs") + return elapsed + + +async def main() -> None: + suites = [ + ("advanced-graph-parallel", build_advanced_parallel()), + ("advanced-graph-sequential", build_advanced_sequential()), + ("state-graph-parallel", build_stategraph_parallel()), + ("state-graph-sequential", build_stategraph_sequential()), + ] + print( + f"runs={RUNS}, middle_nodes={MIDDLE_COUNT}, sleep={SLEEP_SECONDS}s, " + f"state_bytes={STATE_BYTES}" + ) + for name, compiled in suites: + elapsed = await run_benchmark(name, compiled) + print(f"{name}: {elapsed:.3f}s") + + +if __name__ == "__main__": + asyncio.run(main())