From 2c2575370e492552ba80d086a05e3b468f315849 Mon Sep 17 00:00:00 2001 From: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com> Date: Wed, 17 Dec 2025 23:06:05 +0900 Subject: [PATCH] listen up --- libs/langgraph/langgraph/pregel/main.py | 90 +++++++++++++++++++++ libs/langgraph/langgraph/pregel/protocol.py | 20 ++++- libs/langgraph/tests/test_pregel.py | 58 +++++++++++++ 3 files changed, 167 insertions(+), 1 deletion(-) diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index 37e8125f9..e033618cc 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -654,6 +654,8 @@ class Pregel( name: str = "LangGraph", **deprecated_kwargs: Unpack[DeprecatedKwargs], ) -> None: + root_listeners = deprecated_kwargs.pop("_root_listeners", None) + aroot_listeners = deprecated_kwargs.pop("_aroot_listeners", None) if ( config_type := deprecated_kwargs.get("config_type", MISSING) ) is not MISSING: @@ -698,6 +700,14 @@ class Pregel( self.config = config self.trigger_to_nodes = trigger_to_nodes or {} self.name = name + self._root_listeners = cast( + "tuple[Callable[[Any], Any] | None, Callable[[Any], Any] | None, Callable[[Any], Any] | None] | None", + root_listeners, + ) + self._aroot_listeners = cast( + "tuple[Callable[[Any], Awaitable[Any]] | None, Callable[[Any], Awaitable[Any]] | None, Callable[[Any], Awaitable[Any]] | None] | None", + aroot_listeners, + ) if auto_validate: self.validate() @@ -789,6 +799,26 @@ class Pregel( {"config": merge_configs(self.config, config, cast(RunnableConfig, kwargs))} ) + def with_listeners( + self, + *, + on_start: Callable[[Any], Any] | None = None, + on_end: Callable[[Any], Any] | None = None, + on_error: Callable[[Any], Any] | None = None, + ) -> Self: + """Create a copy of the Pregel object with updated sync listeners.""" + return self.copy({"_root_listeners": (on_start, on_end, on_error)}) + + def with_alisteners( + self, + *, + on_start: Callable[[Any], Awaitable[Any]] | None = None, + on_end: Callable[[Any], Awaitable[Any]] | None = None, + on_error: Callable[[Any], Awaitable[Any]] | None = None, + ) -> Self: + """Create a copy of the Pregel object with updated async listeners.""" + return self.copy({"_aroot_listeners": (on_start, on_end, on_error)}) + def validate(self) -> Self: validate_graph( self.nodes, @@ -2497,6 +2527,26 @@ class Pregel( stream = SyncQueue() config = ensure_config(self.config, config) + if self._root_listeners is not None: + on_start, on_end, on_error = self._root_listeners + from langchain_core.tracers.root_listeners import RootListenersTracer + + tracer = RootListenersTracer( + config=config, + on_start=on_start, + on_end=on_end, + on_error=on_error, + ) + existing_callbacks = config.get("callbacks") + config = config.copy() + if existing_callbacks is None: + config["callbacks"] = [tracer] + elif isinstance(existing_callbacks, list): + config["callbacks"] = [*existing_callbacks, tracer] + else: + callbacks_manager = existing_callbacks.copy() + callbacks_manager.add_handler(tracer, inherit=True) + config["callbacks"] = callbacks_manager callback_manager = get_callback_manager_for_config(config) run_manager = callback_manager.on_chain_start( None, @@ -2776,6 +2826,46 @@ class Pregel( ) config = ensure_config(self.config, config) + if self._root_listeners is not None: + on_start, on_end, on_error = self._root_listeners + from langchain_core.tracers.root_listeners import RootListenersTracer + + sync_tracer = RootListenersTracer( + config=config, + on_start=on_start, + on_end=on_end, + on_error=on_error, + ) + existing_callbacks = config.get("callbacks") + config = config.copy() + if existing_callbacks is None: + config["callbacks"] = [sync_tracer] + elif isinstance(existing_callbacks, list): + config["callbacks"] = [*existing_callbacks, sync_tracer] + else: + callbacks_manager = existing_callbacks.copy() + callbacks_manager.add_handler(sync_tracer, inherit=True) + config["callbacks"] = callbacks_manager + if self._aroot_listeners is not None: + on_start, on_end, on_error = self._aroot_listeners + from langchain_core.tracers.root_listeners import AsyncRootListenersTracer + + async_tracer = AsyncRootListenersTracer( + config=config, + on_start=on_start, + on_end=on_end, + on_error=on_error, + ) + existing_callbacks = config.get("callbacks") + config = config.copy() + if existing_callbacks is None: + config["callbacks"] = [async_tracer] + elif isinstance(existing_callbacks, list): + config["callbacks"] = [*existing_callbacks, async_tracer] + else: + callbacks_manager = existing_callbacks.copy() + callbacks_manager.add_handler(async_tracer, inherit=True) + config["callbacks"] = callbacks_manager callback_manager = get_async_callback_manager_for_config(config) run_manager = await callback_manager.on_chain_start( None, diff --git a/libs/langgraph/langgraph/pregel/protocol.py b/libs/langgraph/langgraph/pregel/protocol.py index c9bf6e5ff..44de197f5 100644 --- a/libs/langgraph/langgraph/pregel/protocol.py +++ b/libs/langgraph/langgraph/pregel/protocol.py @@ -1,7 +1,7 @@ from __future__ import annotations from abc import abstractmethod -from collections.abc import AsyncIterator, Callable, Iterator, Sequence +from collections.abc import AsyncIterator, Awaitable, Callable, Iterator, Sequence from typing import Any, Generic, cast from langchain_core.runnables import Runnable, RunnableConfig @@ -20,6 +20,24 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou self, config: RunnableConfig | None = None, **kwargs: Any ) -> Self: ... + @abstractmethod + def with_listeners( + self, + *, + on_start: Callable[[Any], Any] | None = None, + on_end: Callable[[Any], Any] | None = None, + on_error: Callable[[Any], Any] | None = None, + ) -> Self: ... + + @abstractmethod + def with_alisteners( + self, + *, + on_start: Callable[[Any], Awaitable[Any]] | None = None, + on_end: Callable[[Any], Awaitable[Any]] | None = None, + on_error: Callable[[Any], Awaitable[Any]] | None = None, + ) -> Self: ... + @abstractmethod def get_graph( self, diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 97928ec6f..2473605f4 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -252,6 +252,64 @@ def test_context_json_schema() -> None: } +def test_pregel_with_listeners() -> None: + chain = NodeBuilder().subscribe_only("input").do(lambda x: x + 1).write_to("output") + + app = Pregel( + nodes={"one": chain}, + channels={ + "input": LastValue(int), + "output": LastValue(int), + }, + input_channels="input", + output_channels="output", + ) + + events: list[tuple[str, list[str] | None]] = [] + + def on_start(_: Any, config: RunnableConfig) -> None: + events.append(("start", list(config.get("tags") or []))) + + def on_end(_: Any, config: RunnableConfig) -> None: + events.append(("end", list(config.get("tags") or []))) + + out = app.with_listeners(on_start=on_start, on_end=on_end).invoke( + 1, {"tags": ["tag-a"]} + ) + + assert out == 2 + assert events == [("start", ["tag-a"]), ("end", ["tag-a"])] + + +async def test_pregel_with_alisteners() -> None: + chain = NodeBuilder().subscribe_only("input").do(lambda x: x + 1).write_to("output") + + app = Pregel( + nodes={"one": chain}, + channels={ + "input": LastValue(int), + "output": LastValue(int), + }, + input_channels="input", + output_channels="output", + ) + + events: list[tuple[str, list[str] | None]] = [] + + async def on_start(_: Any, config: RunnableConfig) -> None: + events.append(("start", list(config.get("tags") or []))) + + async def on_end(_: Any, config: RunnableConfig) -> None: + events.append(("end", list(config.get("tags") or []))) + + out = await app.with_alisteners(on_start=on_start, on_end=on_end).ainvoke( + 1, {"tags": ["tag-a"]} + ) + + assert out == 2 + assert events == [("start", ["tag-a"]), ("end", ["tag-a"])] + + def test_node_schemas_custom_output() -> None: class State(TypedDict): hello: str