Wire async pump into MessagesTransformer streams

AsyncChatModelStream projections deadlocked when iterated inside
the outer run.messages cursor: the inner stream awaited an
asyncio.Event that nothing was driving while the outer cursor was
suspended.

Plumb the langchain-core async pump hook down to each stream:
- MessagesTransformer gains _bind_apump (mirror of _bind_pump) and
  prefers async wiring in _make_stream.
- AsyncGraphRunStream._wire_arequest_more calls _bind_apump on any
  transformer that exposes it.

Test helpers: EventLog.push is a no-op before subscription; the
unit-test helpers and TestViaMux setups now pre-subscribe the log
(simulating what run.messages iteration does in production) and
verify pushed items via log._items. Flip the known-failure nested
iteration test to pass.
This commit is contained in:
Nick Hollon
2026-04-18 13:20:06 -04:00
parent b6a196fac6
commit ad0146a4de
3 changed files with 107 additions and 25 deletions
+11 -1
View File
@@ -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.
@@ -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":
@@ -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.