Merge branch 'nh/streaming-transformer' into nh/messages-content-blocks

# Conflicts:
#	libs/langgraph/langgraph/stream/streaming_handler.py
This commit is contained in:
Nick Hollon
2026-04-20 16:13:08 -04:00
6 changed files with 269 additions and 295 deletions
+7
View File
@@ -1045,6 +1045,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
interrupt_after: All | list[str] | None = None,
debug: bool = False,
name: str | None = None,
transformers: Sequence[Callable[[], Any]] | None = None,
) -> CompiledStateGraph[StateT, ContextT, InputT, OutputT]:
"""Compiles the `StateGraph` into a `CompiledStateGraph` object.
@@ -1077,6 +1078,11 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
interrupt_after: An optional list of node names to interrupt after.
debug: A flag indicating whether to enable debug mode.
name: The name to use for the compiled graph.
transformers: Optional sequence of zero-arg factories returning
`StreamTransformer` instances. Registered on the compiled
graph and instantiated per-run whenever `stream_v2` /
`astream_v2` is called. Appended after the built-in
`ValuesTransformer` and `MessagesTransformer`.
Returns:
CompiledStateGraph: The compiled `StateGraph`.
@@ -1159,6 +1165,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
store=store,
cache=cache,
name=name or "LangGraph",
stream_transformers=transformers,
)
compiled._serde_allowlist = serde_allowlist
+149
View File
@@ -344,6 +344,23 @@ class NodeBuilder:
)
def _merge_v2_messages_flag(
config: RunnableConfig | None,
) -> RunnableConfig:
"""Return a config with the v2 messages flag set in `configurable`.
Signals to pregel that `stream_mode="messages"` should attach
`StreamMessagesHandlerV2` for this call so invoke-time model runs
route through the v2 event generator and their protocol events
reach the messages channel.
"""
merged: RunnableConfig = dict(config or {}) # type: ignore[assignment]
configurable = dict(merged.get(CONF) or {})
configurable[CONFIG_KEY_STREAM_MESSAGES_V2] = True
merged[CONF] = configurable
return merged
class Pregel(
PregelProtocol[StateT, ContextT, InputT, OutputT],
Generic[StateT, ContextT, InputT, OutputT],
@@ -675,6 +692,7 @@ class Pregel(
config: RunnableConfig | None = None,
trigger_to_nodes: Mapping[str, Sequence[str]] | None = None,
name: str = "LangGraph",
stream_transformers: Sequence[Callable[[], Any]] | None = None,
**deprecated_kwargs: Unpack[DeprecatedKwargs],
) -> None:
if (
@@ -721,6 +739,9 @@ class Pregel(
self.config = config
self.trigger_to_nodes = trigger_to_nodes or {}
self.name = name
self._stream_transformers: tuple[Callable[[], Any], ...] = tuple(
stream_transformers or ()
)
self._serde_allowlist: set[tuple[str, ...]] | None = None
if auto_validate:
self.validate()
@@ -3251,6 +3272,134 @@ class Pregel(
await asyncio.shield(run_manager.on_chain_error(e))
raise
def stream_v2(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
transformers: Sequence[Any] | None = None,
) -> Any:
"""Start a sync v2 streaming run driven by transformer projections.
Builds a `StreamMux` from the built-in `ValuesTransformer` /
`MessagesTransformer`, this graph's compile-time
`stream_transformers`, and any additional `transformers=`
supplied at the call site. Returns a `GraphRunStream` that the
caller drives by iterating any projection — no background
thread.
Args:
input: Graph input.
config: Optional runnable config forwarded to the graph.
interrupt_before: Nodes to interrupt before, if any.
interrupt_after: Nodes to interrupt after, if any.
transformers: Extra transformer instances appended after
compile-time `stream_transformers`.
Returns:
A `GraphRunStream` the caller iterates to drive the run.
"""
from langgraph.stream._mux import StreamMux
from langgraph.stream.run_stream import GraphRunStream
from langgraph.stream.transformers import (
MessagesTransformer,
ValuesTransformer,
)
values_t = ValuesTransformer()
compiled_instances = [f() for f in self._stream_transformers]
mux = StreamMux(
[
values_t,
MessagesTransformer(),
*compiled_instances,
*(transformers or ()),
],
is_async=False,
)
graph_iter = iter(
self.stream(
input,
_merge_v2_messages_flag(config),
stream_mode=[
"values",
"updates",
"messages",
"custom",
"checkpoints",
"tasks",
"debug",
],
subgraphs=True,
version="v2",
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
)
)
return GraphRunStream(graph_iter, mux, values_t)
async def astream_v2(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
transformers: Sequence[Any] | None = None,
) -> Any:
"""Async counterpart to `stream_v2`.
Returns an `AsyncGraphRunStream` whose projections can be awaited
concurrently; each subscribed cursor drives the pump when its
buffer is empty.
Args:
input: Graph input.
config: Optional runnable config forwarded to the graph.
interrupt_before: Nodes to interrupt before, if any.
interrupt_after: Nodes to interrupt after, if any.
transformers: Extra transformer instances appended after
compile-time `stream_transformers`.
"""
from langgraph.stream._mux import StreamMux
from langgraph.stream.run_stream import AsyncGraphRunStream
from langgraph.stream.transformers import (
MessagesTransformer,
ValuesTransformer,
)
values_t = ValuesTransformer()
compiled_instances = [f() for f in self._stream_transformers]
mux = StreamMux(
[
values_t,
MessagesTransformer(),
*compiled_instances,
*(transformers or ()),
],
is_async=True,
)
graph_aiter = self.astream(
input,
_merge_v2_messages_flag(config),
stream_mode=[
"values",
"updates",
"messages",
"custom",
"checkpoints",
"tasks",
"debug",
],
subgraphs=True,
version="v2",
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
).__aiter__()
return AsyncGraphRunStream(graph_aiter, mux, values_t)
@overload
def invoke(
self,
+3 -4
View File
@@ -1,14 +1,14 @@
"""Streaming infrastructure for LangGraph.
Provides a `StreamingHandler` that wraps a compiled graph and exposes
ergonomic streaming projections through a transformer pipeline.
Compile a graph with `transformers=[...]` and call `graph.stream_v2()` /
`graph.astream_v2()` to drive a transformer pipeline that projects the
graph's raw events into ergonomic per-channel streams.
"""
from langgraph.stream._event_log import EventLog
from langgraph.stream._types import ProtocolEvent, StreamTransformer
from langgraph.stream.run_stream import AsyncGraphRunStream, GraphRunStream
from langgraph.stream.stream_channel import StreamChannel
from langgraph.stream.streaming_handler import StreamingHandler
__all__ = [
"AsyncGraphRunStream",
@@ -17,5 +17,4 @@ __all__ = [
"ProtocolEvent",
"StreamChannel",
"StreamTransformer",
"StreamingHandler",
]
@@ -1,168 +0,0 @@
from __future__ import annotations
from collections.abc import Sequence
from typing import Any
from langchain_core.runnables import RunnableConfig
from langgraph._internal._constants import CONF, CONFIG_KEY_STREAM_MESSAGES_V2
from langgraph.pregel import Pregel
from langgraph.stream._mux import StreamMux
from langgraph.stream._types import StreamTransformer
from langgraph.stream.run_stream import AsyncGraphRunStream, GraphRunStream
from langgraph.stream.transformers import MessagesTransformer, ValuesTransformer
from langgraph.types import All, StreamMode
def _merge_v2_messages_flag(
config: RunnableConfig | None,
) -> RunnableConfig:
"""Return a config with the v2 messages flag set in `configurable`.
Signals to pregel that `stream_mode="messages"` should attach
`StreamMessagesHandlerV2` for this call so invoke-time model runs
route through the v2 event generator and their protocol events
reach the messages channel.
"""
merged: RunnableConfig = dict(config or {}) # type: ignore[assignment]
configurable = dict(merged.get(CONF) or {})
configurable[CONFIG_KEY_STREAM_MESSAGES_V2] = True
merged[CONF] = configurable
return merged
# All stream modes to request from the graph.
STREAM_V2_MODES: list[StreamMode] = [
"values",
"updates",
"messages",
"custom",
"checkpoints",
"tasks",
"debug",
]
class StreamingHandler:
"""Wrap a compiled graph with ergonomic streaming projections.
Example:
```python
handler = StreamingHandler(graph)
# Sync
run = handler.stream(input_data)
for state in run.values:
print(state)
output = run.output
# Async — terminal accessors are methods so a missing `await`
# fails loudly instead of silently yielding a coroutine.
run = await handler.astream(input_data)
async for state in run.values:
print(state)
output = await run.output()
```
"""
def __init__(self, graph: Pregel) -> None:
"""Initialize the handler.
Args:
graph: A compiled LangGraph graph to stream from.
"""
self._graph = graph
def stream(
self,
input: Any,
config: RunnableConfig | None = None,
*,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
transformers: list[StreamTransformer] | None = None,
) -> GraphRunStream:
"""Start a sync streaming run.
Returns a GraphRunStream immediately. The caller's iteration on
any projection drives the graph forward — no background thread
is used. This matches v1's model where the caller's `for` loop
is the pump.
Args:
input: Graph input.
config: Optional runnable config forwarded to the graph.
interrupt_before: Nodes to interrupt before, if any.
interrupt_after: Nodes to interrupt after, if any.
transformers: User transformers appended after the built-in
`ValuesTransformer` and `MessagesTransformer`.
Returns:
A GraphRunStream the caller can iterate to drive the run.
"""
values_t = ValuesTransformer()
mux = StreamMux(
[values_t, MessagesTransformer(), *(transformers or ())],
is_async=False,
)
graph_iter = iter(
self._graph.stream(
input,
_merge_v2_messages_flag(config),
stream_mode=STREAM_V2_MODES,
subgraphs=True,
version="v2",
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
)
)
return GraphRunStream(graph_iter, mux, values_t)
async def astream(
self,
input: Any,
config: RunnableConfig | None = None,
*,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
transformers: list[StreamTransformer] | None = None,
) -> AsyncGraphRunStream:
"""Start an async streaming run.
Returns an AsyncGraphRunStream immediately. The caller's
iteration on any projection drives the graph forward — there
is no background task. Concurrent consumers share a
single-flight pump via an internal `asyncio.Lock`.
Args:
input: Graph input.
config: Optional runnable config forwarded to the graph.
interrupt_before: Nodes to interrupt before, if any.
interrupt_after: Nodes to interrupt after, if any.
transformers: User transformers appended after the built-in
`ValuesTransformer` and `MessagesTransformer`.
Returns:
An AsyncGraphRunStream whose projections can be awaited
concurrently; each subscribed cursor drives the pump when
its buffer is empty.
"""
values_t = ValuesTransformer()
mux = StreamMux(
[values_t, MessagesTransformer(), *(transformers or ())],
is_async=True,
)
graph_aiter = self._graph.astream(
input,
_merge_v2_messages_flag(config),
stream_mode=STREAM_V2_MODES,
subgraphs=True,
version="v2",
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
).__aiter__()
return AsyncGraphRunStream(graph_aiter, mux, values_t)
@@ -1,4 +1,4 @@
"""Tests for the StreamingHandler and its supporting infrastructure."""
"""Tests for `Pregel.stream_v2` / `astream_v2` and the transformer pipeline."""
from __future__ import annotations
@@ -17,7 +17,6 @@ from langgraph.graph import StateGraph
from langgraph.stream import (
EventLog,
StreamChannel,
StreamingHandler,
StreamTransformer,
)
from langgraph.stream._convert import convert_to_protocol_event
@@ -404,15 +403,15 @@ class TestStreamChannel:
# ---------------------------------------------------------------------------
# StreamingHandler sync tests
# stream_v2 sync tests
# ---------------------------------------------------------------------------
class TestStreamingHandlerSync:
class TestStreamV2Sync:
def test_values_projection(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
snapshots = list(run.values)
# Should have at least the initial + per-node snapshots.
assert len(snapshots) >= 1
@@ -423,8 +422,8 @@ class TestStreamingHandlerSync:
def test_output(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
output = run.output
assert output is not None
assert output["value"] == "xAB"
@@ -432,8 +431,8 @@ class TestStreamingHandlerSync:
def test_raw_event_iteration(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
events = list(run)
assert len(events) > 0
for event in events:
@@ -444,8 +443,8 @@ class TestStreamingHandlerSync:
def test_extensions_has_native_keys(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
# Drain events so the run completes.
_ = run.output
assert "values" in run.extensions
@@ -457,8 +456,8 @@ class TestStreamingHandlerSync:
def test_extensions_is_read_only(self) -> None:
"""`run.extensions` must reject mutations so users can't corrupt mux state."""
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
with pytest.raises(TypeError):
run.extensions["new_key"] = object() # type: ignore[index]
with pytest.raises(TypeError):
@@ -466,8 +465,8 @@ class TestStreamingHandlerSync:
def test_custom_stream_events(self) -> None:
graph = _build_custom_stream_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
custom_events = [e for e in run if e["method"] == "custom"]
assert len(custom_events) == 2
assert custom_events[0]["params"]["data"] == {"step": "start"}
@@ -475,8 +474,8 @@ class TestStreamingHandlerSync:
def test_interleave_values_and_messages(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
tagged = list(run.interleave("values", "messages"))
names = [name for name, _ in tagged]
@@ -489,8 +488,8 @@ class TestStreamingHandlerSync:
def test_abort_marks_exhausted_and_closes_mux(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
values_iter = iter(run.values)
# Consume one item so the pump advances.
_ = next(values_iter)
@@ -503,64 +502,64 @@ class TestStreamingHandlerSync:
def test_context_manager_calls_abort_on_exit(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
with handler.stream({"value": "x", "items": []}) as run:
handler = graph
with handler.stream_v2({"value": "x", "items": []}) as run:
values_iter = iter(run.values)
_ = next(values_iter)
assert run._exhausted is True
def test_interleave_unknown_projection(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
with pytest.raises(KeyError):
list(run.interleave("values", "does_not_exist"))
class TestStreamingHandlerSyncErrors:
class TestStreamV2SyncErrors:
def test_error_propagation_output(self) -> None:
graph = _build_error_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
with pytest.raises(ValueError, match="boom"):
_ = run.output
def test_error_propagation_values(self) -> None:
graph = _build_error_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
with pytest.raises(ValueError, match="boom"):
list(run.values)
def test_error_propagation_raw_events(self) -> None:
graph = _build_error_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
with pytest.raises(ValueError, match="boom"):
list(run)
def test_error_propagation_interrupted(self) -> None:
"""`run.interrupted` should raise on a failed run, not silently return False."""
graph = _build_error_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
with pytest.raises(ValueError, match="boom"):
_ = run.interrupted
def test_error_propagation_interrupts(self) -> None:
"""`run.interrupts` should raise on a failed run."""
graph = _build_error_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
with pytest.raises(ValueError, match="boom"):
_ = run.interrupts
class TestStreamingHandlerSyncInterrupt:
class TestStreamV2SyncInterrupt:
def test_interrupted(self) -> None:
graph = _build_interrupt_graph()
handler = StreamingHandler(graph)
run = handler.stream(
handler = graph
run = handler.stream_v2(
{"value": "x", "items": []},
{"configurable": {"thread_id": "t1"}},
)
@@ -570,17 +569,17 @@ class TestStreamingHandlerSyncInterrupt:
# ---------------------------------------------------------------------------
# StreamingHandler async tests
# astream_v2 async tests
# ---------------------------------------------------------------------------
class TestStreamingHandlerAsync:
class TestStreamV2Async:
@pytest.mark.anyio
@NEEDS_CONTEXTVARS
async def test_values_projection(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
snapshots = [s async for s in run.values]
assert len(snapshots) >= 1
last = snapshots[-1]
@@ -591,8 +590,8 @@ class TestStreamingHandlerAsync:
@NEEDS_CONTEXTVARS
async def test_output(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
output = await run.output()
assert output is not None
assert output["value"] == "xAB"
@@ -602,8 +601,8 @@ class TestStreamingHandlerAsync:
@NEEDS_CONTEXTVARS
async def test_raw_event_iteration(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
events = [e async for e in run]
assert len(events) > 0
for event in events:
@@ -613,8 +612,8 @@ class TestStreamingHandlerAsync:
@NEEDS_CONTEXTVARS
async def test_abort_marks_exhausted_and_closes_mux(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
values_iter = aiter(run.values)
_ = await anext(values_iter)
await run.abort()
@@ -628,8 +627,8 @@ class TestStreamingHandlerAsync:
@NEEDS_CONTEXTVARS
async def test_async_context_manager_calls_abort_on_exit(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
async with run:
values_iter = aiter(run.values)
_ = await anext(values_iter)
@@ -639,8 +638,8 @@ class TestStreamingHandlerAsync:
@NEEDS_CONTEXTVARS
async def test_extensions_has_native_keys(self) -> None:
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
_ = await run.output()
assert "values" in run.extensions
assert "messages" in run.extensions
@@ -648,13 +647,13 @@ class TestStreamingHandlerAsync:
assert run.messages is run.extensions["messages"]
class TestStreamingHandlerAsyncErrors:
class TestStreamV2AsyncErrors:
@pytest.mark.anyio
@NEEDS_CONTEXTVARS
async def test_error_propagation_output(self) -> None:
graph = _build_error_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
with pytest.raises(ValueError, match="boom"):
await run.output()
@@ -662,8 +661,8 @@ class TestStreamingHandlerAsyncErrors:
@NEEDS_CONTEXTVARS
async def test_error_propagation_values(self) -> None:
graph = _build_error_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
with pytest.raises(ValueError, match="boom"):
async for _ in run.values:
pass
@@ -672,8 +671,8 @@ class TestStreamingHandlerAsyncErrors:
@NEEDS_CONTEXTVARS
async def test_error_propagation_raw_events(self) -> None:
graph = _build_error_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
with pytest.raises(ValueError, match="boom"):
async for _ in run:
pass
@@ -683,8 +682,8 @@ class TestStreamingHandlerAsyncErrors:
async def test_error_propagation_interrupted(self) -> None:
"""`await run.interrupted()` should raise on a failed async run."""
graph = _build_error_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
with pytest.raises(ValueError, match="boom"):
await run.interrupted()
@@ -693,19 +692,19 @@ class TestStreamingHandlerAsyncErrors:
async def test_error_propagation_interrupts(self) -> None:
"""`await run.interrupts()` should raise on a failed async run."""
graph = _build_error_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
with pytest.raises(ValueError, match="boom"):
await run.interrupts()
class TestStreamingHandlerAsyncInterrupt:
class TestStreamV2AsyncInterrupt:
@pytest.mark.anyio
@NEEDS_CONTEXTVARS
async def test_interrupted(self) -> None:
graph = _build_interrupt_graph()
handler = StreamingHandler(graph)
run = await handler.astream(
handler = graph
run = await handler.astream_v2(
{"value": "x", "items": []},
{"configurable": {"thread_id": "t2"}},
)
@@ -714,13 +713,13 @@ class TestStreamingHandlerAsyncInterrupt:
assert len(await run.interrupts()) > 0
class TestStreamingHandlerAsyncCustom:
class TestStreamV2AsyncCustom:
@pytest.mark.anyio
@NEEDS_CONTEXTVARS
async def test_custom_stream_events(self) -> None:
graph = _build_custom_stream_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
events = [e async for e in run]
custom_events = [e for e in events if e["method"] == "custom"]
assert len(custom_events) == 2
@@ -1100,9 +1099,9 @@ class TestCustomTransformer:
return True
graph = _build_simple_graph()
handler = StreamingHandler(graph)
handler = graph
counter_t = CounterTransformer()
run = handler.stream({"value": "x", "items": []}, transformers=[counter_t])
run = handler.stream_v2({"value": "x", "items": []}, transformers=[counter_t])
assert "counter" in run.extensions
# Subscribe before driving the run so channel pushes are retained.
counter_iter = iter(run.extensions["counter"])
@@ -1131,9 +1130,9 @@ class TestCustomTransformer:
return True
graph = _build_simple_graph()
handler = StreamingHandler(graph)
handler = graph
foo_t = FooTransformer()
run = handler.stream({"value": "x", "items": []}, transformers=[foo_t])
run = handler.stream_v2({"value": "x", "items": []}, transformers=[foo_t])
# Subscribe before driving the run.
foo_iter = iter(run.foo)
_ = run.output
@@ -1161,8 +1160,8 @@ class TestCustomTransformer:
return True
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = handler.stream(
handler = graph
run = handler.stream_v2(
{"value": "x", "items": []}, transformers=[EmitterTransformer()]
)
events = list(run)
@@ -1220,12 +1219,12 @@ class TestCustomTransformer:
return True
graph = _build_simple_graph()
handler = StreamingHandler(graph)
handler = graph
with pytest.raises(
ValueError,
match=r"conflict.*'values'.*ValuesTransformer",
):
handler.stream(
handler.stream_v2(
{"value": "x", "items": []},
transformers=[ConflictTransformer()],
)
@@ -1322,9 +1321,9 @@ class TestEventLogAutoLifecycle:
return True
graph = _build_simple_graph()
handler = StreamingHandler(graph)
handler = graph
t = MinimalTransformer()
run = handler.stream({"value": "x", "items": []}, transformers=[t])
run = handler.stream_v2({"value": "x", "items": []}, transformers=[t])
minimal_iter = iter(run.extensions["minimal"])
_ = run.output
items = list(minimal_iter)
@@ -1623,8 +1622,8 @@ class TestAsyncTransformerLane:
self._log.close()
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = await handler.astream(
handler = graph
run = await handler.astream_v2(
{"value": "x", "items": []},
transformers=[Scorer()],
)
@@ -1649,8 +1648,8 @@ class TestMemoryBounds:
"""With a single sync consumer, the pump produces exactly one event
per cursor advance, so the buffer never holds more than one."""
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
events_iter = iter(run)
max_buffered = 0
count = 0
@@ -1667,8 +1666,8 @@ class TestMemoryBounds:
"""Projections without a subscriber drop pushes silently —
their buffers stay empty regardless of run length."""
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
# Subscribe to main events only; leave values and messages unsubscribed.
list(run)
values_log = run.extensions["values"]
@@ -1683,8 +1682,8 @@ class TestMemoryBounds:
process() without populating the log, so the values log buffer
stays empty even across a full run."""
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
_ = run.output
values_log = run.extensions["values"]
assert len(values_log._items) == 0
@@ -1693,8 +1692,8 @@ class TestMemoryBounds:
def test_drained_subscriber_buffer_returns_to_empty(self) -> None:
"""After fully draining a subscribed log, the internal deque is empty."""
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = handler.stream({"value": "x", "items": []})
handler = graph
run = handler.stream_v2({"value": "x", "items": []})
values_log = run.extensions["values"]
list(run.values)
assert len(values_log._items) == 0
@@ -1704,8 +1703,8 @@ class TestMemoryBounds:
async def test_async_single_consumer_buffer_stays_at_most_one(self) -> None:
"""Same drain-on-consume guarantee for the async lane."""
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
max_buffered = 0
count = 0
async for _ in run:
@@ -1719,8 +1718,8 @@ class TestMemoryBounds:
async def test_async_unsubscribed_projections_never_accumulate(self) -> None:
"""Projections with no async subscriber stay empty under astream."""
graph = _build_simple_graph()
handler = StreamingHandler(graph)
run = await handler.astream({"value": "x", "items": []})
handler = graph
run = await handler.astream_v2({"value": "x", "items": []})
_ = await run.output()
values_log = run.extensions["values"]
messages_log = run.extensions["messages"]
@@ -24,7 +24,6 @@ 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.streaming_handler import StreamingHandler
from langgraph.stream.transformers import MessagesTransformer, ValuesTransformer
TS = int(time.time() * 1000)
@@ -499,7 +498,7 @@ class TestViaMux:
# ---------------------------------------------------------------------------
# End-to-end: full graph → StreamingHandler → run.messages
# End-to-end: full graph → stream_v2 → run.messages
# ---------------------------------------------------------------------------
@@ -533,8 +532,7 @@ class TestEndToEnd:
.compile()
)
handler = StreamingHandler(graph)
run = handler.stream({"messages": "hi"})
run = graph.stream_v2({"messages": "hi"})
streams = list(run.messages)
assert len(streams) == 1
@@ -557,8 +555,7 @@ class TestEndToEnd:
.compile()
)
handler = StreamingHandler(graph)
run = handler.stream({"messages": "go"})
run = graph.stream_v2({"messages": "go"})
# Pull the stream handle out, then iterate its text deltas.
(stream,) = list(run.messages)
@@ -579,8 +576,7 @@ class TestEndToEnd:
.compile()
)
handler = StreamingHandler(graph)
run = handler.stream({"messages": "hi"})
run = graph.stream_v2({"messages": "hi"})
streams = list(run.messages)
assert len(streams) == 1
@@ -603,8 +599,7 @@ class TestEndToEnd:
.compile()
)
handler = StreamingHandler(graph)
run = await handler.astream({"messages": "hi"})
run = await graph.astream_v2({"messages": "hi"})
streams = []
async for stream in run.messages:
@@ -641,8 +636,7 @@ class TestEndToEnd:
.compile()
)
handler = StreamingHandler(graph)
run = await handler.astream({"messages": "hi"})
run = await graph.astream_v2({"messages": "hi"})
async def consume_nested() -> list[str]:
collected: list[str] = []
@@ -656,11 +650,11 @@ class TestEndToEnd:
class TestEndToEndV2Invoke:
"""Nodes call `model.invoke()`; `StreamingHandler` routes through v2.
"""Nodes call `model.invoke()`; `stream_v2` routes through v2.
Exercises the auto-routing path added in
`feat(core): route invoke through v2 event path for
_V2StreamingCallbackHandler`: `StreamingHandler` injects
_V2StreamingCallbackHandler`: `stream_v2` injects
`CONFIG_KEY_STREAM_MESSAGES_V2` into the config, pregel attaches
`StreamMessagesHandlerV2`, `BaseChatModel._should_stream_v2` sees the
v2 marker and drives the protocol event generator, and
@@ -682,8 +676,7 @@ class TestEndToEndV2Invoke:
.compile()
)
handler = StreamingHandler(graph)
run = handler.stream({"messages": "hi"})
run = graph.stream_v2({"messages": "hi"})
streams = list(run.messages)
assert len(streams) == 1, (
@@ -709,8 +702,7 @@ class TestEndToEndV2Invoke:
.compile()
)
handler = StreamingHandler(graph)
run = handler.stream({"messages": "go"})
run = graph.stream_v2({"messages": "go"})
(stream,) = list(run.messages)
events = list(stream)
@@ -743,8 +735,7 @@ class TestEndToEndV2Invoke:
.compile()
)
handler = StreamingHandler(graph)
run = handler.stream({"messages": "hi"})
run = graph.stream_v2({"messages": "hi"})
(stream,) = list(run.messages)
assembled = "".join(stream.text)
@@ -771,8 +762,7 @@ class TestEndToEndV2Invoke:
.compile()
)
handler = StreamingHandler(graph)
run = handler.stream({"messages": "hi"})
run = graph.stream_v2({"messages": "hi"})
streams = list(run.messages)
assert len(streams) == 2
@@ -802,8 +792,7 @@ class TestEndToEndV2Invoke:
.compile()
)
handler = StreamingHandler(graph)
run = handler.stream({"messages": "hi"})
run = graph.stream_v2({"messages": "hi"})
streams = list(run.messages)
assert len(streams) == 2
@@ -815,7 +804,7 @@ class TestEndToEndV2Invoke:
@pytest.mark.anyio
async def test_ainvoke_with_v2_marker_populates_messages(self) -> None:
"""Async mirror: `model.ainvoke()` + `StreamingHandler.astream()`."""
"""Async mirror: `model.ainvoke()` + `astream_v2`."""
model = GenericFakeChatModel(messages=iter(["async invoke"]))
async def call_model(state: MessagesState) -> dict[str, Any]:
@@ -829,8 +818,7 @@ class TestEndToEndV2Invoke:
.compile()
)
handler = StreamingHandler(graph)
run = await handler.astream({"messages": "hi"})
run = await graph.astream_v2({"messages": "hi"})
streams = []
async for stream in run.messages:
@@ -844,8 +832,8 @@ class TestEndToEndV2Invoke:
class TestDirectMessagesModeStaysV1:
"""Regression guard: direct `graph.stream(stream_mode="messages")`
(no `StreamingHandler`) must keep the v1 `(AIMessageChunk, metadata)`
tuple shape. The v2 flag is only injected by `StreamingHandler`.
(no `stream_v2`) must keep the v1 `(AIMessageChunk, metadata)`
tuple shape. The v2 flag is only injected by `stream_v2` / `astream_v2`.
"""
def test_direct_graph_stream_messages_yields_ai_message_chunks(self) -> None:
@@ -870,7 +858,7 @@ class TestDirectMessagesModeStaysV1:
payload, _metadata = part
assert isinstance(payload, AIMessageChunk), (
"direct graph.stream(stream_mode='messages') leaked v2 "
"event dicts — StreamingHandler flag bled through."
"event dicts — stream_v2 flag bled through."
)
assembled = "".join(
p[0].content for p in parts if isinstance(p[0].content, str)