mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 21:45:08 +02:00
langgraph: expose tags in the metadata for streamed message chunks (#3238)
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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"],
|
||||
},
|
||||
)
|
||||
]
|
||||
|
||||
@@ -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"],
|
||||
},
|
||||
)
|
||||
]
|
||||
|
||||
Reference in New Issue
Block a user