From 08666353fc390bcab1df91ac02d6a1b88307a46e Mon Sep 17 00:00:00 2001 From: Nick Hollon Date: Wed, 29 Apr 2026 09:08:21 -0400 Subject: [PATCH] fix(langgraph): decouple run.output/interrupted/interrupts from ValuesTransformer (#7639) --- libs/langgraph/langgraph/pregel/main.py | 44 +++++----- libs/langgraph/langgraph/stream/run_stream.py | 80 ++++++++++++------- .../langgraph/stream/transformers.py | 50 +++++------- libs/langgraph/tests/test_pregel_stream_v2.py | 38 +++++++++ .../tests/test_stream_messages_transformer.py | 4 +- .../tests/test_stream_subgraph_transformer.py | 12 --- 6 files changed, 135 insertions(+), 93 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index 33a135d25..177fd1d8b 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -3394,25 +3394,25 @@ class Pregel( ) -> 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. + Builds a `StreamMux` from the built-in transformers, 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. + + `run.output`, `run.interrupted` and `run.interrupts` work + regardless of which transformers are registered. Note: Nesting v1 `stream(stream_mode="messages")` inside a node of a `stream_v2` run is not fully supported. The outer v2 - messages handler is inheritable, so it sits in the inner - chat model's callback chain; `BaseChatModel.invoke` then - routes through the v2 event protocol and the inner v1 - messages handler does not see `on_llm_new_token` chunks. - The inner stream still yields a finalized message via - `on_llm_end`, but token-by-token output is lost. Use - `stream_v2` for the inner graph as well, or call - `chat_model.stream(...)` explicitly inside the node, to - get token-level streaming. + messages handler reroutes `BaseChatModel.invoke` through + the v2 event protocol, so the inner v1 handler does not see + `on_llm_new_token` chunks. The inner stream still yields a + finalized message via `on_llm_end`. Use `stream_v2` for + the inner graph as well, or call + `chat_model.stream(...)` explicitly, to get token-level + streaming. Args: input: Graph input. @@ -3444,7 +3444,6 @@ class Pregel( scope=parent_ns, is_async=False, ) - values_t = cast(ValuesTransformer, mux.transformer_by_key("values")) graph_iter = iter( self.stream( input, @@ -3456,7 +3455,7 @@ class Pregel( interrupt_after=interrupt_after, ) ) - return GraphRunStream(graph_iter, mux, values_t) + return GraphRunStream(graph_iter, mux) async def astream_v2( self, @@ -3478,11 +3477,9 @@ class Pregel( `astream(stream_mode="messages")` inside a node of an `astream_v2` run drops `on_llm_new_token` chunks because the outer v2 handler reroutes `BaseChatModel.invoke` - through the v2 event protocol. The inner stream still - yields a finalized message at end-of-call. Use - `astream_v2` for the inner graph as well, or call - `chat_model.astream(...)` explicitly inside the node, to - get token-level streaming. + through the v2 event protocol. Use `astream_v2` for the + inner graph as well, or call `chat_model.astream(...)` + explicitly, to get token-level streaming. Args: input: Graph input. @@ -3511,7 +3508,6 @@ class Pregel( scope=parent_ns, is_async=True, ) - values_t = cast(ValuesTransformer, mux.transformer_by_key("values")) graph_aiter = self.astream( input, patch_configurable(config, {CONFIG_KEY_STREAM_MESSAGES_V2: True}), @@ -3521,7 +3517,7 @@ class Pregel( interrupt_before=interrupt_before, interrupt_after=interrupt_after, ).__aiter__() - return AsyncGraphRunStream(graph_aiter, mux, values_t) + return AsyncGraphRunStream(graph_aiter, mux) @overload def invoke( diff --git a/libs/langgraph/langgraph/stream/run_stream.py b/libs/langgraph/langgraph/stream/run_stream.py index e712fedee..0b17d0cbd 100644 --- a/libs/langgraph/langgraph/stream/run_stream.py +++ b/libs/langgraph/langgraph/stream/run_stream.py @@ -10,7 +10,7 @@ from langgraph.stream._mux import StreamMux from langgraph.stream._types import ProtocolEvent if TYPE_CHECKING: - from langgraph.stream.transformers import SubgraphStatus, ValuesTransformer + from langgraph.stream.transformers import SubgraphStatus def _drive_until_done(pump: Callable[[], bool]) -> None: @@ -44,7 +44,6 @@ class GraphRunStream: self, graph_iter: Iterator[Any] | None, mux: StreamMux, - values_transformer: ValuesTransformer, *, wire_pump: bool = True, ) -> None: @@ -55,8 +54,6 @@ class GraphRunStream: or `None` for nested run streams whose pump is driven by an outer run (e.g. `SubgraphRunStream`). mux: The StreamMux owning projections and the main log. - values_transformer: The built-in values transformer - providing `output` / `interrupted` / `interrupts`. wire_pump: When True (default), bind `_pump_next` as the mux's pump callable. Subclasses that inherit a parent pump via `StreamMux._make_child` should pass False to @@ -65,8 +62,11 @@ class GraphRunStream: self._graph_iter = graph_iter self._mux = mux self.extensions: Mapping[str, Any] = MappingProxyType(mux.extensions) - self._values_transformer = values_transformer self._exhausted = False + self._latest: dict[str, Any] | None = None + self._interrupted = False + self._interrupts: list[Any] = [] + self._scope_list: list[str] = list(mux.scope) for key in mux.native_keys: setattr(self, key, mux.extensions[key]) if wire_pump: @@ -83,6 +83,19 @@ class GraphRunStream: """ mux.bind_pump(self._pump_next) + def _observe_event(self, event: ProtocolEvent) -> None: + """Track values-event state for output/interrupted/interrupts.""" + if event["method"] != "values": + return + params = event["params"] + if params["namespace"] != self._scope_list: + return + self._latest = params["data"] + interrupts = params.get("interrupts", ()) + if interrupts: + self._interrupted = True + self._interrupts.extend(interrupts) + def _pump_next(self) -> bool: """Pull one event from the graph and push it through the mux. @@ -95,7 +108,9 @@ class GraphRunStream: return False try: part = next(self._graph_iter) - self._mux.push(convert_to_protocol_event(part)) + event = convert_to_protocol_event(part) + self._observe_event(event) + self._mux.push(event) return True except StopIteration: self._mux.close() @@ -136,9 +151,9 @@ class GraphRunStream: def output(self) -> dict[str, Any] | None: """Drive the run to completion and return the final state.""" _drive_until_done(self._pump_next) - if (err := self._values_transformer.error) is not None: + if (err := self._mux._events._error) is not None: raise err - return self._values_transformer._latest + return self._latest @property def interrupted(self) -> bool: @@ -149,9 +164,9 @@ class GraphRunStream: BaseException: If the run ended with an error. """ _drive_until_done(self._pump_next) - if (err := self._values_transformer.error) is not None: + if (err := self._mux._events._error) is not None: raise err - return self._values_transformer._interrupted + return self._interrupted @property def interrupts(self) -> list[Any]: @@ -161,9 +176,9 @@ class GraphRunStream: BaseException: If the run ended with an error. """ _drive_until_done(self._pump_next) - if (err := self._values_transformer.error) is not None: + if (err := self._mux._events._error) is not None: raise err - return self._values_transformer._interrupts + return self._interrupts def __iter__(self) -> Iterator[ProtocolEvent]: """Subscribe to the main event log and iterate protocol events.""" @@ -247,7 +262,6 @@ class AsyncGraphRunStream: self, graph_aiter: AsyncIterator[Any] | None, mux: StreamMux, - values_transformer: ValuesTransformer, *, wire_pump: bool = True, ) -> None: @@ -258,8 +272,6 @@ class AsyncGraphRunStream: `None` for nested run streams whose pump is driven by an outer run (e.g. `AsyncSubgraphRunStream`). mux: The StreamMux owning projections and the main log. - values_transformer: The built-in values transformer - providing `output` / `interrupted` / `interrupts`. wire_pump: When True (default), bind `_apump_next` as the mux's async pump callable. Subclasses that inherit a parent pump via `StreamMux._make_child` should pass @@ -268,8 +280,11 @@ class AsyncGraphRunStream: self._graph_aiter = graph_aiter self._mux = mux self.extensions: Mapping[str, Any] = MappingProxyType(mux.extensions) - self._values_transformer = values_transformer self._exhausted = False + self._latest: dict[str, Any] | None = None + self._interrupted = False + self._interrupts: list[Any] = [] + self._scope_list: list[str] = list(mux.scope) self._pump_cond = asyncio.Condition() self._pumping = False for key in mux.native_keys: @@ -277,6 +292,19 @@ class AsyncGraphRunStream: if wire_pump: self._wire_arequest_more(mux) + def _observe_event(self, event: ProtocolEvent) -> None: + """Track values-event state for output/interrupted/interrupts.""" + if event["method"] != "values": + return + params = event["params"] + if params["namespace"] != self._scope_list: + return + self._latest = params["data"] + interrupts = params.get("interrupts", ()) + if interrupts: + self._interrupted = True + self._interrupts.extend(interrupts) + def _wire_arequest_more(self, mux: StreamMux) -> None: """Wire the async pull callback through the mux. @@ -319,7 +347,9 @@ class AsyncGraphRunStream: try: try: part = await self._graph_aiter.__anext__() - await self._mux.apush(convert_to_protocol_event(part)) + event = convert_to_protocol_event(part) + self._observe_event(event) + await self._mux.apush(event) return True except StopAsyncIteration: self._exhausted = True @@ -378,9 +408,9 @@ class AsyncGraphRunStream: BaseException: If the run ended with an error. """ await _adrive_until_done(self._apump_next) - if (err := self._values_transformer.error) is not None: + if (err := self._mux._events._error) is not None: raise err - return self._values_transformer._latest + return self._latest async def interrupted(self) -> bool: """Drive the run to completion and return whether it was @@ -390,9 +420,9 @@ class AsyncGraphRunStream: BaseException: If the run ended with an error. """ await _adrive_until_done(self._apump_next) - if (err := self._values_transformer.error) is not None: + if (err := self._mux._events._error) is not None: raise err - return self._values_transformer._interrupted + return self._interrupted async def interrupts(self) -> list[Any]: """Drive the run to completion and return interrupt payloads. @@ -401,9 +431,9 @@ class AsyncGraphRunStream: BaseException: If the run ended with an error. """ await _adrive_until_done(self._apump_next) - if (err := self._values_transformer.error) is not None: + if (err := self._mux._events._error) is not None: raise err - return self._values_transformer._interrupts + return self._interrupts def __aiter__(self) -> AsyncIterator[ProtocolEvent]: """Subscribe to the main event log and iterate protocol events.""" @@ -444,7 +474,6 @@ class SubgraphRunStream(GraphRunStream, _SubgraphRunStreamMixin): def __init__( self, mux: StreamMux, - values_transformer: ValuesTransformer, *, path: tuple[str, ...], graph_name: str | None = None, @@ -456,7 +485,6 @@ class SubgraphRunStream(GraphRunStream, _SubgraphRunStreamMixin): super().__init__( graph_iter=None, mux=mux, - values_transformer=values_transformer, wire_pump=False, ) self.path = path @@ -489,7 +517,6 @@ class AsyncSubgraphRunStream(AsyncGraphRunStream, _SubgraphRunStreamMixin): def __init__( self, mux: StreamMux, - values_transformer: ValuesTransformer, *, path: tuple[str, ...], graph_name: str | None = None, @@ -499,7 +526,6 @@ class AsyncSubgraphRunStream(AsyncGraphRunStream, _SubgraphRunStreamMixin): super().__init__( graph_aiter=None, mux=mux, - values_transformer=values_transformer, wire_pump=False, ) self.path = path diff --git a/libs/langgraph/langgraph/stream/transformers.py b/libs/langgraph/langgraph/stream/transformers.py index a72664ff1..fc2009b13 100644 --- a/libs/langgraph/langgraph/stream/transformers.py +++ b/libs/langgraph/langgraph/stream/transformers.py @@ -28,10 +28,9 @@ _logger = logging.getLogger(__name__) class ValuesTransformer(StreamTransformer): """Capture values events as a drainable stream of state snapshots. - Keeps `_latest` / `_interrupted` / `_interrupts` as scalar state - regardless of whether the log has a subscriber — so `run.output()` - and `run.interrupted` work without forcing the caller to iterate - `run.values`. Log pushes are silent no-ops when unsubscribed. + Provides the `run.values` projection. `run.output`, + `run.interrupted` and `run.interrupts` are tracked directly + by the run stream and do not depend on this transformer. Native transformer — projection keys are exposed as direct attributes on the run stream (e.g. `run.values`). @@ -39,9 +38,9 @@ class ValuesTransformer(StreamTransformer): Only values events at the run's own level are captured; snapshots from deeper subgraphs are left in the main event log but excluded from the projection. "Own level" is defined by `scope`, which - `stream_v2` / `astream_v2` populate from the caller's checkpoint - namespace so that a nested `stream_v2` call still sees its own - root snapshots. + `stream_v2` / `astream_v2` populate from the caller's + checkpoint namespace so that a nested `stream_v2` call still + sees its own root snapshots. """ _native = True @@ -549,17 +548,10 @@ class SubgraphTransformer(_TasksLifecycleBase): try: child_mux = self._mux._make_child(ns) except RuntimeError: - # Mux wasn't built from factories — no mini-mux navigation - # available. Skip; LifecycleTransformer still tracks the - # subgraph via the flat event stream. - return - values_t = child_mux.transformer_by_key("values") - if not isinstance(values_t, ValuesTransformer): return handle_cls = AsyncSubgraphRunStream if child_mux.is_async else SubgraphRunStream handle = handle_cls( mux=child_mux, - values_transformer=values_t, path=ns, graph_name=graph_name, trigger_call_id=trigger_call_id, @@ -630,7 +622,9 @@ class SubgraphTransformer(_TasksLifecycleBase): else: await handle._mux.aclose() - def _child_mux_for_event(self, event: ProtocolEvent) -> StreamMux | None: + def _handle_for_event( + self, event: ProtocolEvent + ) -> SubgraphRunStream | AsyncSubgraphRunStream | None: ns = tuple(event["params"]["namespace"]) depth = len(self.scope) if len(ns) < depth + 1: @@ -638,22 +632,21 @@ class SubgraphTransformer(_TasksLifecycleBase): handle = self._handles.get(ns[: depth + 1]) if handle is None or handle._mux is None or handle._mux._events._closed: return None - return handle._mux + return handle def process(self, event: ProtocolEvent) -> bool: - # Discover / update terminal status before forwarding so a - # `started` handle exists by the time the child mini-mux sees - # its own first event. + # Run tasks bookkeeping first so a `started` handle exists + # by the time we forward the event to the child mini-mux. keep = super().process(event) - child_mux = self._child_mux_for_event(event) - if child_mux is not None: - child_mux.push(event) + handle = self._handle_for_event(event) + if handle is not None: + handle._observe_event(event) + handle._mux.push(event) return keep async def aprocess(self, event: ProtocolEvent) -> bool: - # Async counterpart to `process`: repeat the tasks bookkeeping - # here instead of delegating to `process`, so child mini-muxes - # receive events through their async lane. + # Async counterpart: repeats the tasks bookkeeping here so + # child mini-muxes receive events through their async lane. if event["method"] == "tasks": ns = tuple(event["params"]["namespace"]) data = event["params"]["data"] @@ -665,9 +658,10 @@ class SubgraphTransformer(_TasksLifecycleBase): keep = False else: keep = True - child_mux = self._child_mux_for_event(event) - if child_mux is not None: - await child_mux.apush(event) + handle = self._handle_for_event(event) + if handle is not None: + handle._observe_event(event) + await handle._mux.apush(event) return keep def _complete_open_handles(self) -> BaseException | None: diff --git a/libs/langgraph/tests/test_pregel_stream_v2.py b/libs/langgraph/tests/test_pregel_stream_v2.py index e4a62a6f4..17d4013bc 100644 --- a/libs/langgraph/tests/test_pregel_stream_v2.py +++ b/libs/langgraph/tests/test_pregel_stream_v2.py @@ -21,6 +21,7 @@ from langgraph.stream import ( from langgraph.stream._convert import convert_to_protocol_event from langgraph.stream._mux import StreamMux from langgraph.stream._types import ProtocolEvent +from langgraph.stream.run_stream import AsyncGraphRunStream, GraphRunStream from langgraph.stream.transformers import MessagesTransformer, ValuesTransformer from langgraph.types import StreamWriter, interrupt @@ -779,6 +780,43 @@ class TestValuesTransformer: assert len(t._interrupts) == 2 +class TestOutputWithoutValuesTransformer: + """run.output / run.interrupted / run.interrupts must work even when + ValuesTransformer is not registered.""" + + def _stream_part( + self, method: str, data: Any, namespace: tuple[str, ...] = () + ) -> dict[str, Any]: + return {"type": method, "ns": namespace, "data": data} + + def test_output_without_values_transformer(self) -> None: + mux = StreamMux(factories=[MessagesTransformer], is_async=False) + run = GraphRunStream( + iter([self._stream_part("values", {"v": "final"})]), + mux, + ) + assert "values" not in run.extensions + assert run.output == {"v": "final"} + + def test_interrupts_without_values_transformer(self) -> None: + part = self._stream_part("values", {"v": 1}) + part["interrupts"] = ({"value": "pause"},) + mux = StreamMux(factories=[MessagesTransformer], is_async=False) + run = GraphRunStream(iter([part]), mux) + assert run.interrupted is True + assert len(run.interrupts) == 1 + + @pytest.mark.anyio + async def test_async_output_without_values_transformer(self) -> None: + async def _parts() -> Any: + yield {"type": "values", "ns": (), "data": {"v": "async_final"}} + + mux = StreamMux(factories=[MessagesTransformer], is_async=True) + run = AsyncGraphRunStream(_parts(), mux) + assert "values" not in run.extensions + assert await run.output() == {"v": "async_final"} + + class TestMessagesTransformer: def test_captures_root_messages(self) -> None: t = MessagesTransformer() diff --git a/libs/langgraph/tests/test_stream_messages_transformer.py b/libs/langgraph/tests/test_stream_messages_transformer.py index 0a767dcad..7f5c209b8 100644 --- a/libs/langgraph/tests/test_stream_messages_transformer.py +++ b/libs/langgraph/tests/test_stream_messages_transformer.py @@ -403,7 +403,7 @@ class TestWireRequestMore: mux = StreamMux([values_t, messages_t], is_async=False) assert messages_t._pump_fn is None - run = GraphRunStream(iter([]), mux, values_t) + run = GraphRunStream(iter([]), mux) assert messages_t._pump_fn is not None assert messages_t._pump_fn() is False assert run._exhausted @@ -412,7 +412,7 @@ class TestWireRequestMore: values_t = ValuesTransformer() messages_t = MessagesTransformer() mux = StreamMux([values_t, messages_t], is_async=False) - GraphRunStream(iter([]), mux, values_t) + GraphRunStream(iter([]), mux) log: StreamChannel[ChatModelStream] = mux.extensions["messages"] log._subscribed = True diff --git a/libs/langgraph/tests/test_stream_subgraph_transformer.py b/libs/langgraph/tests/test_stream_subgraph_transformer.py index ce46d0c3a..d96e16621 100644 --- a/libs/langgraph/tests/test_stream_subgraph_transformer.py +++ b/libs/langgraph/tests/test_stream_subgraph_transformer.py @@ -546,8 +546,6 @@ def test_child_forwarding_errors_fail_sync_run() -> None: ], is_async=False, ) - values_t = mux.transformer_by_key("values") - assert isinstance(values_t, ValuesTransformer) run = GraphRunStream( iter( [ @@ -565,7 +563,6 @@ def test_child_forwarding_errors_fail_sync_run() -> None: ] ), mux, - values_t, ) handle = next(iter(run.subgraphs)) @@ -587,8 +584,6 @@ async def test_child_forwarding_errors_fail_async_run() -> None: ], is_async=True, ) - values_t = mux.transformer_by_key("values") - assert isinstance(values_t, ValuesTransformer) run = AsyncGraphRunStream( _astream_parts( _stream_part( @@ -604,7 +599,6 @@ async def test_child_forwarding_errors_fail_async_run() -> None: _stream_part("values", ("agent:abc",), {"x": 1}), ), mux, - values_t, ) handle = await run.subgraphs.__aiter__().__anext__() @@ -625,8 +619,6 @@ def test_child_finalize_errors_propagate_to_sync_run() -> None: ], is_async=False, ) - values_t = mux.transformer_by_key("values") - assert isinstance(values_t, ValuesTransformer) run = GraphRunStream( iter( [ @@ -643,7 +635,6 @@ def test_child_finalize_errors_propagate_to_sync_run() -> None: ] ), mux, - values_t, ) with pytest.raises(RuntimeError, match="child finalize boom"): @@ -662,8 +653,6 @@ async def test_child_finalize_errors_propagate_to_async_run() -> None: ], is_async=True, ) - values_t = mux.transformer_by_key("values") - assert isinstance(values_t, ValuesTransformer) run = AsyncGraphRunStream( _astream_parts( _stream_part( @@ -678,7 +667,6 @@ async def test_child_finalize_errors_propagate_to_async_run() -> None: ) ), mux, - values_t, ) with pytest.raises(RuntimeError, match="child afinalize boom"):