feat(langgraph): route invoke messages through v2 via StreamingHandler

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.
This commit is contained in:
Nick Hollon
2026-04-17 16:13:11 -04:00
parent 5f24a0356a
commit acaa767542
5 changed files with 978 additions and 5 deletions
@@ -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,
@@ -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
+17 -3
View File
@@ -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,
@@ -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",
@@ -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."
)