mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-27 20:15:00 +02:00
more
This commit is contained in:
@@ -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())
|
||||
Reference in New Issue
Block a user