From 39e65a1a6231261d92d8740069ab150f9a3824f7 Mon Sep 17 00:00:00 2001 From: Vadym Barda Date: Wed, 29 Jan 2025 15:24:32 -0500 Subject: [PATCH] langgraph: expose tags in the metadata for streamed message chunks (#3238) --- libs/langgraph/langgraph/pregel/messages.py | 4 +++ libs/langgraph/tests/test_pregel.py | 37 +++++++++++++++++++ libs/langgraph/tests/test_pregel_async.py | 40 +++++++++++++++++++++ 3 files changed, 81 insertions(+) diff --git a/libs/langgraph/langgraph/pregel/messages.py b/libs/langgraph/langgraph/pregel/messages.py index 989f116dc..ffb07eddb 100644 --- a/libs/langgraph/langgraph/pregel/messages.py +++ b/libs/langgraph/langgraph/pregel/messages.py @@ -76,11 +76,15 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): chunk: Optional[ChatGenerationChunk] = None, run_id: UUID, parent_run_id: Optional[UUID] = None, + tags: Optional[list[str]] = None, **kwargs: Any, ) -> Any: if not isinstance(chunk, ChatGenerationChunk): return if meta := self.metadata.get(run_id): + filtered_tags = [t for t in (tags or []) if not t.startswith("seq:step")] + if filtered_tags: + meta[1]["tags"] = filtered_tags self._emit(meta, chunk.message) def on_llm_end( diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 4aca6edf2..eda5c7b51 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -29,6 +29,7 @@ from typing import ( import httpx import pytest +from langchain_core.language_models import GenericFakeChatModel from langchain_core.runnables import ( RunnableConfig, RunnableLambda, @@ -80,6 +81,7 @@ from tests.conftest import ( from tests.memory_assert import MemorySaverAssertCheckpointMetadata from tests.messages import ( _AnyIdAIMessage, + _AnyIdAIMessageChunk, _AnyIdHumanMessage, _AnyIdToolMessage, ) @@ -6288,3 +6290,38 @@ def test_named_tasks_functional() -> None: {"qux": "foo|bar|baz|custom_baz|qux"}, {"workflow": "foo|bar|baz|custom_baz|qux"}, ] + + +def test_tags_stream_mode_messages() -> None: + model = GenericFakeChatModel(messages=iter(["foo"]), tags=["meow"]) + graph = ( + StateGraph(MessagesState) + .add_node( + "call_model", lambda state: {"messages": model.invoke(state["messages"])} + ) + .add_edge(START, "call_model") + .compile() + ) + assert list( + graph.stream( + { + "messages": "hi", + }, + stream_mode="messages", + ) + ) == [ + ( + _AnyIdAIMessageChunk(content="foo"), + { + "langgraph_step": 1, + "langgraph_node": "call_model", + "langgraph_triggers": ["start:call_model"], + "langgraph_path": ("__pregel_pull", "call_model"), + "langgraph_checkpoint_ns": AnyStr("call_model:"), + "checkpoint_ns": AnyStr("call_model:"), + "ls_provider": "genericfakechatmodel", + "ls_model_type": "chat", + "tags": ["meow"], + }, + ) + ] diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index ac387b015..38627d12d 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -26,6 +26,7 @@ from uuid import UUID import httpx import pytest +from langchain_core.language_models import GenericFakeChatModel from langchain_core.runnables import ( RunnableConfig, RunnableLambda, @@ -82,6 +83,7 @@ from tests.memory_assert import ( ) from tests.messages import ( _AnyIdAIMessage, + _AnyIdAIMessageChunk, _AnyIdHumanMessage, _AnyIdToolMessage, ) @@ -7466,3 +7468,41 @@ async def test_overriding_injectable_args_with_async_task() -> None: return "OK" assert await main.ainvoke({}) == "OK" + + +async def test_tags_stream_mode_messages() -> None: + model = GenericFakeChatModel(messages=iter(["foo"]), tags=["meow"]) + + async def call_model(state, config): + return {"messages": await model.ainvoke(state["messages"], config)} + + graph = ( + StateGraph(MessagesState) + .add_node(call_model) + .add_edge(START, "call_model") + .compile() + ) + assert [ + c + async for c in graph.astream( + { + "messages": "hi", + }, + stream_mode="messages", + ) + ] == [ + ( + _AnyIdAIMessageChunk(content="foo"), + { + "langgraph_step": 1, + "langgraph_node": "call_model", + "langgraph_triggers": ["start:call_model"], + "langgraph_path": ("__pregel_pull", "call_model"), + "langgraph_checkpoint_ns": AnyStr("call_model:"), + "checkpoint_ns": AnyStr("call_model:"), + "ls_provider": "genericfakechatmodel", + "ls_model_type": "chat", + "tags": ["meow"], + }, + ) + ]