diff --git a/libs/langgraph/langgraph/pregel/_messages.py b/libs/langgraph/langgraph/pregel/_messages.py index 5f7d3af16..ed48fe0fc 100644 --- a/libs/langgraph/langgraph/pregel/_messages.py +++ b/libs/langgraph/langgraph/pregel/_messages.py @@ -10,7 +10,7 @@ from typing import ( from uuid import UUID, uuid4 from langchain_core.callbacks import BaseCallbackHandler -from langchain_core.messages import BaseMessage +from langchain_core.messages import BaseMessage, ToolMessage from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, LLMResult from pydantic import BaseModel @@ -303,6 +303,35 @@ class StreamMessagesHandlerV2(StreamMessagesHandler, _V2StreamingCallbackHandler super().__init__(stream, subgraphs, parent_ns=parent_ns) self._streamed_run_ids: set[UUID] = set() + def _find_and_emit_messages(self, meta: Meta, response: Any) -> None: + """Like the v1 handler, but skip ToolMessage from node outputs. + + Tool results belong on the tools channel / state in v3; v2-flagged streams + must not replay finalized ToolMessages as chat tokens (see MessagesTransformer). + Legacy v1-only `stream_mode="messages"` still emits ToolMessages (see subgraph + streaming tests). + """ + if isinstance(response, BaseMessage) and not isinstance(response, ToolMessage): + self._emit(meta, response, dedupe=True) + elif isinstance(response, Sequence): + for value in response: + if isinstance(value, BaseMessage) and not isinstance( + value, ToolMessage + ): + self._emit(meta, value, dedupe=True) + else: + for value in _state_values(response): + if isinstance(value, BaseMessage) and not isinstance( + value, ToolMessage + ): + self._emit(meta, value, dedupe=True) + elif isinstance(value, Sequence): + for item in value: + if isinstance(item, BaseMessage) and not isinstance( + item, ToolMessage + ): + self._emit(meta, item, dedupe=True) + def on_llm_end( self, response: LLMResult, diff --git a/libs/langgraph/langgraph/stream/transformers.py b/libs/langgraph/langgraph/stream/transformers.py index 0ce420af7..6159634a2 100644 --- a/libs/langgraph/langgraph/stream/transformers.py +++ b/libs/langgraph/langgraph/stream/transformers.py @@ -8,7 +8,7 @@ from langchain_core.language_models.chat_model_stream import ( AsyncChatModelStream, ChatModelStream, ) -from langchain_core.messages import AIMessageChunk, BaseMessage +from langchain_core.messages import AIMessageChunk, BaseMessage, ToolMessage from langchain_protocol.protocol import MessagesData from typing_extensions import NotRequired, TypedDict @@ -203,6 +203,7 @@ class MessagesTransformer(StreamTransformer): # Correlate protocol events back to a ChatModelStream by run_id # (attached to the event's metadata by StreamMessagesHandler). self._by_run: dict[str, ChatModelStream] = {} + self._ignored_runs: set[str] = set() self._pump_fn: Callable[[], bool] | None = None self._apump_fn: Callable[[], Awaitable[bool]] | None = None # Cached as a list once for cheap equality with the protocol @@ -276,8 +277,10 @@ class MessagesTransformer(StreamTransformer): self._route_protocol_event( cast("MessagesData", payload), run_id=run_id, node=node ) - elif isinstance(payload, BaseMessage) and not isinstance( - payload, AIMessageChunk + elif ( + isinstance(payload, BaseMessage) + and not isinstance(payload, AIMessageChunk) + and not isinstance(payload, ToolMessage) ): self._route_whole_message(payload, node=node) # Legacy AIMessageChunk tuples (from on_llm_new_token) are ignored; @@ -295,6 +298,11 @@ class MessagesTransformer(StreamTransformer): ) -> None: event_type = event.get("event") if event_type == "message-start": + # Tool results are exposed on the tools projection and state + # snapshots; run.messages is the chat-token projection. + if event.get("role") == "tool": + self._ignored_runs.add(run_id) + return message_id = event.get("message_id") stream = self._make_stream( namespace=[], @@ -304,6 +312,9 @@ class MessagesTransformer(StreamTransformer): self._by_run[run_id] = stream self._log.push(stream) stream.dispatch(event) + elif run_id in self._ignored_runs: + if event_type == "message-finish": + self._ignored_runs.discard(run_id) elif run_id in self._by_run: stream = self._by_run[run_id] stream.dispatch(event) @@ -319,12 +330,14 @@ class MessagesTransformer(StreamTransformer): def finalize(self) -> None: """Clear any routing state — streams close themselves via `message-finish`.""" self._by_run.clear() + self._ignored_runs.clear() def fail(self, err: BaseException) -> None: """Propagate run error to any streams still open when the graph fails.""" for stream in list(self._by_run.values()): stream.fail(err) self._by_run.clear() + self._ignored_runs.clear() SubgraphStatus = Literal["started", "completed", "failed", "interrupted", "drained"] diff --git a/libs/langgraph/tests/test_stream_messages_transformer.py b/libs/langgraph/tests/test_stream_messages_transformer.py index 3f37bb72c..6e13ef47e 100644 --- a/libs/langgraph/tests/test_stream_messages_transformer.py +++ b/libs/langgraph/tests/test_stream_messages_transformer.py @@ -12,7 +12,7 @@ from langchain_core.language_models.chat_model_stream import ( AsyncChatModelStream, ChatModelStream, ) -from langchain_core.messages import AIMessage, AIMessageChunk +from langchain_core.messages import AIMessage, AIMessageChunk, ToolMessage from langchain_core.runnables import RunnableConfig from typing_extensions import TypedDict @@ -213,6 +213,23 @@ class TestProtocolEventRouting: log.close() assert _unstamped(log._items) == [] + def test_tool_role_protocol_events_are_ignored(self) -> None: + t, log = _make_sync_transformer() + for evt in [ + {"event": "message-start", "role": "tool", "message_id": "tool-msg-1"}, + { + "event": "content-block-delta", + "index": 0, + "content_block": {"type": "text", "text": "[]"}, + }, + {"event": "message-finish", "reason": "stop"}, + ]: + t.process(_proto_event(evt, run_id="tool-run")) + + log.close() + assert _unstamped(log._items) == [] + assert t._ignored_runs == set() + def test_concurrent_streams_routed_by_run_id(self) -> None: t, log = _make_sync_transformer() life_a = _lifecycle(text="aaaa", message_id="run-a") @@ -273,6 +290,29 @@ class TestWholeMessageFallback: assert stream.done assert stream.output.text == "the full answer" + def test_whole_tool_message_is_ignored(self) -> None: + t, log = _make_sync_transformer() + t.process( + { + "type": "event", + "method": "messages", + "params": { + "namespace": [], + "timestamp": TS, + "data": ( + ToolMessage( + content="[]", + id="tool-msg-1", + tool_call_id="call_1", + ), + {"langgraph_node": "tools"}, + ), + }, + } + ) + log.close() + assert _unstamped(log._items) == [] + def test_whole_message_has_full_lifecycle(self) -> None: t, log = _make_sync_transformer() t.process(_whole_msg("full")) @@ -849,6 +889,23 @@ class TestStreamMessagesHandlerV2Unit: assert emitted == [] + def test_on_chain_end_does_not_emit_tool_messages(self) -> None: + from uuid import uuid4 + + from langgraph.pregel._messages import StreamMessagesHandlerV2 + + emitted: list[Any] = [] + handler = StreamMessagesHandlerV2(emitted.append, subgraphs=False) + run_id = uuid4() + handler.metadata[run_id] = ((), {"langgraph_node": "tools"}) + + handler.on_chain_end( + {"messages": [ToolMessage(content="[]", tool_call_id="call_1")]}, + run_id=run_id, + ) + + assert emitted == [] + def test_on_llm_end_dedupes_when_final_message_id_differs(self) -> None: """A streamed v2 message should not be emitted again from the final AIMessage fallback when its final id does not match `message-start`.""" diff --git a/libs/prebuilt/langgraph/prebuilt/_tool_call_transformer.py b/libs/prebuilt/langgraph/prebuilt/_tool_call_transformer.py index 696dcdcab..37ed1b8e8 100644 --- a/libs/prebuilt/langgraph/prebuilt/_tool_call_transformer.py +++ b/libs/prebuilt/langgraph/prebuilt/_tool_call_transformer.py @@ -5,12 +5,42 @@ from __future__ import annotations from collections.abc import Awaitable, Callable from typing import Any +from langchain_core.messages import ToolMessage from langgraph.stream._types import ProtocolEvent, StreamTransformer from langgraph.stream.stream_channel import StreamChannel from langgraph.prebuilt._tool_call_stream import ToolCallStream +def _is_serialized_tool_message(value: Any) -> bool: + """Detect a serialized LangChain `ToolMessage` payload. + + Example: + { + "lc": 1, + "type": "constructor", + "id": ["langchain_core", "messages", "ToolMessage"], + "kwargs": {"content": "raw tool result", "tool_call_id": "call_1"}, + } + """ + return ( + isinstance(value, dict) + and value.get("type") == "constructor" + and isinstance(value.get("id"), list) + and value["id"][-1] == "ToolMessage" + ) + + +def _normalize_tool_output(output: Any) -> Any: + if isinstance(output, ToolMessage): + return output.content + if _is_serialized_tool_message(output): + kwargs = output.get("kwargs") + if isinstance(kwargs, dict): + return kwargs.get("content") + return output + + class ToolCallTransformer(StreamTransformer): """Project `tools` channel events into `ToolCallStream` handles. @@ -109,7 +139,7 @@ class ToolCallTransformer(StreamTransformer): elif event_type == "tool-finished": stream = self._active.pop(tool_call_id, None) if stream is not None: - stream._finish(data.get("output")) + stream._finish(_normalize_tool_output(data.get("output"))) elif event_type == "tool-error": stream = self._active.pop(tool_call_id, None) if stream is not None: diff --git a/libs/prebuilt/tests/test_tool_call_transformer.py b/libs/prebuilt/tests/test_tool_call_transformer.py index 1287d4a49..ae8fedd67 100644 --- a/libs/prebuilt/tests/test_tool_call_transformer.py +++ b/libs/prebuilt/tests/test_tool_call_transformer.py @@ -6,7 +6,7 @@ import time from typing import Annotated, Any import pytest -from langchain_core.messages import AIMessage +from langchain_core.messages import AIMessage, ToolMessage from langchain_core.tools import tool from langgraph.constants import END, START from langgraph.graph import StateGraph @@ -128,6 +128,42 @@ class TestToolCallTransformerUnit: assert stream.error is None assert "tc1" not in transformer._active + def test_finish_unwraps_tool_message_output(self) -> None: + mux, transformer = _mux() + mux.push(_tool_event("tool-started", "tc1", tool_name="echo")) + stream = transformer._active["tc1"] + mux.push( + _tool_event( + "tool-finished", + "tc1", + output=ToolMessage(content="done", tool_call_id="tc1"), + ) + ) + assert stream.completed is True + assert stream.output == "done" + + def test_finish_unwraps_serialized_tool_message_output(self) -> None: + mux, transformer = _mux() + mux.push(_tool_event("tool-started", "tc1", tool_name="echo")) + stream = transformer._active["tc1"] + mux.push( + _tool_event( + "tool-finished", + "tc1", + output={ + "lc": 1, + "type": "constructor", + "id": ["langchain_core", "messages", "ToolMessage"], + "kwargs": { + "content": "serialized done", + "tool_call_id": "tc1", + }, + }, + ) + ) + assert stream.completed is True + assert stream.output == "serialized done" + def test_error_closes_stream(self) -> None: mux, transformer = _mux() mux.push(_tool_event("tool-started", "tc1", tool_name="boom"))