From acaa76754254ec1e9a8ef885b4a4f98aaa50bd59 Mon Sep 17 00:00:00 2001 From: Nick Hollon Date: Fri, 17 Apr 2026 16:13:11 -0400 Subject: [PATCH] feat(langgraph): route invoke messages through v2 via StreamingHandler MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit When `StreamingHandler(graph).stream()` is used, content-block (v2) protocol events now flow through `stream_mode="messages"` for every `model.invoke()` call inside a node — with no node-level code changes. Adds `StreamMessagesHandlerV2`, a `StreamMessagesHandler` subclass that also inherits `_V2StreamingCallbackHandler` from langchain-core. The marker base flips `BaseChatModel.invoke` to drive the protocol event generator (firing `on_stream_event`) instead of `_stream` (firing `on_llm_new_token`). The handler inherits `on_stream_event` from the parent — events forward onto the messages channel unchanged — and overrides `on_llm_new_token` to no-op so a node calling `model.stream()` directly on a v2-flagged run can't leak AIMessageChunks onto the same channel. Opt-in is scoped to `StreamingHandler`: it merges a new internal `CONFIG_KEY_STREAM_MESSAGES_V2=True` into `config.configurable` before dispatching to `graph.stream` / `graph.astream`. Pregel reads the flag at handler-construction time in both sync and async stream paths and attaches the v2 subclass only when set. Direct `graph.stream(stream_mode="messages")` callers keep the v1 `(AIMessageChunk, metadata)` shape — confirmed by a regression test. Existing dedupe between the streamed v2 lifecycle and a node returning the same assembled `AIMessage` transfers for free: the handler populates `self.seen` from `message-start` events (via the inherited `on_stream_event` body), and `on_chain_end`'s `_find_and_emit_messages` already gates on `seen` — so an invoking node surfaces as exactly one `ChatModelStream`, not two. Test coverage in `tests/test_stream_messages_transformer.py`: - `TestEndToEndV2Invoke` — node calling `model.invoke()` produces a single `ChatModelStream` with the full v2 event lifecycle, text projection accumulates correctly, multi-node graphs produce one stream per model call, constructed-message nodes still replay via `message_to_events`, async mirror via `ainvoke` + `astream`. - `TestDirectMessagesModeStaysV1` — regression guard: direct `graph.stream(stream_mode="messages")` still yields AIMessageChunk tuples (not event dicts). - `TestStreamMessagesHandlerV2Unit` — direct unit test that the v2 handler's `on_llm_new_token` does not emit. --- .../langgraph/_internal/_constants.py | 4 + libs/langgraph/langgraph/pregel/_messages.py | 74 ++ libs/langgraph/langgraph/pregel/main.py | 20 +- .../langgraph/stream/streaming_handler.py | 23 +- .../tests/test_stream_messages_transformer.py | 862 ++++++++++++++++++ 5 files changed, 978 insertions(+), 5 deletions(-) create mode 100644 libs/langgraph/tests/test_stream_messages_transformer.py diff --git a/libs/langgraph/langgraph/_internal/_constants.py b/libs/langgraph/langgraph/_internal/_constants.py index 68cb48fe8..d28289053 100644 --- a/libs/langgraph/langgraph/_internal/_constants.py +++ b/libs/langgraph/langgraph/_internal/_constants.py @@ -66,6 +66,9 @@ CONFIG_KEY_RUNTIME = sys.intern("__pregel_runtime") # holds a `Runtime` instance with context, store, stream writer, etc. CONFIG_KEY_RESUME_MAP = sys.intern("__pregel_resume_map") # holds a mapping of task ns -> resume value for resuming tasks +CONFIG_KEY_STREAM_MESSAGES_V2 = sys.intern("__pregel_stream_messages_v2") +# when True, attach StreamMessagesHandlerV2 so content-block (v2) events +# flow through stream_mode="messages"; set by StreamingHandler only. # --- Other constants --- PUSH = sys.intern("__pregel_push") @@ -107,6 +110,7 @@ RESERVED = { CONFIG_KEY_CHECKPOINT_ID, CONFIG_KEY_CHECKPOINT_NS, CONFIG_KEY_RESUME_MAP, + CONFIG_KEY_STREAM_MESSAGES_V2, # other constants PUSH, PULL, diff --git a/libs/langgraph/langgraph/pregel/_messages.py b/libs/langgraph/langgraph/pregel/_messages.py index 6cda6b0a1..56eb829e3 100644 --- a/libs/langgraph/langgraph/pregel/_messages.py +++ b/libs/langgraph/langgraph/pregel/_messages.py @@ -24,6 +24,11 @@ try: except ImportError: _StreamingCallbackHandler = object # type: ignore +try: + from langchain_core.tracers._streaming import _V2StreamingCallbackHandler +except ImportError: + _V2StreamingCallbackHandler = object # type: ignore + T = TypeVar("T") Meta = tuple[tuple[str, ...], dict[str, Any]] @@ -150,6 +155,37 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): stream_metadata["tags"] = filtered_tags self.metadata[run_id] = (ns, stream_metadata) + def on_stream_event( + self, + event: dict[str, Any], + *, + run_id: UUID, + parent_run_id: UUID | None = None, + tags: list[str] | None = None, + **kwargs: Any, + ) -> Any: + """Forward a protocol event from `stream_v2` as a messages stream part. + + Fires once per `MessagesData` event (`message-start`, per-block + `content-block-*`, `message-finish`). The transformer layer + correlates events back to a single `ChatModelStream` via + `metadata["run_id"]` — attached here so the v1 + `stream_mode="messages"` output (which emits + `(AIMessageChunk, metadata)` via `on_llm_new_token`) keeps its + original metadata shape. + """ + if meta := self.metadata.get(run_id): + # Record message_id on message-start so on_chain_end's + # dedupe skips the finalized AIMessage the node returns + # (otherwise the messages projection double-counts: once + # from streaming, once from the chain output). + if event.get("event") == "message-start": + msg_id = event.get("message_id") + if msg_id: + self.seen.add(msg_id) + v2_meta = {**meta[1], "run_id": str(run_id)} + self.stream((meta[0], "messages", (event, v2_meta))) + def on_llm_new_token( self, token: str, @@ -256,3 +292,41 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): **kwargs: Any, ) -> Any: self.metadata.pop(run_id, None) + + +class StreamMessagesHandlerV2(StreamMessagesHandler, _V2StreamingCallbackHandler): + """v2 variant of `StreamMessagesHandler`. + + Declaring `_V2StreamingCallbackHandler` as a base flips + `BaseChatModel.invoke` to route through `_stream_chat_model_events` + (firing `on_stream_event`) instead of `_stream` (firing + `on_llm_new_token`). Inherits `on_stream_event` from the parent, + which forwards protocol events onto the messages stream channel. + + Pregel attaches this class instead of the v1 handler only when + `StreamingHandler` opts in via the internal + `CONFIG_KEY_STREAM_MESSAGES_V2` config key; direct + `graph.stream(stream_mode="messages")` callers keep the v1 + AIMessageChunk shape. + """ + + def on_llm_new_token( + self, + token: str, + *, + chunk: ChatGenerationChunk | None = None, + run_id: UUID, + parent_run_id: UUID | None = None, + tags: list[str] | None = None, + **kwargs: Any, + ) -> Any: + """Suppress v1 chunk emission on the messages channel. + + The v2 marker already steers `invoke` to the event generator, + so `on_llm_new_token` should not fire under normal routing. + This override guards against any caller (e.g. a node that + calls `model.stream()` directly, which still fires the v1 + callback) leaking AIMessageChunks onto a v2-flagged messages + stream. + """ + return None diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index c440e77f9..6c0471970 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -73,6 +73,7 @@ from langgraph._internal._constants import ( CONFIG_KEY_RUNTIME, CONFIG_KEY_SEND, CONFIG_KEY_STREAM, + CONFIG_KEY_STREAM_MESSAGES_V2, CONFIG_KEY_TASK_ID, CONFIG_KEY_THREAD_ID, ERROR, @@ -133,7 +134,10 @@ from langgraph.pregel._loop import ( AsyncPregelLoop, SyncPregelLoop, ) -from langgraph.pregel._messages import StreamMessagesHandler +from langgraph.pregel._messages import ( + StreamMessagesHandler, + StreamMessagesHandlerV2, +) from langgraph.pregel._read import DEFAULT_BOUND, PregelNode from langgraph.pregel._retry import RetryPolicy from langgraph.pregel._runner import PregelRunner @@ -2626,8 +2630,13 @@ class Pregel( # set up messages stream mode if "messages" in stream_modes: ns_ = cast(str | None, config[CONF].get(CONFIG_KEY_CHECKPOINT_NS)) + messages_handler_cls = ( + StreamMessagesHandlerV2 + if config[CONF].get(CONFIG_KEY_STREAM_MESSAGES_V2) + else StreamMessagesHandler + ) run_manager.inheritable_handlers.append( - StreamMessagesHandler( + messages_handler_cls( stream.put, subgraphs, parent_ns=tuple(ns_.split(NS_SEP)) if ns_ else None, @@ -3003,8 +3012,13 @@ class Pregel( if "messages" in stream_modes: # namespace can be None in a root level graph? ns_ = cast(str | None, config[CONF].get(CONFIG_KEY_CHECKPOINT_NS)) + messages_handler_cls = ( + StreamMessagesHandlerV2 + if config[CONF].get(CONFIG_KEY_STREAM_MESSAGES_V2) + else StreamMessagesHandler + ) run_manager.inheritable_handlers.append( - StreamMessagesHandler( + messages_handler_cls( stream_put, subgraphs, parent_ns=tuple(ns_.split(NS_SEP)) if ns_ else None, diff --git a/libs/langgraph/langgraph/stream/streaming_handler.py b/libs/langgraph/langgraph/stream/streaming_handler.py index 9748c48b5..e51e810af 100644 --- a/libs/langgraph/langgraph/stream/streaming_handler.py +++ b/libs/langgraph/langgraph/stream/streaming_handler.py @@ -6,6 +6,7 @@ from typing import Any from langchain_core.runnables import RunnableConfig +from langgraph._internal._constants import CONF, CONFIG_KEY_STREAM_MESSAGES_V2 from langgraph.pregel import Pregel from langgraph.stream._convert import convert_to_protocol_event from langgraph.stream._mux import StreamMux @@ -14,6 +15,24 @@ from langgraph.stream.run_stream import AsyncGraphRunStream, GraphRunStream from langgraph.stream.transformers import MessagesTransformer, ValuesTransformer from langgraph.types import All, StreamMode + +def _merge_v2_messages_flag( + config: RunnableConfig | None, +) -> RunnableConfig: + """Return a config with the v2 messages flag set in `configurable`. + + Signals to pregel that `stream_mode="messages"` should attach + `StreamMessagesHandlerV2` for this call so invoke-time model runs + route through the v2 event generator and their protocol events + reach the messages channel. + """ + merged: RunnableConfig = dict(config or {}) # type: ignore[assignment] + configurable = dict(merged.get(CONF) or {}) + configurable[CONFIG_KEY_STREAM_MESSAGES_V2] = True + merged[CONF] = configurable + return merged + + # All stream modes to request from the graph. STREAM_V2_MODES: list[StreamMode] = [ "values", @@ -100,7 +119,7 @@ class StreamingHandler: graph_iter = iter( self._graph.stream( input, - config, + _merge_v2_messages_flag(config), stream_mode=STREAM_V2_MODES, subgraphs=True, version="v2", @@ -152,7 +171,7 @@ class StreamingHandler: try: async for part in self._graph.astream( input, - config, + _merge_v2_messages_flag(config), stream_mode=STREAM_V2_MODES, subgraphs=True, version="v2", diff --git a/libs/langgraph/tests/test_stream_messages_transformer.py b/libs/langgraph/tests/test_stream_messages_transformer.py new file mode 100644 index 000000000..c8a3275ca --- /dev/null +++ b/libs/langgraph/tests/test_stream_messages_transformer.py @@ -0,0 +1,862 @@ +"""Tests for the MessagesTransformer content-block upgrade (B2). + +Verifies that `MessagesTransformer` routes protocol events (emitted by +`stream_v2` via `on_stream_event`) to `ChatModelStream` objects keyed by +run_id, and replays whole `AIMessage` payloads via `message_to_events`. +Legacy v1 `AIMessageChunk` tuples (from `on_llm_new_token`) are ignored. +""" + +from __future__ import annotations + +import time +from typing import Any + +import pytest +from langchain_core.language_models import GenericFakeChatModel +from langchain_core.language_models.chat_model_stream import ( + AsyncChatModelStream, + ChatModelStream, +) +from langchain_core.messages import AIMessage, AIMessageChunk + +from langgraph.constants import END, START +from langgraph.graph import MessagesState, StateGraph +from langgraph.stream._event_log import EventLog +from langgraph.stream._mux import StreamMux +from langgraph.stream.run_stream import GraphRunStream +from langgraph.stream.streaming_handler import StreamingHandler +from langgraph.stream.transformers import MessagesTransformer, ValuesTransformer + +TS = int(time.time() * 1000) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _proto_event( + event: dict[str, Any], + *, + run_id: str = "run-1", + node: str = "llm", +) -> dict[str, Any]: + """Build a messages ProtocolEvent carrying a protocol event dict (v2 path).""" + metadata: dict[str, Any] = {"langgraph_node": node, "run_id": run_id} + return { + "type": "event", + "method": "messages", + "params": { + "namespace": [], + "timestamp": TS, + "data": (event, metadata), + }, + } + + +def _v1_chunk( + text: str, + msg_id: str = "msg-1", + *, + finish: bool = False, + node: str = "llm", +) -> dict[str, Any]: + """Build a messages ProtocolEvent carrying a v1 AIMessageChunk tuple.""" + rm: dict[str, Any] = {} + if finish: + rm["finish_reason"] = "stop" + message = AIMessageChunk(content=text, id=msg_id, response_metadata=rm) + metadata: dict[str, Any] = {"langgraph_node": node} + return { + "type": "event", + "method": "messages", + "params": { + "namespace": [], + "timestamp": TS, + "data": (message, metadata), + }, + } + + +def _whole_msg( + text: str, + msg_id: str = "msg-10", + *, + node: str = "node", +) -> dict[str, Any]: + """Build a messages ProtocolEvent carrying a completed AIMessage.""" + message = AIMessage(content=text, id=msg_id) + metadata: dict[str, Any] = {"langgraph_node": node} + return { + "type": "event", + "method": "messages", + "params": { + "namespace": [], + "timestamp": TS, + "data": (message, metadata), + }, + } + + +def _make_sync_transformer() -> tuple[MessagesTransformer, EventLog[ChatModelStream]]: + t = MessagesTransformer() + proj = t.init() + log: EventLog[ChatModelStream] = proj["messages"] + log._bind(is_async=False) + t._bind_pump(lambda: False) + return t, log + + +def _make_async_transformer() -> tuple[MessagesTransformer, EventLog[ChatModelStream]]: + t = MessagesTransformer() + proj = t.init() + log: EventLog[ChatModelStream] = proj["messages"] + log._bind(is_async=True) + return t, log + + +# Standard lifecycle events for one streaming LLM call. +def _lifecycle( + *, + text: str = "hello world", + message_id: str = "run-1", +) -> list[dict[str, Any]]: + """Produce a valid protocol event lifecycle: start, delta, finish, end.""" + # Split text into two deltas to exercise delta accumulation. + half = len(text) // 2 + first, second = text[:half], text[half:] + return [ + {"event": "message-start", "role": "ai", "message_id": message_id}, + { + "event": "content-block-start", + "index": 0, + "content_block": {"type": "text", "text": ""}, + }, + { + "event": "content-block-delta", + "index": 0, + "content_block": {"type": "text", "text": first}, + }, + { + "event": "content-block-delta", + "index": 0, + "content_block": {"type": "text", "text": second}, + }, + { + "event": "content-block-finish", + "index": 0, + "content_block": {"type": "text", "text": text}, + }, + {"event": "message-finish", "reason": "stop"}, + ] + + +# --------------------------------------------------------------------------- +# Primary path: protocol event routing +# --------------------------------------------------------------------------- + + +class TestProtocolEventRouting: + def test_message_start_creates_stream(self) -> None: + t, log = _make_sync_transformer() + t.process( + _proto_event( + {"event": "message-start", "role": "ai", "message_id": "run-1"}, + run_id="run-1", + ) + ) + # Stream is in the log immediately. + log.close() + streams = list(log) + assert len(streams) == 1 + assert isinstance(streams[0], ChatModelStream) + assert streams[0].message_id == "run-1" + + def test_full_lifecycle_yields_done_stream(self) -> None: + t, log = _make_sync_transformer() + for evt in _lifecycle(text="hello world"): + t.process(_proto_event(evt, run_id="run-1")) + log.close() + (stream,) = list(log) + assert stream.done + assert stream.output.content == "hello world" + + def test_message_finish_cleans_up_routing(self) -> None: + t, log = _make_sync_transformer() + for evt in _lifecycle(): + t.process(_proto_event(evt, run_id="run-1")) + assert t._by_run == {} + + def test_events_without_prior_start_are_ignored(self) -> None: + """Orphan delta events (no preceding message-start) are dropped silently.""" + t, log = _make_sync_transformer() + t.process( + _proto_event( + { + "event": "content-block-delta", + "index": 0, + "content_block": {"type": "text", "text": "orphan"}, + }, + run_id="unknown", + ) + ) + log.close() + assert list(log) == [] + + def test_concurrent_streams_routed_by_run_id(self) -> None: + """Two interleaved LLM calls each produce their own stream.""" + t, log = _make_sync_transformer() + # Interleave events from two different run_ids. + life_a = _lifecycle(text="aaaa", message_id="run-a") + life_b = _lifecycle(text="bbbb", message_id="run-b") + for a, b in zip(life_a, life_b): + t.process(_proto_event(a, run_id="run-a")) + t.process(_proto_event(b, run_id="run-b")) + log.close() + streams = list(log) + assert len(streams) == 2 + by_id = {s.message_id: s for s in streams} + assert by_id["run-a"].output.content == "aaaa" + assert by_id["run-b"].output.content == "bbbb" + + def test_text_deltas_accumulated_on_stream(self) -> None: + t, log = _make_sync_transformer() + for evt in _lifecycle(text="abcdef"): + t.process(_proto_event(evt)) + log.close() + (stream,) = list(log) + deltas = list(stream._text_proj._deltas) + assert "".join(deltas) == "abcdef" + + def test_stream_pushed_on_message_start_not_finish(self) -> None: + """Consumer can see the stream before it finishes.""" + t, log = _make_sync_transformer() + t.process( + _proto_event( + {"event": "message-start", "role": "ai", "message_id": "run-1"}, + run_id="run-1", + ) + ) + # The log has the stream immediately — even though message-finish + # hasn't arrived yet. + assert len(log._items) == 1 + + def test_node_metadata_set_on_stream(self) -> None: + t, log = _make_sync_transformer() + t.process( + _proto_event( + {"event": "message-start", "role": "ai", "message_id": "run-1"}, + run_id="run-1", + node="my_llm", + ) + ) + (stream,) = [*log._items] + assert stream.node == "my_llm" + + +# --------------------------------------------------------------------------- +# Non-streaming (whole AIMessage) fallback +# --------------------------------------------------------------------------- + + +class TestWholeMessageFallback: + def test_whole_ai_message_produces_complete_stream(self) -> None: + t, log = _make_sync_transformer() + t.process(_whole_msg("the full answer")) + log.close() + (stream,) = list(log) + assert stream.done + assert stream.output.content == "the full answer" + + def test_whole_message_has_full_lifecycle(self) -> None: + t, log = _make_sync_transformer() + t.process(_whole_msg("full")) + log.close() + (stream,) = list(log) + event_types = [e["event"] for e in stream._events] + assert event_types == [ + "message-start", + "content-block-start", + "content-block-delta", + "content-block-finish", + "message-finish", + ] + + +# --------------------------------------------------------------------------- +# Legacy v1 chunks are ignored (users must migrate to stream_v2) +# --------------------------------------------------------------------------- + + +class TestLegacyChunksIgnored: + def test_aimessage_chunk_tuple_is_dropped(self) -> None: + t, log = _make_sync_transformer() + t.process(_v1_chunk("hello")) + t.process(_v1_chunk(" world", finish=True)) + log.close() + assert list(log) == [] + + +# --------------------------------------------------------------------------- +# Filtering behaviors +# --------------------------------------------------------------------------- + + +class TestFiltering: + def test_non_messages_events_pass_through(self) -> None: + t, _ = _make_sync_transformer() + values_event = { + "type": "event", + "method": "values", + "params": {"namespace": [], "timestamp": TS, "data": {"x": 1}}, + } + assert t.process(values_event) is True + + def test_subgraph_namespace_dropped(self) -> None: + t, log = _make_sync_transformer() + t.process( + { + "type": "event", + "method": "messages", + "params": { + "namespace": ["subgraph"], + "timestamp": TS, + "data": ( + {"event": "message-start", "message_id": "run-x"}, + {"run_id": "run-x"}, + ), + }, + } + ) + log.close() + assert list(log) == [] + + +# --------------------------------------------------------------------------- +# Lifecycle: finalize / fail +# --------------------------------------------------------------------------- + + +class TestLifecycle: + def test_fail_propagates_to_open_streams(self) -> None: + t, log = _make_sync_transformer() + t.process( + _proto_event( + {"event": "message-start", "message_id": "run-1"}, + run_id="run-1", + ) + ) + streams = list(log._items) + err = RuntimeError("graph died") + t.fail(err) + assert t._by_run == {} + assert streams[0]._error is err + + def test_finalize_clears_routing_state(self) -> None: + t, _ = _make_sync_transformer() + t.process( + _proto_event( + {"event": "message-start", "message_id": "run-1"}, + run_id="run-1", + ) + ) + assert "run-1" in t._by_run + t.finalize() + assert t._by_run == {} + + +# --------------------------------------------------------------------------- +# Async mode (AsyncChatModelStream) +# --------------------------------------------------------------------------- + + +class TestAsyncMode: + def test_async_mode_creates_async_stream(self) -> None: + t, log = _make_async_transformer() + for evt in _lifecycle(text="async stream"): + t.process(_proto_event(evt)) + streams = list(log._items) + assert len(streams) == 1 + assert isinstance(streams[0], AsyncChatModelStream) + + @pytest.mark.anyio + async def test_async_text_projection_yields_deltas(self) -> None: + t, log = _make_async_transformer() + for evt in _lifecycle(text="hello world"): + t.process(_proto_event(evt)) + (stream,) = list(log._items) + assert isinstance(stream, AsyncChatModelStream) + collected = [] + async for delta in stream.text: + collected.append(delta) + assert "".join(collected) == "hello world" + + @pytest.mark.anyio + async def test_async_output_awaitable(self) -> None: + t, log = _make_async_transformer() + for evt in _lifecycle(text="async"): + t.process(_proto_event(evt)) + (stream,) = list(log._items) + msg = await stream.output + assert msg.content == "async" + + +# --------------------------------------------------------------------------- +# GraphRunStream integration +# --------------------------------------------------------------------------- + + +class TestWireRequestMore: + def test_bind_pump_called_on_wire(self) -> None: + values_t = ValuesTransformer() + messages_t = MessagesTransformer() + mux = StreamMux([values_t, messages_t], is_async=False) + + assert messages_t._pump_fn is None + run = GraphRunStream(iter([]), mux, values_t) + # After wire, the transformer's pump callback is set. + assert messages_t._pump_fn is not None + # And calling it invokes GraphRunStream._pump_next (drains an empty + # graph_iter, returns False). + assert messages_t._pump_fn() is False + assert run._exhausted + + def test_created_streams_have_request_more(self) -> None: + values_t = ValuesTransformer() + messages_t = MessagesTransformer() + mux = StreamMux([values_t, messages_t], is_async=False) + + GraphRunStream(iter([]), mux, values_t) + + for evt in _lifecycle(): + messages_t.process(_proto_event(evt)) + + log: EventLog[ChatModelStream] = mux.extensions["messages"] + (stream,) = list(log._items) + # Pump was threaded through: the stream's _request_more points at + # the same callable the transformer was bound with. + assert stream._request_more is messages_t._pump_fn + + +# --------------------------------------------------------------------------- +# End-to-end via StreamMux +# --------------------------------------------------------------------------- + + +class TestViaMux: + def test_streaming_via_mux(self) -> None: + t = MessagesTransformer() + v = ValuesTransformer() + mux = StreamMux([v, t], is_async=False) + t._bind_pump(lambda: False) + + for evt in _lifecycle(text="mux stream"): + mux.push(_proto_event(evt)) + mux.close() + + log: EventLog[ChatModelStream] = mux.extensions["messages"] + (stream,) = list(log) + assert stream.output.content == "mux stream" + + def test_whole_message_via_mux(self) -> None: + t = MessagesTransformer() + v = ValuesTransformer() + mux = StreamMux([v, t], is_async=False) + t._bind_pump(lambda: False) + + mux.push(_whole_msg("result")) + mux.close() + + log: EventLog[ChatModelStream] = mux.extensions["messages"] + (stream,) = list(log) + assert stream.output.content == "result" + + @pytest.mark.anyio + async def test_async_streaming_via_mux(self) -> None: + t = MessagesTransformer() + v = ValuesTransformer() + mux = StreamMux([v, t], is_async=True) + + for evt in _lifecycle(text="async mux"): + await mux.apush(_proto_event(evt)) + + log: EventLog[ChatModelStream] = mux.extensions["messages"] + streams = list(log._items) + assert len(streams) == 1 + msg = await streams[0].output + assert msg.content == "async mux" + await mux.aclose() + + +# --------------------------------------------------------------------------- +# End-to-end: full graph → StreamingHandler → run.messages +# --------------------------------------------------------------------------- + + +class TestEndToEnd: + """Prove the full pipeline works when a node calls `model.stream_v2()`. + + These tests exercise the path that the new messages projection is + designed for: a user node invokes `stream_v2` on a chat model, + `on_stream_event` fires on `StreamMessagesHandler`, the handler + forwards to the mux, and the transformer routes events into a + `ChatModelStream` exposed on `run.messages`. + + Nothing in Pregel calls `stream_v2` automatically yet; the planned + `graph.stream_v2()` API (B4) and the `create_react_agent` + integration (C2) will wire that up. Until then, populating the + messages projection is opt-in at the node level. + """ + + def test_node_calling_stream_v2_populates_messages(self) -> None: + model = GenericFakeChatModel(messages=iter(["hello world"])) + + def call_model(state: MessagesState) -> dict[str, Any]: + stream = model.stream_v2(state["messages"]) + return {"messages": stream.output} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + handler = StreamingHandler(graph) + run = handler.stream({"messages": "hi"}) + streams = list(run.messages) + + assert len(streams) == 1 + assert isinstance(streams[0], ChatModelStream) + assert streams[0].output.content == "hello world" + + def test_node_stream_v2_text_deltas_iterate(self) -> None: + """Consumer can iterate `.text` on the streamed message in real time.""" + model = GenericFakeChatModel(messages=iter(["streamed answer"])) + + def call_model(state: MessagesState) -> dict[str, Any]: + stream = model.stream_v2(state["messages"]) + return {"messages": stream.output} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + handler = StreamingHandler(graph) + run = handler.stream({"messages": "go"}) + + # Pull the stream handle out, then iterate its text deltas. + (stream,) = list(run.messages) + text = "".join(stream.text) + assert text == "streamed answer" + + def test_non_llm_message_returned_from_node(self) -> None: + """Node returns a finalized AIMessage directly — whole-message fallback.""" + + def return_message(state: MessagesState) -> dict[str, Any]: + return {"messages": AIMessage(content="hardcoded", id="msg-abc")} + + graph = ( + StateGraph(MessagesState) + .add_node("return_message", return_message) + .add_edge(START, "return_message") + .add_edge("return_message", END) + .compile() + ) + + handler = StreamingHandler(graph) + run = handler.stream({"messages": "hi"}) + streams = list(run.messages) + + assert len(streams) == 1 + assert streams[0].output.content == "hardcoded" + + @pytest.mark.anyio + async def test_async_node_calling_astream_v2(self) -> None: + model = GenericFakeChatModel(messages=iter(["async answer"])) + + async def call_model(state: MessagesState) -> dict[str, Any]: + stream = await model.astream_v2(state["messages"]) + msg = await stream + return {"messages": msg} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + handler = StreamingHandler(graph) + run = await handler.astream({"messages": "hi"}) + + streams = [] + async for stream in run.messages: + streams.append(stream) + + assert len(streams) == 1 + assert isinstance(streams[0], AsyncChatModelStream) + msg = await streams[0].output + assert msg.content == "async answer" + + +class TestEndToEndV2Invoke: + """Nodes call `model.invoke()`; `StreamingHandler` routes through v2. + + Exercises the auto-routing path added in + `feat(core): route invoke through v2 event path for + _V2StreamingCallbackHandler`: `StreamingHandler` injects + `CONFIG_KEY_STREAM_MESSAGES_V2` into the config, pregel attaches + `StreamMessagesHandlerV2`, `BaseChatModel._should_stream_v2` sees the + v2 marker and drives the protocol event generator, and + `on_stream_event` forwards each event onto the messages channel. + """ + + def test_invoke_with_v2_marker_populates_messages(self) -> None: + """Node calling `model.invoke()` produces one ChatModelStream with v2 events.""" + model = GenericFakeChatModel(messages=iter(["hello world"])) + + def call_model(state: MessagesState) -> dict[str, Any]: + return {"messages": model.invoke(state["messages"])} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + handler = StreamingHandler(graph) + run = handler.stream({"messages": "hi"}) + streams = list(run.messages) + + assert len(streams) == 1, ( + "Expected exactly one ChatModelStream — the streamed invoke and " + "the node's return of the same AIMessage must dedupe." + ) + stream = streams[0] + assert isinstance(stream, ChatModelStream) + assert stream.output.content == "hello world" + + def test_invoke_v2_emits_protocol_events(self) -> None: + """Iterating the stream yields the full v2 lifecycle (not v1 chunks).""" + model = GenericFakeChatModel(messages=iter(["streamed answer"])) + + def call_model(state: MessagesState) -> dict[str, Any]: + return {"messages": model.invoke(state["messages"])} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + handler = StreamingHandler(graph) + run = handler.stream({"messages": "go"}) + (stream,) = list(run.messages) + + events = list(stream) + event_types = [e.get("event") for e in events] + assert "message-start" in event_types + assert "content-block-start" in event_types + assert "content-block-delta" in event_types + assert "content-block-finish" in event_types + assert "message-finish" in event_types + # Sanity: every event is a dict carrying an "event" key — not an + # AIMessageChunk tuple from the v1 path. + for event in events: + assert isinstance(event, dict) + assert "event" in event + # Typed projection still assembles the final text. + assert stream.output.content == "streamed answer" + + def test_invoke_text_deltas_iterate_live(self) -> None: + """`.text` projection yields deltas in order.""" + model = GenericFakeChatModel(messages=iter(["delta streaming works"])) + + def call_model(state: MessagesState) -> dict[str, Any]: + return {"messages": model.invoke(state["messages"])} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + handler = StreamingHandler(graph) + run = handler.stream({"messages": "hi"}) + (stream,) = list(run.messages) + + assembled = "".join(stream.text) + assert assembled == "delta streaming works" + + def test_invoke_dedupe_survives_multi_node_graph(self) -> None: + """Two model-invoking nodes produce exactly two streams, each once.""" + model_a = GenericFakeChatModel(messages=iter(["alpha"])) + model_b = GenericFakeChatModel(messages=iter(["beta"])) + + def node_a(state: MessagesState) -> dict[str, Any]: + return {"messages": model_a.invoke(state["messages"])} + + def node_b(state: MessagesState) -> dict[str, Any]: + return {"messages": model_b.invoke(state["messages"])} + + graph = ( + StateGraph(MessagesState) + .add_node("node_a", node_a) + .add_node("node_b", node_b) + .add_edge(START, "node_a") + .add_edge("node_a", "node_b") + .add_edge("node_b", END) + .compile() + ) + + handler = StreamingHandler(graph) + run = handler.stream({"messages": "hi"}) + streams = list(run.messages) + + assert len(streams) == 2 + contents = {s.output.content for s in streams} + assert contents == {"alpha", "beta"} + + def test_invoke_plus_constructed_message_two_streams(self) -> None: + """A v2-streamed node + a node that returns a constructed AIMessage + produces two ChatModelStreams — one from the live event lifecycle, + one synthesized from the constructed message via `message_to_events`. + """ + model = GenericFakeChatModel(messages=iter(["live stream"])) + + def streaming_node(state: MessagesState) -> dict[str, Any]: + return {"messages": model.invoke(state["messages"])} + + def constructed_node(state: MessagesState) -> dict[str, Any]: + return {"messages": [AIMessage(content="hardcoded", id="constructed-1")]} + + graph = ( + StateGraph(MessagesState) + .add_node("streaming_node", streaming_node) + .add_node("constructed_node", constructed_node) + .add_edge(START, "streaming_node") + .add_edge("streaming_node", "constructed_node") + .add_edge("constructed_node", END) + .compile() + ) + + handler = StreamingHandler(graph) + run = handler.stream({"messages": "hi"}) + streams = list(run.messages) + + assert len(streams) == 2 + assert streams[0].node == "streaming_node" + assert streams[0].output.content == "live stream" + assert streams[1].node == "constructed_node" + assert streams[1].output.content == "hardcoded" + assert streams[1].message_id == "constructed-1" + + @pytest.mark.anyio + async def test_ainvoke_with_v2_marker_populates_messages(self) -> None: + """Async mirror: `model.ainvoke()` + `StreamingHandler.astream()`.""" + model = GenericFakeChatModel(messages=iter(["async invoke"])) + + async def call_model(state: MessagesState) -> dict[str, Any]: + return {"messages": await model.ainvoke(state["messages"])} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + handler = StreamingHandler(graph) + run = await handler.astream({"messages": "hi"}) + + streams = [] + async for stream in run.messages: + streams.append(stream) + + assert len(streams) == 1 + assert isinstance(streams[0], AsyncChatModelStream) + msg = await streams[0].output + assert msg.content == "async invoke" + + +class TestDirectMessagesModeStaysV1: + """Regression guard: direct `graph.stream(stream_mode="messages")` + (no `StreamingHandler`) must keep the v1 `(AIMessageChunk, metadata)` + tuple shape. The v2 flag is only injected by `StreamingHandler`. + """ + + def test_direct_graph_stream_messages_yields_ai_message_chunks(self) -> None: + model = GenericFakeChatModel(messages=iter(["legacy path"])) + + def call_model(state: MessagesState) -> dict[str, Any]: + return {"messages": model.invoke(state["messages"])} + + graph = ( + StateGraph(MessagesState) + .add_node("call_model", call_model) + .add_edge(START, "call_model") + .add_edge("call_model", END) + .compile() + ) + + parts = list(graph.stream({"messages": "hi"}, stream_mode="messages")) + # Should have at least one streamed chunk; each part is + # (AIMessageChunk, metadata) — not a v2 event dict. + assert parts, "expected stream_mode='messages' to emit tuples" + for part in parts: + payload, _metadata = part + assert isinstance(payload, AIMessageChunk), ( + "direct graph.stream(stream_mode='messages') leaked v2 " + "event dicts — StreamingHandler flag bled through." + ) + assembled = "".join( + p[0].content for p in parts if isinstance(p[0].content, str) + ) + assert assembled == "legacy path" + + +class TestStreamMessagesHandlerV2Unit: + """Unit tests on the handler class itself.""" + + def test_on_llm_new_token_is_noop(self) -> None: + """v2 handler must not emit v1 chunks even if `on_llm_new_token` fires + (e.g. from a node calling `model.stream()` directly on a v2-flagged run). + """ + from uuid import uuid4 + + from langchain_core.outputs import ChatGenerationChunk + + from langgraph.pregel._messages import StreamMessagesHandlerV2 + + emitted: list[Any] = [] + handler = StreamMessagesHandlerV2(emitted.append, subgraphs=False) + run_id = uuid4() + # Register a fake run so `self.metadata.get(run_id)` would succeed for + # other callbacks — this makes sure the no-op is unconditional, not a + # side effect of missing metadata. + handler.metadata[run_id] = ((), {"langgraph_node": "x"}) + + handler.on_llm_new_token( + "hello", + chunk=ChatGenerationChunk(message=AIMessageChunk(content="hello")), + run_id=run_id, + ) + + assert emitted == [], ( + "StreamMessagesHandlerV2.on_llm_new_token must not push to the " + "messages stream — it's the v2 marker's guarantee." + )