mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-17 21:25:46 +02:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2c2575370e |
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user