Compare commits

...
Author SHA1 Message Date
William Fu-Hinthorn 2c2575370e listen up 2025-12-17 23:06:05 +09:00
3 changed files with 167 additions and 1 deletions
+90
View File
@@ -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,
+19 -1
View File
@@ -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,
+58
View File
@@ -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