mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 04:25:08 +02:00
add test
This commit is contained in:
@@ -1,7 +1,8 @@
|
||||
import pytest
|
||||
from langchain_core.callbacks import AsyncCallbackManager
|
||||
from langchain_core.callbacks import AsyncCallbackManager, BaseCallbackHandler
|
||||
|
||||
from langgraph._internal._config import get_async_callback_manager_for_config
|
||||
from langgraph.graph import StateGraph
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
@@ -17,3 +18,21 @@ def test_new_async_manager_merges_tags_with_config() -> None:
|
||||
config = {"callbacks": None, "tags": ["a"]}
|
||||
manager = get_async_callback_manager_for_config(config, tags=["b"])
|
||||
assert manager.inheritable_tags == ["a", "b"]
|
||||
|
||||
|
||||
async def test_with_config_callbacks_preserved_in_astream_events() -> None:
|
||||
class TrackingCallback(BaseCallbackHandler):
|
||||
def __init__(self) -> None:
|
||||
self.called = False
|
||||
|
||||
def on_chain_start(self, *args, **kwargs) -> None:
|
||||
self.called = True
|
||||
|
||||
builder = StateGraph(dict)
|
||||
builder.add_node("node", lambda state: state)
|
||||
builder.add_edge("__start__", "node")
|
||||
cb = TrackingCallback()
|
||||
graph = builder.compile().with_config({"callbacks": [cb]})
|
||||
async for _ in graph.astream_events({}, version="v2"):
|
||||
pass
|
||||
assert cb.called
|
||||
|
||||
Reference in New Issue
Block a user