diff --git a/libs/langgraph/langgraph/stream/run_stream.py b/libs/langgraph/langgraph/stream/run_stream.py index d14829703..3f2d5b0ac 100644 --- a/libs/langgraph/langgraph/stream/run_stream.py +++ b/libs/langgraph/langgraph/stream/run_stream.py @@ -270,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 685269cf1..70a981310 100644 --- a/libs/langgraph/langgraph/stream/transformers.py +++ b/libs/langgraph/langgraph/stream/transformers.py @@ -14,7 +14,7 @@ from langgraph.stream._event_log import EventLog from langgraph.stream._types import ProtocolEvent, StreamTransformer if TYPE_CHECKING: - from collections.abc import Callable + from collections.abc import Awaitable, Callable class ValuesTransformer(StreamTransformer): @@ -111,6 +111,7 @@ class MessagesTransformer(StreamTransformer): # (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} @@ -119,6 +120,15 @@ class MessagesTransformer(StreamTransformer): """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, *, @@ -126,7 +136,21 @@ class MessagesTransformer(StreamTransformer): node: str | None, message_id: str | None, ) -> ChatModelStream: - """Create a ChatModelStream (sync) or AsyncChatModelStream (async).""" + """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, @@ -134,13 +158,12 @@ class MessagesTransformer(StreamTransformer): message_id=message_id, ) stream.set_request_more(self._pump_fn) - else: - stream = AsyncChatModelStream( - namespace=namespace, - node=node, - message_id=message_id, - ) - return stream + return stream + return AsyncChatModelStream( + namespace=namespace, + node=node, + message_id=message_id, + ) def process(self, event: ProtocolEvent) -> bool: if event["method"] != "messages": 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.