This commit is contained in:
Quanzheng Long
2026-03-14 10:29:42 -07:00
parent 56e9fe1b10
commit b3474d2db1
2 changed files with 244 additions and 6 deletions
@@ -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 []
@@ -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())