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

This commit is contained in:
Nick Hollon
2026-04-24 16:59:07 -04:00
committed by GitHub
parent 40055e92cc
commit 0eac6626b3
11 changed files with 1697 additions and 698 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]]
@@ -256,3 +261,126 @@ 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:
"""Intentional no-op — v1 chunks are not used on v2-flagged runs.
The v2 marker already steers `invoke` to the event generator, so
`on_llm_new_token` should not fire under normal routing. This
override stays a pass-through (no call to `super()`) to make
the intent explicit and to guard 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.
"""
# Intentionally empty: v2 handler does not forward v1 chunks.
def __init__(
self,
stream: Callable[[StreamChunk], None],
subgraphs: bool,
*,
parent_ns: tuple[str, ...] | None = None,
) -> None:
super().__init__(stream, subgraphs, parent_ns=parent_ns)
self._streamed_run_ids: set[UUID] = set()
def on_llm_end(
self,
response: LLMResult,
*,
run_id: UUID,
parent_run_id: UUID | None = None,
**kwargs: Any,
) -> Any:
if meta := self.metadata.get(run_id):
if response.generations and response.generations[0]:
gen = response.generations[0][0]
if isinstance(gen, ChatGeneration):
if run_id in self._streamed_run_ids:
if gen.message.id is None:
gen.message.id = str(uuid4())
self.seen.add(gen.message.id)
else:
self._emit(meta, gen.message, dedupe=True)
self._streamed_run_ids.discard(run_id)
self.metadata.pop(run_id, None)
def on_llm_error(
self,
error: BaseException,
*,
run_id: UUID,
parent_run_id: UUID | None = None,
**kwargs: Any,
) -> Any:
self._streamed_run_ids.discard(run_id)
super().on_llm_error(
error,
run_id=run_id,
parent_run_id=parent_run_id,
**kwargs,
)
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.
Lives on the v2 handler rather than the v1 base: content-block
events are a v2-only concept, and forwarding them only when the
v2 handler is attached keeps the message channel's shape
predictable for v1 callers.
"""
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":
self._streamed_run_ids.add(run_id)
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)))
+142 -69
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
@@ -340,6 +344,17 @@ class NodeBuilder:
)
_STREAM_V2_MODES: list[StreamMode] = [
"values",
"updates",
"messages",
"custom",
"checkpoints",
"tasks",
"debug",
]
class Pregel(
PregelProtocol[StateT, ContextT, InputT, OutputT],
Generic[StateT, ContextT, InputT, OutputT],
@@ -2586,19 +2601,7 @@ class Pregel(
stream = SyncQueue()
config = ensure_config(self.config, config)
callback_manager = get_callback_manager_for_config(config)
if "ls_integration" not in callback_manager.metadata:
callback_manager.add_metadata({"ls_integration": "langgraph"})
run_manager = callback_manager.on_chain_start(
None,
input,
name=config.get("run_name", self.get_name()),
run_id=config.get("run_id"),
)
graph_callback_manager = get_sync_graph_callback_manager_for_config(
config,
run_id=run_manager.run_id,
)
run_manager = None
try:
# assign defaults
(
@@ -2619,6 +2622,36 @@ class Pregel(
interrupt_after=interrupt_after,
durability=durability,
)
callback_manager = get_callback_manager_for_config(config)
if "messages" in stream_modes and version != "v2":
# Strip any inherited v2 messages handler so a v1 stream
# does not get routed through the content-block event
# protocol. Leave v1 handlers in place — an outer
# stream(stream_mode="messages", subgraphs=True) relies
# on its inheritable handler to observe events emitted
# by inner stream(stream_mode="messages") calls.
callback_manager.handlers = [
h
for h in callback_manager.handlers
if not isinstance(h, StreamMessagesHandlerV2)
]
callback_manager.inheritable_handlers = [
h
for h in callback_manager.inheritable_handlers
if not isinstance(h, StreamMessagesHandlerV2)
]
if "ls_integration" not in callback_manager.metadata:
callback_manager.add_metadata({"ls_integration": "langgraph"})
run_manager = callback_manager.on_chain_start(
None,
input,
name=config.get("run_name", self.get_name()),
run_id=config.get("run_id"),
)
graph_callback_manager = get_sync_graph_callback_manager_for_config(
config,
run_id=run_manager.run_id,
)
if checkpointer is None and durability is not None:
warnings.warn(
"`durability` has no effect when no checkpointer is present.",
@@ -2630,8 +2663,16 @@ class Pregel(
# set up messages stream mode
if "messages" in stream_modes:
ns_ = cast(str | None, config[CONF].get(CONFIG_KEY_CHECKPOINT_NS))
use_stream_messages_v2 = bool(
version == "v2" and config[CONF].get(CONFIG_KEY_STREAM_MESSAGES_V2)
)
messages_handler_cls = (
StreamMessagesHandlerV2
if use_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,
@@ -2808,7 +2849,8 @@ class Pregel(
# set final channel values as run output
run_manager.on_chain_end(loop.output)
except BaseException as e:
run_manager.on_chain_error(e)
if run_manager is not None:
run_manager.on_chain_error(e)
raise
@overload
@@ -2948,33 +2990,7 @@ class Pregel(
)
config = ensure_config(self.config, config)
callback_manager = get_async_callback_manager_for_config(config)
if "ls_integration" not in callback_manager.metadata:
callback_manager.add_metadata({"ls_integration": "langgraph"})
run_manager = await callback_manager.on_chain_start(
None,
input,
name=config.get("run_name", self.get_name()),
run_id=config.get("run_id"),
)
graph_callback_manager = get_async_graph_callback_manager_for_config(
config,
run_id=run_manager.run_id,
)
# if running from astream_log() run each proc with streaming
do_stream = (
next(
(
True
for h in run_manager.handlers
if isinstance(h, _StreamingCallbackHandler)
and not isinstance(h, StreamMessagesHandler)
),
False,
)
if _StreamingCallbackHandler is not None
else False
)
run_manager = None
try:
# assign defaults
(
@@ -2995,6 +3011,50 @@ class Pregel(
interrupt_after=interrupt_after,
durability=durability,
)
callback_manager = get_async_callback_manager_for_config(config)
if "messages" in stream_modes and version != "v2":
# Strip any inherited v2 messages handler so a v1 stream
# does not get routed through the content-block event
# protocol. Leave v1 handlers in place — an outer
# astream(stream_mode="messages", subgraphs=True) relies
# on its inheritable handler to observe events emitted
# by inner astream(stream_mode="messages") calls.
callback_manager.handlers = [
h
for h in callback_manager.handlers
if not isinstance(h, StreamMessagesHandlerV2)
]
callback_manager.inheritable_handlers = [
h
for h in callback_manager.inheritable_handlers
if not isinstance(h, StreamMessagesHandlerV2)
]
if "ls_integration" not in callback_manager.metadata:
callback_manager.add_metadata({"ls_integration": "langgraph"})
run_manager = await callback_manager.on_chain_start(
None,
input,
name=config.get("run_name", self.get_name()),
run_id=config.get("run_id"),
)
graph_callback_manager = get_async_graph_callback_manager_for_config(
config,
run_id=run_manager.run_id,
)
# if running from astream_log() run each proc with streaming
do_stream = (
next(
(
True
for h in run_manager.handlers
if isinstance(h, _StreamingCallbackHandler)
and not isinstance(h, StreamMessagesHandler)
),
False,
)
if _StreamingCallbackHandler is not None
else False
)
if checkpointer is None and durability is not None:
warnings.warn(
"`durability` has no effect when no checkpointer is present.",
@@ -3007,8 +3067,16 @@ 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))
use_stream_messages_v2 = bool(
version == "v2" and config[CONF].get(CONFIG_KEY_STREAM_MESSAGES_V2)
)
messages_handler_cls = (
StreamMessagesHandlerV2
if use_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,
@@ -3238,7 +3306,8 @@ class Pregel(
# set final channel values as run output
await run_manager.on_chain_end(loop.output)
except BaseException as e:
await asyncio.shield(run_manager.on_chain_error(e))
if run_manager is not None:
await asyncio.shield(run_manager.on_chain_error(e))
raise
def stream_v2(
@@ -3277,12 +3346,13 @@ class Pregel(
ValuesTransformer,
)
values_t = ValuesTransformer()
parent_ns = _resolve_parent_ns(self.config, config)
values_t = ValuesTransformer(parent_ns=parent_ns)
compiled_instances = [f() for f in self._stream_transformers]
mux = StreamMux(
[
values_t,
MessagesTransformer(),
MessagesTransformer(parent_ns=parent_ns),
*compiled_instances,
*(transformers or ()),
],
@@ -3291,16 +3361,8 @@ class Pregel(
graph_iter = iter(
self.stream(
input,
config,
stream_mode=[
"values",
"updates",
"messages",
"custom",
"checkpoints",
"tasks",
"debug",
],
patch_configurable(config, {CONFIG_KEY_STREAM_MESSAGES_V2: True}),
stream_mode=_STREAM_V2_MODES,
subgraphs=True,
version="v2",
interrupt_before=interrupt_before,
@@ -3339,12 +3401,13 @@ class Pregel(
ValuesTransformer,
)
values_t = ValuesTransformer()
parent_ns = _resolve_parent_ns(self.config, config)
values_t = ValuesTransformer(parent_ns=parent_ns)
compiled_instances = [f() for f in self._stream_transformers]
mux = StreamMux(
[
values_t,
MessagesTransformer(),
MessagesTransformer(parent_ns=parent_ns),
*compiled_instances,
*(transformers or ()),
],
@@ -3352,16 +3415,8 @@ class Pregel(
)
graph_aiter = self.astream(
input,
config,
stream_mode=[
"values",
"updates",
"messages",
"custom",
"checkpoints",
"tasks",
"debug",
],
patch_configurable(config, {CONFIG_KEY_STREAM_MESSAGES_V2: True}),
stream_mode=_STREAM_V2_MODES,
subgraphs=True,
version="v2",
interrupt_before=interrupt_before,
@@ -3844,6 +3899,24 @@ def _coerce_checkpoint_values(payload: Any, mapper: Callable[[Any], Any]) -> Non
payload["values"] = mapper(payload["values"])
def _resolve_parent_ns(
graph_config: RunnableConfig | None, call_config: RunnableConfig | None
) -> tuple[str, ...]:
"""Return the checkpoint namespace the caller is running under.
`stream_v2` uses this to scope its native projections
(`ValuesTransformer`, `MessagesTransformer`) to events emitted at
the run's own level. A root call resolves to `()`; a call made
from inside a node carries the outer graph's task namespace so the
projection still matches its own root-level events.
"""
merged = ensure_config(graph_config, call_config)
ns = merged.get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_NS)
if not ns:
return ()
return tuple(ns.split(NS_SEP))
def _build_server_info(
config: RunnableConfig, parent_runtime: Runtime[Any]
) -> ServerInfo | None:
+25 -9
View File
@@ -64,14 +64,23 @@ class GraphRunStream:
self._wire_request_more(mux)
def _wire_request_more(self, mux: StreamMux) -> None:
"""Install `_request_more` on every sync EventLog so cursors
can drive the pump when their buffer catches up."""
"""Install `_request_more` on every sync EventLog so cursors can
drive the pump when their buffer catches up.
Also calls `_bind_pump` on any transformer that exposes it, so
transformers producing ChatModelStream objects (e.g.
MessagesTransformer) can wire the pull callback on each stream
as it's created.
"""
mux._events._request_more = self._pump_next
for value in mux.extensions.values():
if isinstance(value, EventLog):
value._request_more = self._pump_next
elif isinstance(value, StreamChannel):
value._log._request_more = self._pump_next
for transformer in mux._transformers:
if hasattr(transformer, "_bind_pump"):
transformer._bind_pump(self._pump_next)
def _pump_next(self) -> bool:
"""Pull one event from the graph and push it through the mux.
@@ -125,8 +134,7 @@ class GraphRunStream:
def output(self) -> dict[str, Any] | None:
"""Drive the run to completion and return the final state."""
_drive_until_done(self._pump_next)
err = self._values_transformer.error
if err is not None:
if (err := self._values_transformer.error) is not None:
raise err
return self._values_transformer._latest
@@ -139,8 +147,7 @@ class GraphRunStream:
BaseException: If the run ended with an error.
"""
_drive_until_done(self._pump_next)
err = self._values_transformer.error
if err is not None:
if (err := self._values_transformer.error) is not None:
raise err
return self._values_transformer._interrupted
@@ -152,8 +159,7 @@ class GraphRunStream:
BaseException: If the run ended with an error.
"""
_drive_until_done(self._pump_next)
err = self._values_transformer.error
if err is not None:
if (err := self._values_transformer.error) is not None:
raise err
return self._values_transformer._interrupts
@@ -262,13 +268,23 @@ class AsyncGraphRunStream:
def _wire_arequest_more(self, mux: StreamMux) -> None:
"""Install `_arequest_more` on every async EventLog so cursors
can drive the pump when their buffer catches up."""
can drive the pump when their buffer catches up.
Also calls `_bind_apump` on any transformer that exposes it,
so transformers producing `AsyncChatModelStream` objects (e.g.
`MessagesTransformer`) can fan the pull callback out to each
stream's projections. Mirrors the sync `_wire_request_more`
plumbing.
"""
mux._events._arequest_more = self._apump_next
for value in mux.extensions.values():
if isinstance(value, EventLog):
value._arequest_more = self._apump_next
elif isinstance(value, StreamChannel):
value._log._arequest_more = self._apump_next
for transformer in mux._transformers:
if hasattr(transformer, "_bind_apump"):
transformer._bind_apump(self._apump_next)
async def _apump_next(self) -> bool:
"""Drive one pump step, or wait for the active pumper to drive one.
+176 -19
View File
@@ -1,10 +1,21 @@
from __future__ import annotations
from typing import Any
from typing import TYPE_CHECKING, Any, cast
from langchain_core.language_models._compat_bridge import message_to_events
from langchain_core.language_models.chat_model_stream import (
AsyncChatModelStream,
ChatModelStream,
)
from langchain_core.messages import AIMessageChunk, BaseMessage
from langchain_protocol.protocol import MessagesData
from langgraph.stream._event_log import EventLog
from langgraph.stream._types import ProtocolEvent, StreamTransformer
if TYPE_CHECKING:
from collections.abc import Awaitable, Callable
class ValuesTransformer(StreamTransformer):
"""Capture values events as a drainable stream of state snapshots.
@@ -17,17 +28,22 @@ class ValuesTransformer(StreamTransformer):
Native transformer — projection keys are exposed as direct
attributes on the run stream (e.g. `run.values`).
Only root-namespace values events are captured; subgraph state
snapshots are ignored.
Only values events at the run's own level are captured; snapshots
from deeper subgraphs are left in the main event log but excluded
from the projection. "Own level" is defined by `parent_ns`, which
`stream_v2` / `astream_v2` populate from the caller's checkpoint
namespace so that a nested `stream_v2` call still sees its own
root snapshots.
"""
_native = True
def __init__(self) -> None:
def __init__(self, *, parent_ns: tuple[str, ...] = ()) -> None:
self._log: EventLog[dict[str, Any]] = EventLog()
self._latest: dict[str, Any] | None = None
self._interrupted = False
self._interrupts: list[Any] = []
self._parent_ns: list[str] = list(parent_ns)
def init(self) -> dict[str, Any]:
return {"values": self._log}
@@ -44,7 +60,7 @@ class ValuesTransformer(StreamTransformer):
if event["method"] != "values":
return True
params = event["params"]
if params["namespace"]:
if params["namespace"] != self._parent_ns:
return True
self._latest = params["data"]
interrupts = params.get("interrupts", ())
@@ -56,34 +72,175 @@ class ValuesTransformer(StreamTransformer):
class MessagesTransformer(StreamTransformer):
"""Pass through raw (chunk, metadata) tuples from messages events.
"""Capture messages events as ChatModelStream objects.
This is the same shape as today's `stream_mode="messages"` output.
A follow-on PR will replace this with a richer transformer that
produces ChatModelStream objects using the protocol handler.
The messages projection yields one `ChatModelStream` (or
`AsyncChatModelStream`) per LLM call. Consumers iterate
`run.messages` to get stream handles, then use each handle's typed
projections (`.text`, `.reasoning`, `.tool_calls`, `.usage`,
`.output`) for per-message content.
Only root-namespace messages events are captured; tokens emitted
from subgraphs are dropped from the `messages` projection.
Consumers that need subgraph tokens should iterate the raw event
stream or register a custom transformer.
Two input shapes are handled (via `params["data"] = (payload,
metadata)` from `StreamMessagesHandler`):
Native transformer — projection keys are exposed as direct
attributes on the run stream (e.g. `run.messages`).
1. Protocol event (dict with `"event"` key) — emitted by
`stream_v2()` / `astream_v2()` via the `on_stream_event`
callback. Routed to an existing `ChatModelStream` by
`metadata["run_id"]`. A `message-start` event creates a new
stream; `message-finish` closes it.
2. Whole `AIMessage` — emitted from `on_chain_end` when a node
returns a finalized message. Replayed as a synthetic protocol
event lifecycle via `message_to_events`, then the
already-complete stream is pushed to the log.
V1 `AIMessageChunk` tuples (from `on_llm_new_token`) are not
streamed into this projection: chat models that want to populate
`run.messages` with content-block streaming must use
`stream_v2()` / `astream_v2()`. Models called via the legacy
`stream()` method still surface their final `AIMessage` via
`on_chain_end` when a node returns it as state.
Only events at the run's own level are projected; tokens from
deeper subgraphs are left in the main event log but excluded from
`.messages`. "Own level" is defined by `parent_ns`, which
`stream_v2` / `astream_v2` populate from the caller's checkpoint
namespace so that a `stream_v2` call inside a node still sees its
own root chat model streams on `.messages`. Consumers that need
subgraph tokens should iterate the raw event stream or register a
custom transformer.
Native transformer — the `messages` projection is exposed as a
direct attribute on the run stream.
"""
_native = True
def __init__(self) -> None:
self._log: EventLog[tuple[Any, dict[str, Any]]] = EventLog()
def __init__(self, *, parent_ns: tuple[str, ...] = ()) -> None:
self._log: EventLog[ChatModelStream] = EventLog()
# 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._pump_fn: Callable[[], bool] | None = None
self._apump_fn: Callable[[], Awaitable[bool]] | None = None
# Root scope for this projection. Only chat model streams whose
# emitted namespace matches `parent_ns` are surfaced on
# `.messages`; events from deeper subgraphs stay in the main
# event log for other consumers but are not projected here.
self._parent_ns: list[str] = list(parent_ns)
def init(self) -> dict[str, Any]:
return {"messages": self._log}
def _bind_pump(self, fn: Callable[[], bool]) -> None:
"""Wire the sync pull callback. Called by GraphRunStream._wire_request_more."""
self._pump_fn = fn
def _bind_apump(self, fn: Callable[[], Awaitable[bool]]) -> None:
"""Wire the async pull callback.
Called by `AsyncGraphRunStream._wire_arequest_more` so each
`AsyncChatModelStream` this transformer creates can drive the
shared graph pump from its projection cursors.
"""
self._apump_fn = fn
def _make_stream(
self,
*,
namespace: list[str],
node: str | None,
message_id: str | None,
) -> ChatModelStream:
"""Create a ChatModelStream (sync) or AsyncChatModelStream (async).
Wires whichever pump is bound. Prefers the async pump so nested
iteration under `AsyncGraphRunStream` drives the graph forward
without a background task. The unwired fallback (no pump bound)
is used by unit tests that dispatch events manually.
"""
if self._apump_fn is not None:
astream = AsyncChatModelStream(
namespace=namespace,
node=node,
message_id=message_id,
)
astream.set_arequest_more(self._apump_fn)
return astream
if self._pump_fn is not None:
stream: ChatModelStream = ChatModelStream(
namespace=namespace,
node=node,
message_id=message_id,
)
stream.set_request_more(self._pump_fn)
return stream
return AsyncChatModelStream(
namespace=namespace,
node=node,
message_id=message_id,
)
def process(self, event: ProtocolEvent) -> bool:
if event["method"] != "messages":
return True
params = event["params"]
if params["namespace"]:
if params["namespace"] != self._parent_ns:
return True
self._log.push(params["data"])
payload, metadata = params["data"]
node: str | None = metadata.get("langgraph_node")
run_id = str(metadata.get("run_id", "")) if metadata else ""
if isinstance(payload, dict) and "event" in payload:
self._route_protocol_event(
cast("MessagesData", payload), run_id=run_id, node=node
)
elif isinstance(payload, BaseMessage) and not isinstance(
payload, AIMessageChunk
):
self._route_whole_message(payload, node=node)
# Legacy AIMessageChunk tuples (from on_llm_new_token) are ignored;
# v1 streaming callers must switch to stream_v2() to populate this
# projection.
return True
def _route_protocol_event(
self,
event: MessagesData,
*,
run_id: str,
node: str | None,
) -> None:
event_type = event.get("event")
if event_type == "message-start":
message_id = event.get("message_id")
stream = self._make_stream(
namespace=[],
node=node,
message_id=str(message_id) if message_id is not None else None,
)
self._by_run[run_id] = stream
self._log.push(stream)
stream.dispatch(event)
elif run_id in self._by_run:
stream = self._by_run[run_id]
stream.dispatch(event)
if event_type == "message-finish":
del self._by_run[run_id]
def _route_whole_message(self, message: BaseMessage, *, node: str | None) -> None:
stream = self._make_stream(namespace=[], node=node, message_id=message.id)
for evt in message_to_events(message, message_id=message.id):
stream.dispatch(evt)
self._log.push(stream)
def finalize(self) -> None:
"""Clear any routing state — streams close themselves via `message-finish`."""
self._by_run.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()
+1 -1
View File
@@ -24,7 +24,7 @@ classifiers = [
'Programming Language :: Python :: 3.13',
]
dependencies = [
"langchain-core==1.3.0a2",
"langchain-core==1.3.2",
"langgraph-checkpoint>=2.1.0,<5.0.0",
"langgraph-sdk>=0.3.0,<0.4.0",
"langgraph-prebuilt>=1.0.9,<1.1.0",
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,872 @@
"""Tests for MessagesTransformer: protocol event routing, whole-message fallback,
legacy v1 chunk filtering, and end-to-end via stream_v2 / astream_v2."""
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 langchain_core.runnables import RunnableConfig
from typing_extensions import TypedDict
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.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)."""
return {
"type": "event",
"method": "messages",
"params": {
"namespace": [],
"timestamp": TS,
"data": (event, {"langgraph_node": node, "run_id": run_id}),
},
}
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] = {"finish_reason": "stop"} if finish else {}
return {
"type": "event",
"method": "messages",
"params": {
"namespace": [],
"timestamp": TS,
"data": (
AIMessageChunk(content=text, id=msg_id, response_metadata=rm),
{"langgraph_node": node},
),
},
}
def _whole_msg(
text: str,
msg_id: str = "msg-10",
*,
node: str = "node",
) -> dict[str, Any]:
"""Build a messages ProtocolEvent carrying a completed AIMessage."""
return {
"type": "event",
"method": "messages",
"params": {
"namespace": [],
"timestamp": TS,
"data": (AIMessage(content=text, id=msg_id), {"langgraph_node": node}),
},
}
def _make_sync_transformer() -> tuple[MessagesTransformer, EventLog[ChatModelStream]]:
t = MessagesTransformer()
log: EventLog[ChatModelStream] = t.init()["messages"]
log._bind(is_async=False)
# Subscribe up front so pushes during process() are retained.
log._subscribed = True
t._bind_pump(lambda: False)
return t, log
def _make_async_transformer() -> tuple[MessagesTransformer, EventLog[ChatModelStream]]:
t = MessagesTransformer()
log: EventLog[ChatModelStream] = t.init()["messages"]
log._bind(is_async=True)
log._subscribed = True
return t, log
def _lifecycle(
*, text: str = "hello world", message_id: str = "run-1"
) -> list[dict[str, Any]]:
"""Produce a valid protocol event lifecycle: start, delta, finish."""
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"},
]
def _simple_graph():
def call_model(state: MessagesState) -> dict[str, Any]:
model = GenericFakeChatModel(messages=iter(["hello world"]))
stream = model.stream_v2(state["messages"])
return {"messages": stream.output}
return (
StateGraph(MessagesState)
.add_node("call_model", call_model)
.add_edge(START, "call_model")
.add_edge("call_model", END)
.compile()
)
# ---------------------------------------------------------------------------
# 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",
)
)
log.close()
(stream,) = list(log._items)
assert isinstance(stream, ChatModelStream)
assert stream.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._items)
assert stream.done
assert stream.output.text == "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:
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._items) == []
def test_concurrent_streams_routed_by_run_id(self) -> None:
t, log = _make_sync_transformer()
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._items)
assert len(streams) == 2
by_id = {s.message_id: s for s in streams}
assert by_id["run-a"].output.text == "aaaa"
assert by_id["run-b"].output.text == "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._items)
assert "".join(stream._text_proj._deltas) == "abcdef"
def test_stream_pushed_on_message_start_not_finish(self) -> None:
# Consumer can see the stream before message-finish arrives.
t, log = _make_sync_transformer()
t.process(
_proto_event(
{"event": "message-start", "role": "ai", "message_id": "run-1"},
run_id="run-1",
)
)
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,) = list(log._items)
assert stream.node == "my_llm"
# ---------------------------------------------------------------------------
# Whole-message 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._items)
assert stream.done
assert stream.output.text == "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._items)
assert [e["event"] for e in stream._events] == [
"message-start",
"content-block-start",
"content-block-delta",
"content-block-finish",
"message-finish",
]
# ---------------------------------------------------------------------------
# Filtering
# ---------------------------------------------------------------------------
class TestFiltering:
def test_non_messages_events_pass_through(self) -> None:
t, _ = _make_sync_transformer()
assert (
t.process(
{
"type": "event",
"method": "values",
"params": {"namespace": [], "timestamp": TS, "data": {"x": 1}},
}
)
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._items) == []
def test_legacy_v1_chunks_ignored(self) -> None:
# v1 AIMessageChunk tuples (from on_llm_new_token) are not streamed
# into this projection; callers must migrate to stream_v2.
t, log = _make_sync_transformer()
t.process(_v1_chunk("hello"))
t.process(_v1_chunk(" world", finish=True))
log.close()
assert list(log._items) == []
# ---------------------------------------------------------------------------
# Lifecycle: fail / finalize
# ---------------------------------------------------------------------------
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
# ---------------------------------------------------------------------------
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))
assert isinstance(list(log._items)[0], AsyncChatModelStream)
@pytest.mark.anyio
async def test_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)
assert "".join([d async for d in stream.text]) == "hello world"
@pytest.mark.anyio
async def test_output_awaitable(self) -> None:
t, log = _make_async_transformer()
for evt in _lifecycle(text="async"):
t.process(_proto_event(evt))
(stream,) = list(log._items)
assert (await stream.output).text == "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)
assert messages_t._pump_fn is not None
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)
log: EventLog[ChatModelStream] = mux.extensions["messages"]
log._subscribed = True
for evt in _lifecycle():
messages_t.process(_proto_event(evt))
(stream,) = list(log._items)
assert stream._request_more is messages_t._pump_fn
# ---------------------------------------------------------------------------
# End-to-end via StreamMux
# ---------------------------------------------------------------------------
class TestViaMux:
def _make_mux(
self,
) -> tuple[MessagesTransformer, StreamMux, EventLog[ChatModelStream]]:
t = MessagesTransformer()
v = ValuesTransformer()
mux = StreamMux([v, t], is_async=False)
t._bind_pump(lambda: False)
log: EventLog[ChatModelStream] = mux.extensions["messages"]
log._subscribed = True
return t, mux, log
def test_streaming_via_mux(self) -> None:
t, mux, log = self._make_mux()
for evt in _lifecycle(text="mux stream"):
mux.push(_proto_event(evt))
mux.close()
(stream,) = list(log._items)
assert stream.output.text == "mux stream"
def test_whole_message_via_mux(self) -> None:
t, mux, log = self._make_mux()
mux.push(_whole_msg("result"))
mux.close()
(stream,) = list(log._items)
assert stream.output.text == "result"
@pytest.mark.anyio
async def test_async_streaming_via_mux(self) -> None:
t = MessagesTransformer()
v = ValuesTransformer()
mux = StreamMux([v, t], is_async=True)
log: EventLog[ChatModelStream] = mux.extensions["messages"]
log._subscribed = True
for evt in _lifecycle(text="async mux"):
await mux.apush(_proto_event(evt))
(stream,) = list(log._items)
assert (await stream.output).text == "async mux"
await mux.aclose()
# ---------------------------------------------------------------------------
# End-to-end: graph → stream_v2 → run.messages (node calls stream_v2)
# ---------------------------------------------------------------------------
class TestEndToEnd:
"""stream_v2 path: node calls model.stream_v2() explicitly."""
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()
)
run = graph.stream_v2({"messages": "hi"})
(stream,) = list(run.messages)
assert isinstance(stream, ChatModelStream)
assert stream.output.text == "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()
)
run = graph.stream_v2({"messages": "go"})
(stream,) = list(run.messages)
assert "".join(stream.text) == "streamed answer"
def test_non_llm_message_returned_from_node(self) -> None:
"""Whole-message fallback: node returns a finalized AIMessage directly."""
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()
)
run = graph.stream_v2({"messages": "hi"})
(stream,) = list(run.messages)
assert stream.output.text == "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"])
return {"messages": await stream}
graph = (
StateGraph(MessagesState)
.add_node("call_model", call_model)
.add_edge(START, "call_model")
.add_edge("call_model", END)
.compile()
)
run = await graph.astream_v2({"messages": "hi"})
streams = [s async for s in run.messages]
assert len(streams) == 1
assert isinstance(streams[0], AsyncChatModelStream)
assert (await streams[0].output).text == "async answer"
@pytest.mark.anyio
async def test_nested_async_iteration_yields_text_deltas(self) -> None:
"""Inner stream.text drives the shared graph pump via the async pump binding."""
import asyncio
model = GenericFakeChatModel(messages=iter(["hello world"]))
async def call_model(state: MessagesState) -> dict[str, Any]:
stream = await model.astream_v2(state["messages"])
return {"messages": await stream}
graph = (
StateGraph(MessagesState)
.add_node("call_model", call_model)
.add_edge(START, "call_model")
.add_edge("call_model", END)
.compile()
)
run = await graph.astream_v2({"messages": "hi"})
async def consume() -> list[str]:
collected: list[str] = []
async for stream in run.messages:
async for delta in stream.text:
collected.append(delta)
return collected
assert "".join(await asyncio.wait_for(consume(), timeout=2.0)) == "hello world"
# ---------------------------------------------------------------------------
# End-to-end: graph → stream_v2 → run.messages (node calls invoke)
# ---------------------------------------------------------------------------
class TestEndToEndV2Invoke:
"""Auto-routing path: stream_v2 injects CONFIG_KEY_STREAM_MESSAGES_V2,
causing BaseChatModel to drive the v2 protocol event generator even for
model.invoke()."""
def _graph(self, model):
def call_model(state: MessagesState) -> dict[str, Any]:
return {"messages": model.invoke(state["messages"])}
return (
StateGraph(MessagesState)
.add_node("call_model", call_model)
.add_edge(START, "call_model")
.add_edge("call_model", END)
.compile()
)
def test_invoke_populates_messages(self) -> None:
run = self._graph(
GenericFakeChatModel(messages=iter(["hello world"]))
).stream_v2({"messages": "hi"})
(stream,) = list(run.messages)
assert isinstance(stream, ChatModelStream)
assert stream.output.text == "hello world"
def test_invoke_emits_protocol_events(self) -> None:
"""Iterating the stream yields the full v2 lifecycle, not v1 chunks."""
run = self._graph(
GenericFakeChatModel(messages=iter(["streamed answer"]))
).stream_v2({"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.text == "streamed answer"
def test_invoke_text_deltas_iterate(self) -> None:
run = self._graph(
GenericFakeChatModel(messages=iter(["delta streaming works"]))
).stream_v2({"messages": "hi"})
(stream,) = list(run.messages)
assert "".join(stream.text) == "delta streaming works"
def test_invoke_two_nodes_two_streams(self) -> None:
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()
)
streams = list(graph.stream_v2({"messages": "hi"}).messages)
assert len(streams) == 2
assert {s.output.text for s in streams} == {"alpha", "beta"}
def test_invoke_plus_constructed_message_two_streams(self) -> None:
"""Live-streamed node + constructed-message node → two ChatModelStreams."""
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()
)
run = graph.stream_v2({"messages": "hi"})
streams = list(run.messages)
assert len(streams) == 2
assert streams[0].node == "streaming_node"
assert streams[0].output.text == "live stream"
assert streams[1].node == "constructed_node"
assert streams[1].output.text == "hardcoded"
assert streams[1].message_id == "constructed-1"
@pytest.mark.anyio
async def test_ainvoke_populates_messages(self) -> None:
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()
)
run = await graph.astream_v2({"messages": "hi"})
streams = [s async for s in run.messages]
assert len(streams) == 1
assert isinstance(streams[0], AsyncChatModelStream)
assert (await streams[0].output).text == "async invoke"
# ---------------------------------------------------------------------------
# Regression: direct stream_mode="messages" must stay v1
# ---------------------------------------------------------------------------
class TestDirectMessagesModeStaysV1:
def test_direct_graph_stream_messages_yields_ai_message_chunks(self) -> None:
"""graph.stream(stream_mode="messages") must not leak v2 event dicts —
the v2 flag is only injected by stream_v2 / astream_v2."""
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"))
assert parts, "expected stream_mode='messages' to emit tuples"
for payload, _metadata in parts:
assert isinstance(payload, AIMessageChunk)
assert (
"".join(p[0].content for p in parts if isinstance(p[0].content, str))
== "legacy path"
)
def test_nested_graph_stream_messages_stays_v1_under_outer_stream_v2(self) -> None:
"""An outer `stream_v2()` run must not flip an inner direct
`stream_mode="messages"` call onto the v2 event protocol."""
model = GenericFakeChatModel(messages=iter(["nested legacy path"]))
def call_model(state: MessagesState) -> dict[str, Any]:
return {"messages": model.invoke(state["messages"])}
inner = (
StateGraph(MessagesState)
.add_node("call_model", call_model)
.add_edge(START, "call_model")
.add_edge("call_model", END)
.compile()
)
class OuterState(TypedDict, total=False):
saw_only_chunks: bool
first_payload_type: str
text: str
def call_subgraph(state: OuterState, config: RunnableConfig) -> dict[str, Any]:
parts = list(
inner.stream(
{"messages": "hi"},
config,
stream_mode="messages",
)
)
assert parts
payloads = [payload for payload, _metadata in parts]
return {
"saw_only_chunks": all(
isinstance(payload, AIMessageChunk) for payload in payloads
),
"first_payload_type": type(payloads[0]).__name__,
"text": "".join(
payload.content
for payload in payloads
if isinstance(payload, AIMessageChunk)
and isinstance(payload.content, str)
),
}
outer = (
StateGraph(OuterState)
.add_node("call_subgraph", call_subgraph)
.add_edge(START, "call_subgraph")
.add_edge("call_subgraph", END)
.compile()
)
result = outer.stream_v2({}).output
assert result is not None
assert result["saw_only_chunks"] is True
assert result["first_payload_type"] == "AIMessageChunk"
assert result["text"] == "nested legacy path"
# ---------------------------------------------------------------------------
# StreamMessagesHandlerV2 unit
# ---------------------------------------------------------------------------
class TestStreamMessagesHandlerV2Unit:
def test_on_llm_new_token_is_noop(self) -> None:
"""v2 handler must not emit v1 chunks even when on_llm_new_token fires."""
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()
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 == []
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`."""
from uuid import uuid4
from langchain_core.outputs import ChatGeneration, LLMResult
from langgraph.pregel._messages import StreamMessagesHandlerV2
emitted: list[Any] = []
handler = StreamMessagesHandlerV2(emitted.append, subgraphs=False)
run_id = uuid4()
handler.metadata[run_id] = ((), {"langgraph_node": "x"})
handler.on_stream_event(
{"event": "message-start", "message_id": "stream-msg-1"},
run_id=run_id,
)
handler.on_llm_end(
LLMResult(
generations=[
[
ChatGeneration(
message=AIMessage(content="hello", id="final-msg-1")
)
]
]
),
run_id=run_id,
)
assert len(emitted) == 1
+17 -4
View File
@@ -1348,10 +1348,11 @@ wheels = [
[[package]]
name = "langchain-core"
version = "1.3.0a2"
version = "1.3.2"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "jsonpatch" },
{ name = "langchain-protocol" },
{ name = "langsmith" },
{ name = "packaging" },
{ name = "pydantic" },
@@ -1360,9 +1361,21 @@ dependencies = [
{ name = "typing-extensions" },
{ name = "uuid-utils" },
]
sdist = { url = "https://files.pythonhosted.org/packages/af/bc/0bff31fcaff174d86031cc713471a3e85ed4ec8e5cd95ad0217f2aced20e/langchain_core-1.3.0a2.tar.gz", hash = "sha256:52d978c84552b74b9a3f16c1fced84f9e27cc96d7a67c601925ce6cbc4ea3cf9", size = 854580, upload-time = "2026-04-13T14:37:55.745Z" }
sdist = { url = "https://files.pythonhosted.org/packages/a8/03/7219502e8ca728d65eb44d7a3eb60239230742a70dbfc9241b9bfd61c4ab/langchain_core-1.3.2.tar.gz", hash = "sha256:fd7a50b2f28ba561fd9d7f5d2760bc9e06cf00cdf820a3ccafe88a94ffa8d5b7", size = 911813, upload-time = "2026-04-24T15:49:23.699Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/0e/14/03c09686602567059f26af29de0c44546a83af2f2aa29925e61040e43ea2/langchain_core-1.3.0a2-py3-none-any.whl", hash = "sha256:9e929a34f0b0c6c1255e395a1de34f8626893ceb4cdae550a22a0bd18c87be54", size = 510233, upload-time = "2026-04-13T14:37:54.277Z" },
{ url = "https://files.pythonhosted.org/packages/7d/d5/8fa4431007cbb7cfed7590f4d6a5dea3ad724f4174d248f6642ef5ce7d05/langchain_core-1.3.2-py3-none-any.whl", hash = "sha256:d44a66127f9f8db735bdfd0ab9661bccb47a97113cfd3f2d89c74864422b7274", size = 542390, upload-time = "2026-04-24T15:49:21.991Z" },
]
[[package]]
name = "langchain-protocol"
version = "0.0.11"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/d0/bb/38b5eaefa41c67735eedd9f9a2568b11c9eb376fa129a5edd7cc3dcde071/langchain_protocol-0.0.11.tar.gz", hash = "sha256:c276e2373b5ac691fc7ac9a72019d55182444ce8e89385c3f7e9f0185d0aace7", size = 6622, upload-time = "2026-04-23T22:13:16.771Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/0b/fa/6a8ecad8472b182f2caf9d83fd89f40fc1590cb96546d90089b7869b7f5e/langchain_protocol-0.0.11-py3-none-any.whl", hash = "sha256:364da1faf6f5d3001413bede792c1a822c0f23ae55d1ce1266ca7d8e80e79011", size = 6778, upload-time = "2026-04-23T22:13:15.712Z" },
]
[[package]]
@@ -1439,7 +1452,7 @@ test = [
[package.metadata]
requires-dist = [
{ name = "langchain-core", specifier = "==1.3.0a2" },
{ name = "langchain-core", specifier = "==1.3.2" },
{ name = "langgraph-checkpoint", editable = "../checkpoint" },
{ name = "langgraph-prebuilt", editable = "../prebuilt" },
{ name = "langgraph-sdk", editable = "../sdk-py" },
+17 -4
View File
@@ -249,10 +249,11 @@ wheels = [
[[package]]
name = "langchain-core"
version = "1.3.0a2"
version = "1.3.2"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "jsonpatch" },
{ name = "langchain-protocol" },
{ name = "langsmith" },
{ name = "packaging" },
{ name = "pydantic" },
@@ -261,9 +262,21 @@ dependencies = [
{ name = "typing-extensions" },
{ name = "uuid-utils" },
]
sdist = { url = "https://files.pythonhosted.org/packages/af/bc/0bff31fcaff174d86031cc713471a3e85ed4ec8e5cd95ad0217f2aced20e/langchain_core-1.3.0a2.tar.gz", hash = "sha256:52d978c84552b74b9a3f16c1fced84f9e27cc96d7a67c601925ce6cbc4ea3cf9", size = 854580, upload-time = "2026-04-13T14:37:55.745Z" }
sdist = { url = "https://files.pythonhosted.org/packages/a8/03/7219502e8ca728d65eb44d7a3eb60239230742a70dbfc9241b9bfd61c4ab/langchain_core-1.3.2.tar.gz", hash = "sha256:fd7a50b2f28ba561fd9d7f5d2760bc9e06cf00cdf820a3ccafe88a94ffa8d5b7", size = 911813, upload-time = "2026-04-24T15:49:23.699Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/0e/14/03c09686602567059f26af29de0c44546a83af2f2aa29925e61040e43ea2/langchain_core-1.3.0a2-py3-none-any.whl", hash = "sha256:9e929a34f0b0c6c1255e395a1de34f8626893ceb4cdae550a22a0bd18c87be54", size = 510233, upload-time = "2026-04-13T14:37:54.277Z" },
{ url = "https://files.pythonhosted.org/packages/7d/d5/8fa4431007cbb7cfed7590f4d6a5dea3ad724f4174d248f6642ef5ce7d05/langchain_core-1.3.2-py3-none-any.whl", hash = "sha256:d44a66127f9f8db735bdfd0ab9661bccb47a97113cfd3f2d89c74864422b7274", size = 542390, upload-time = "2026-04-24T15:49:21.991Z" },
]
[[package]]
name = "langchain-protocol"
version = "0.0.11"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/d0/bb/38b5eaefa41c67735eedd9f9a2568b11c9eb376fa129a5edd7cc3dcde071/langchain_protocol-0.0.11.tar.gz", hash = "sha256:c276e2373b5ac691fc7ac9a72019d55182444ce8e89385c3f7e9f0185d0aace7", size = 6622, upload-time = "2026-04-23T22:13:16.771Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/0b/fa/6a8ecad8472b182f2caf9d83fd89f40fc1590cb96546d90089b7869b7f5e/langchain_protocol-0.0.11-py3-none-any.whl", hash = "sha256:364da1faf6f5d3001413bede792c1a822c0f23ae55d1ce1266ca7d8e80e79011", size = 6778, upload-time = "2026-04-23T22:13:15.712Z" },
]
[[package]]
@@ -281,7 +294,7 @@ dependencies = [
[package.metadata]
requires-dist = [
{ name = "langchain-core", specifier = "==1.3.0a2" },
{ name = "langchain-core", specifier = "==1.3.2" },
{ name = "langgraph-checkpoint", editable = "../checkpoint" },
{ name = "langgraph-prebuilt", editable = "." },
{ name = "langgraph-sdk", editable = "../sdk-py" },
+17 -4
View File
@@ -262,10 +262,11 @@ wheels = [
[[package]]
name = "langchain-core"
version = "1.3.0a2"
version = "1.3.2"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "jsonpatch" },
{ name = "langchain-protocol" },
{ name = "langsmith" },
{ name = "packaging" },
{ name = "pydantic" },
@@ -274,9 +275,21 @@ dependencies = [
{ name = "typing-extensions" },
{ name = "uuid-utils" },
]
sdist = { url = "https://files.pythonhosted.org/packages/af/bc/0bff31fcaff174d86031cc713471a3e85ed4ec8e5cd95ad0217f2aced20e/langchain_core-1.3.0a2.tar.gz", hash = "sha256:52d978c84552b74b9a3f16c1fced84f9e27cc96d7a67c601925ce6cbc4ea3cf9", size = 854580, upload-time = "2026-04-13T14:37:55.745Z" }
sdist = { url = "https://files.pythonhosted.org/packages/a8/03/7219502e8ca728d65eb44d7a3eb60239230742a70dbfc9241b9bfd61c4ab/langchain_core-1.3.2.tar.gz", hash = "sha256:fd7a50b2f28ba561fd9d7f5d2760bc9e06cf00cdf820a3ccafe88a94ffa8d5b7", size = 911813, upload-time = "2026-04-24T15:49:23.699Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/0e/14/03c09686602567059f26af29de0c44546a83af2f2aa29925e61040e43ea2/langchain_core-1.3.0a2-py3-none-any.whl", hash = "sha256:9e929a34f0b0c6c1255e395a1de34f8626893ceb4cdae550a22a0bd18c87be54", size = 510233, upload-time = "2026-04-13T14:37:54.277Z" },
{ url = "https://files.pythonhosted.org/packages/7d/d5/8fa4431007cbb7cfed7590f4d6a5dea3ad724f4174d248f6642ef5ce7d05/langchain_core-1.3.2-py3-none-any.whl", hash = "sha256:d44a66127f9f8db735bdfd0ab9661bccb47a97113cfd3f2d89c74864422b7274", size = 542390, upload-time = "2026-04-24T15:49:21.991Z" },
]
[[package]]
name = "langchain-protocol"
version = "0.0.11"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/d0/bb/38b5eaefa41c67735eedd9f9a2568b11c9eb376fa129a5edd7cc3dcde071/langchain_protocol-0.0.11.tar.gz", hash = "sha256:c276e2373b5ac691fc7ac9a72019d55182444ce8e89385c3f7e9f0185d0aace7", size = 6622, upload-time = "2026-04-23T22:13:16.771Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/0b/fa/6a8ecad8472b182f2caf9d83fd89f40fc1590cb96546d90089b7869b7f5e/langchain_protocol-0.0.11-py3-none-any.whl", hash = "sha256:364da1faf6f5d3001413bede792c1a822c0f23ae55d1ce1266ca7d8e80e79011", size = 6778, upload-time = "2026-04-23T22:13:15.712Z" },
]
[[package]]
@@ -294,7 +307,7 @@ dependencies = [
[package.metadata]
requires-dist = [
{ name = "langchain-core", specifier = "==1.3.0a2" },
{ name = "langchain-core", specifier = "==1.3.2" },
{ name = "langgraph-checkpoint", editable = "../checkpoint" },
{ name = "langgraph-prebuilt", editable = "../prebuilt" },
{ name = "langgraph-sdk", editable = "." },