mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 12:35:08 +02:00
Merge branch 'nh/messages-content-blocks' into nh/subgraph-lifecycle
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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) == []
|
||||
|
||||
|
||||
Reference in New Issue
Block a user