diff --git a/libs/langgraph/langgraph/stream/run_stream.py b/libs/langgraph/langgraph/stream/run_stream.py index 13bd6fb86..3f2d5b0ac 100644 --- a/libs/langgraph/langgraph/stream/run_stream.py +++ b/libs/langgraph/langgraph/stream/run_stream.py @@ -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. @@ -261,13 +270,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: """Pull one event from the graph and push it through the mux. diff --git a/libs/langgraph/langgraph/stream/transformers.py b/libs/langgraph/langgraph/stream/transformers.py index dec52515f..70a981310 100644 --- a/libs/langgraph/langgraph/stream/transformers.py +++ b/libs/langgraph/langgraph/stream/transformers.py @@ -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. @@ -56,34 +67,165 @@ 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 root-namespace events are captured; tokens from subgraphs are + dropped. 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() + 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 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"]: 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() diff --git a/libs/langgraph/tests/test_stream_messages_transformer.py b/libs/langgraph/tests/test_stream_messages_transformer.py index c8a3275ca..35e62f8d1 100644 --- a/libs/langgraph/tests/test_stream_messages_transformer.py +++ b/libs/langgraph/tests/test_stream_messages_transformer.py @@ -103,6 +103,10 @@ def _make_sync_transformer() -> tuple[MessagesTransformer, EventLog[ChatModelStr proj = t.init() log: EventLog[ChatModelStream] = proj["messages"] log._bind(is_async=False) + # Production subscribes via `iter(log)` from the graph consumer — do that + # up front so `push` during `process` isn't a no-op. Tests read buffered + # items via `log._items` directly rather than re-iterating. + log._subscribed = True t._bind_pump(lambda: False) return t, log @@ -112,6 +116,7 @@ def _make_async_transformer() -> tuple[MessagesTransformer, EventLog[ChatModelSt proj = t.init() log: EventLog[ChatModelStream] = proj["messages"] log._bind(is_async=True) + log._subscribed = True return t, log @@ -167,7 +172,7 @@ class TestProtocolEventRouting: ) # Stream is in the log immediately. log.close() - streams = list(log) + streams = list(log._items) assert len(streams) == 1 assert isinstance(streams[0], ChatModelStream) assert streams[0].message_id == "run-1" @@ -177,7 +182,7 @@ class TestProtocolEventRouting: for evt in _lifecycle(text="hello world"): t.process(_proto_event(evt, run_id="run-1")) log.close() - (stream,) = list(log) + (stream,) = list(log._items) assert stream.done assert stream.output.content == "hello world" @@ -201,7 +206,7 @@ class TestProtocolEventRouting: ) ) log.close() - assert list(log) == [] + assert list(log._items) == [] def test_concurrent_streams_routed_by_run_id(self) -> None: """Two interleaved LLM calls each produce their own stream.""" @@ -213,7 +218,7 @@ class TestProtocolEventRouting: t.process(_proto_event(a, run_id="run-a")) t.process(_proto_event(b, run_id="run-b")) log.close() - streams = list(log) + streams = list(log._items) assert len(streams) == 2 by_id = {s.message_id: s for s in streams} assert by_id["run-a"].output.content == "aaaa" @@ -224,7 +229,7 @@ class TestProtocolEventRouting: for evt in _lifecycle(text="abcdef"): t.process(_proto_event(evt)) log.close() - (stream,) = list(log) + (stream,) = list(log._items) deltas = list(stream._text_proj._deltas) assert "".join(deltas) == "abcdef" @@ -264,7 +269,7 @@ class TestWholeMessageFallback: t, log = _make_sync_transformer() t.process(_whole_msg("the full answer")) log.close() - (stream,) = list(log) + (stream,) = list(log._items) assert stream.done assert stream.output.content == "the full answer" @@ -272,7 +277,7 @@ class TestWholeMessageFallback: t, log = _make_sync_transformer() t.process(_whole_msg("full")) log.close() - (stream,) = list(log) + (stream,) = list(log._items) event_types = [e["event"] for e in stream._events] assert event_types == [ "message-start", @@ -294,7 +299,7 @@ class TestLegacyChunksIgnored: t.process(_v1_chunk("hello")) t.process(_v1_chunk(" world", finish=True)) log.close() - assert list(log) == [] + assert list(log._items) == [] # --------------------------------------------------------------------------- @@ -329,7 +334,7 @@ class TestFiltering: } ) log.close() - assert list(log) == [] + assert list(log._items) == [] # --------------------------------------------------------------------------- @@ -427,11 +432,12 @@ class TestWireRequestMore: 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)) - 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. @@ -449,13 +455,15 @@ class TestViaMux: v = ValuesTransformer() mux = StreamMux([v, t], is_async=False) t._bind_pump(lambda: False) + log: EventLog[ChatModelStream] = mux.extensions["messages"] + # Simulate a consumer subscribing (as `run.messages` iteration would). + log._subscribed = True for evt in _lifecycle(text="mux stream"): mux.push(_proto_event(evt)) mux.close() - log: EventLog[ChatModelStream] = mux.extensions["messages"] - (stream,) = list(log) + (stream,) = list(log._items) assert stream.output.content == "mux stream" def test_whole_message_via_mux(self) -> None: @@ -463,12 +471,13 @@ class TestViaMux: v = ValuesTransformer() mux = StreamMux([v, t], is_async=False) t._bind_pump(lambda: False) + log: EventLog[ChatModelStream] = mux.extensions["messages"] + log._subscribed = True mux.push(_whole_msg("result")) mux.close() - log: EventLog[ChatModelStream] = mux.extensions["messages"] - (stream,) = list(log) + (stream,) = list(log._items) assert stream.output.content == "result" @pytest.mark.anyio @@ -476,11 +485,12 @@ class TestViaMux: 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)) - log: EventLog[ChatModelStream] = mux.extensions["messages"] streams = list(log._items) assert len(streams) == 1 msg = await streams[0].output @@ -605,6 +615,45 @@ class TestEndToEnd: msg = await streams[0].output assert msg.content == "async answer" + @pytest.mark.anyio + async def test_nested_async_iteration_yields_text_deltas(self) -> None: + """Iterate `stream.text` inside `async for stream in run.messages`. + + The inner `stream.text` cursor drives the shared graph pump via + `AsyncProjection._arequest_more`, wired by + `MessagesTransformer._bind_apump` and + `AsyncGraphRunStream._wire_arequest_more`. + """ + import asyncio + + model = GenericFakeChatModel(messages=iter(["hello world"])) + + 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"}) + + async def consume_nested() -> list[str]: + collected: list[str] = [] + async for stream in run.messages: + async for delta in stream.text: + collected.append(delta) + return collected + + deltas = await asyncio.wait_for(consume_nested(), timeout=2.0) + assert "".join(deltas) == "hello world" + class TestEndToEndV2Invoke: """Nodes call `model.invoke()`; `StreamingHandler` routes through v2. diff --git a/libs/langgraph/tests/test_streaming_handler.py b/libs/langgraph/tests/test_streaming_handler.py index 835158ed5..3f5cdc8e2 100644 --- a/libs/langgraph/tests/test_streaming_handler.py +++ b/libs/langgraph/tests/test_streaming_handler.py @@ -913,24 +913,42 @@ class TestValuesTransformer: class TestMessagesTransformer: def test_captures_root_messages(self) -> None: + """Protocol-event lifecycle produces a ChatModelStream in the log.""" t = MessagesTransformer() t.init() t._log._bind(is_async=False) + t._bind_pump(lambda: False) it = iter(t._log) - t.process(_event("messages", ("chunk", {"meta": True}))) + meta = {"langgraph_node": "llm", "run_id": "run-1"} + for evt in ( + {"event": "message-start", "role": "ai", "message_id": "run-1"}, + {"event": "message-finish", "reason": "stop"}, + ): + t.process(_event("messages", (evt, meta))) t._log.close() items = list(it) assert len(items) == 1 - assert items[0] == ("chunk", {"meta": True}) + # Items in the messages log are ChatModelStream objects, not raw + # tuples — the content-block-centric projection. + assert hasattr(items[0], "dispatch") + assert items[0].message_id == "run-1" def test_ignores_non_root_namespace(self) -> None: t = MessagesTransformer() t.init() t._log._bind(is_async=False) + t._bind_pump(lambda: False) it = iter(t._log) - t.process(_event("messages", ("chunk", {}), namespace=["sub"])) + meta = {"langgraph_node": "llm", "run_id": "run-1"} + t.process( + _event( + "messages", + ({"event": "message-start", "message_id": "run-1"}, meta), + namespace=["sub"], + ) + ) t._log.close() assert list(it) == []