From 986c1cc2e387987280ad5dd63c36796ddabb9a51 Mon Sep 17 00:00:00 2001 From: Nick Hollon Date: Thu, 16 Apr 2026 10:24:52 -0400 Subject: [PATCH] Auto-close EventLogs, reject projection key conflicts, fix async interrupted/interrupts Three usability fixes: - Mux now auto-closes/fails EventLogs in projections (like StreamChannels), so transformers no longer need finalize/fail boilerplate - StreamingHandler._setup() raises ValueError if a user transformer returns projection keys that collide with already-registered keys - AsyncGraphRunStream.interrupted and .interrupts now await the pump task before returning, matching the output property's behavior --- libs/langgraph/langgraph/stream/_mux.py | 32 ++-- libs/langgraph/langgraph/stream/_types.py | 4 +- libs/langgraph/langgraph/stream/run_stream.py | 32 +++- .../langgraph/stream/streaming_handler.py | 7 + .../langgraph/stream/transformers.py | 12 -- .../langgraph/tests/test_streaming_handler.py | 144 ++++++++++++++++-- 6 files changed, 189 insertions(+), 42 deletions(-) diff --git a/libs/langgraph/langgraph/stream/_mux.py b/libs/langgraph/langgraph/stream/_mux.py index 54666924c..01aa6df8e 100644 --- a/libs/langgraph/langgraph/stream/_mux.py +++ b/libs/langgraph/langgraph/stream/_mux.py @@ -29,6 +29,7 @@ class StreamMux: self._events._bind(is_async=is_async) self._transformers: list[StreamTransformer] = [] self._channels: list[StreamChannel[Any]] = [] + self._logs: list[EventLog[Any]] = [] self._seq = 0 def register(self, transformer: StreamTransformer) -> dict[str, Any]: @@ -72,11 +73,13 @@ class StreamMux: self._events.push(event) def close(self) -> None: - """Finalize all transformers, close all channels and the main log. + """Finalize all transformers, close all projections and the main log. - If any transformer's ``finalize()`` raises, the remaining - transformers, channels, and the main log are still closed. - The first error is re-raised after cleanup completes. + EventLogs and StreamChannels discovered in transformer projections + are auto-closed after ``finalize()`` runs — transformers don't need + to close them manually. If any transformer's ``finalize()`` raises, + the remaining transformers, projections, and the main log are still + closed. The first error is re-raised after cleanup completes. """ first_error: BaseException | None = None for transformer in self._transformers: @@ -85,25 +88,35 @@ class StreamMux: except BaseException as e: if first_error is None: first_error = e + for log in self._logs: + if not log._closed: + log.close() for ch in self._channels: - ch._close() + if not ch._log._closed: + ch._close() self._events.close() if first_error is not None: raise first_error def fail(self, err: BaseException) -> None: - """Fail all transformers, channels, and the main log. + """Fail all transformers, projections, and the main log. - If any transformer's ``fail()`` raises, the remaining - transformers, channels, and the main log are still failed. + EventLogs and StreamChannels discovered in transformer projections + are auto-failed — transformers don't need to fail them manually. + If any transformer's ``fail()`` raises, the remaining transformers, + projections, and the main log are still failed. """ for transformer in self._transformers: try: transformer.fail(err) except BaseException: pass + for log in self._logs: + if not log._closed: + log.fail(err) for ch in self._channels: - ch._fail(err) + if not ch._log._closed: + ch._fail(err) self._events.fail(err) # ------------------------------------------------------------------ @@ -127,6 +140,7 @@ class StreamMux: value._wire(_make_forward(channel_name)) elif isinstance(value, EventLog): value._bind(is_async=self._is_async) + self._logs.append(value) def _forward(self, channel_name: str, item: Any) -> None: """Inject a ProtocolEvent for a StreamChannel push. diff --git a/libs/langgraph/langgraph/stream/_types.py b/libs/langgraph/langgraph/stream/_types.py index 3aec9c12f..513044916 100644 --- a/libs/langgraph/langgraph/stream/_types.py +++ b/libs/langgraph/langgraph/stream/_types.py @@ -41,7 +41,9 @@ class StreamTransformer(ABC): Subclasses must implement `init` and `process`. The `finalize` and `fail` hooks are optional — the default implementations are no-ops. - StreamChannel instances are auto-closed/failed by the mux regardless. + EventLog and StreamChannel instances in the projection dict are + auto-closed/failed by the mux, so most transformers don't need + ``finalize`` or ``fail`` at all. """ @abstractmethod diff --git a/libs/langgraph/langgraph/stream/run_stream.py b/libs/langgraph/langgraph/stream/run_stream.py index 4b15e545a..8788f02d4 100644 --- a/libs/langgraph/langgraph/stream/run_stream.py +++ b/libs/langgraph/langgraph/stream/run_stream.py @@ -151,17 +151,37 @@ class AsyncGraphRunStream: return self._values_transformer._latest @property - def interrupted(self) -> bool: - """Whether the run was interrupted. + def interrupted(self) -> Any: + """Return an awaitable that resolves to whether the run was interrupted. - Only meaningful after the run has completed (after consuming the - stream or awaiting ``output``). + Usage:: + + interrupted = await run.interrupted """ + return self._get_interrupted() + + async def _get_interrupted(self) -> bool: + try: + await self._pump_task + except BaseException: + pass return self._values_transformer._interrupted @property - def interrupts(self) -> list[Any]: - """Interrupt payloads, populated when interrupted is True.""" + def interrupts(self) -> Any: + """Return an awaitable that resolves to interrupt payloads. + + Usage:: + + interrupts = await run.interrupts + """ + return self._get_interrupts() + + async def _get_interrupts(self) -> list[Any]: + try: + await self._pump_task + except BaseException: + pass return self._values_transformer._interrupts def __aiter__(self) -> AsyncIterator[ProtocolEvent]: diff --git a/libs/langgraph/langgraph/stream/streaming_handler.py b/libs/langgraph/langgraph/stream/streaming_handler.py index 4a901bdd3..2a14a17cf 100644 --- a/libs/langgraph/langgraph/stream/streaming_handler.py +++ b/libs/langgraph/langgraph/stream/streaming_handler.py @@ -154,6 +154,13 @@ class StreamingHandler: for t in all_transformers: projection = mux.register(t) + conflicts = set(projection) & set(extensions) + if conflicts: + name = type(t).__name__ + raise ValueError( + f"Transformer {name} returned projection keys that " + f"conflict with already-registered keys: {conflicts}" + ) extensions.update(projection) if getattr(t, "_native", False): native_keys.update(projection.keys()) diff --git a/libs/langgraph/langgraph/stream/transformers.py b/libs/langgraph/langgraph/stream/transformers.py index 9c7bd8444..59a470230 100644 --- a/libs/langgraph/langgraph/stream/transformers.py +++ b/libs/langgraph/langgraph/stream/transformers.py @@ -39,12 +39,6 @@ class ValuesTransformer(StreamTransformer): self._interrupts.extend(interrupts) return True - def finalize(self) -> None: - self._log.close() - - def fail(self, err: BaseException) -> None: - self._log.fail(err) - class MessagesTransformer(StreamTransformer): """Captures messages events and passes through raw (chunk, metadata) tuples. @@ -74,9 +68,3 @@ class MessagesTransformer(StreamTransformer): return True self._log.push(params["data"]) return True - - def finalize(self) -> None: - self._log.close() - - def fail(self, err: BaseException) -> None: - self._log.fail(err) diff --git a/libs/langgraph/tests/test_streaming_handler.py b/libs/langgraph/tests/test_streaming_handler.py index df6a1410e..d051f1810 100644 --- a/libs/langgraph/tests/test_streaming_handler.py +++ b/libs/langgraph/tests/test_streaming_handler.py @@ -560,8 +560,8 @@ class TestStreamingHandlerAsyncInterrupt: {"configurable": {"thread_id": "t2"}}, ) _ = await run.output - assert run.interrupted is True - assert len(run.interrupts) > 0 + assert await run.interrupted is True + assert len(await run.interrupts) > 0 class TestStreamingHandlerAsyncCustom: @@ -729,7 +729,7 @@ class TestValuesTransformer: t.process(_event("values", {"val": "root"})) t.process(_event("values", {"val": "sub"}, namespace=["sub"])) - t.finalize() + t._log.close() items = list(t._log) assert len(items) == 1 assert items[0]["val"] == "root" @@ -742,7 +742,7 @@ class TestValuesTransformer: result = t.process(_event("updates", {"x": 1})) assert result is True # passed through - t.finalize() + t._log.close() assert list(t._log) == [] # but not captured def test_tracks_interrupts(self) -> None: @@ -768,7 +768,7 @@ class TestMessagesTransformer: t._log._bind(is_async=False) t.process(_event("messages", ("chunk", {"meta": True}))) - t.finalize() + t._log.close() items = list(t._log) assert len(items) == 1 assert items[0] == ("chunk", {"meta": True}) @@ -779,7 +779,7 @@ class TestMessagesTransformer: t._log._bind(is_async=False) t.process(_event("messages", ("chunk", {}), namespace=["sub"])) - t.finalize() + t._log.close() assert list(t._log) == [] def test_ignores_non_messages_methods(self) -> None: @@ -789,14 +789,14 @@ class TestMessagesTransformer: result = t.process(_event("values", {"v": 1})) assert result is True - t.finalize() + t._log.close() assert list(t._log) == [] def test_fail_propagates(self) -> None: t = MessagesTransformer() t.init() t._log._bind(is_async=False) - t.fail(ValueError("msg error")) + t._log.fail(ValueError("msg error")) with pytest.raises(ValueError, match="msg error"): list(t._log) @@ -962,12 +962,6 @@ class TestCustomTransformer: self._log.push("saw_values") return True - def finalize(self) -> None: - self._log.close() - - def fail(self, err: BaseException) -> None: - self._log.fail(err) - graph = _build_simple_graph() handler = StreamingHandler(graph) foo_t = FooTransformer() @@ -1040,3 +1034,125 @@ class TestCustomTransformer: # Seq numbers must be strictly increasing. for i in range(1, len(seqs)): assert seqs[i] > seqs[i - 1], f"Seq out of order at index {i}: {seqs}" + + def test_projection_key_conflict_raises(self) -> None: + """User transformer that collides with a built-in key should raise.""" + + class ConflictTransformer(StreamTransformer): + def __init__(self) -> None: + super().__init__() + self._log: EventLog[str] = EventLog() + + def init(self) -> dict[str, Any]: + return {"values": self._log} + + def process(self, event: ProtocolEvent) -> bool: + return True + + graph = _build_simple_graph() + handler = StreamingHandler(graph) + with pytest.raises(ValueError, match="conflict.*{'values'}"): + handler.stream( + {"value": "x", "items": []}, + transformers=[ConflictTransformer()], + ) + + +class TestEventLogAutoLifecycle: + def test_mux_auto_closes_event_logs(self) -> None: + """EventLogs in projections should be auto-closed by mux.close().""" + + class SimpleTransformer(StreamTransformer): + def __init__(self) -> None: + super().__init__() + self._log: EventLog[str] = EventLog() + + def init(self) -> dict[str, Any]: + return {"items": self._log} + + def process(self, event: ProtocolEvent) -> bool: + self._log.push("saw_event") + return True + + mux = StreamMux() + mux.register(SimpleTransformer()) + + mux.push(_event("values")) + mux.close() + + # The log should have been auto-closed — iteration should work. + items = list(mux._events) + assert len(items) == 1 + + def test_mux_auto_fails_event_logs(self) -> None: + """EventLogs in projections should be auto-failed by mux.fail().""" + + class SimpleTransformer(StreamTransformer): + def __init__(self) -> None: + super().__init__() + self._log: EventLog[str] = EventLog() + + def init(self) -> dict[str, Any]: + return {"items": self._log} + + def process(self, event: ProtocolEvent) -> bool: + self._log.push("saw_event") + return True + + t = SimpleTransformer() + mux = StreamMux() + mux.register(t) + + mux.push(_event("values")) + mux.fail(ValueError("boom")) + + # The transformer's log should have been auto-failed. + with pytest.raises(ValueError, match="boom"): + list(t._log) + + def test_no_double_close_if_transformer_closes_own_log(self) -> None: + """If a transformer closes its log in finalize(), mux should not error.""" + + class ManualCloseTransformer(StreamTransformer): + def __init__(self) -> None: + super().__init__() + self._log: EventLog[str] = EventLog() + + def init(self) -> dict[str, Any]: + return {"items": self._log} + + def process(self, event: ProtocolEvent) -> bool: + return True + + def finalize(self) -> None: + self._log.close() + + mux = StreamMux() + mux.register(ManualCloseTransformer()) + # Should not raise even though the log is closed by both + # the transformer and the mux. + mux.close() + + def test_transformer_without_finalize_works(self) -> None: + """Transformer with only init+process should work end-to-end.""" + + class MinimalTransformer(StreamTransformer): + def __init__(self) -> None: + super().__init__() + self._log: EventLog[str] = EventLog() + + def init(self) -> dict[str, Any]: + return {"minimal": self._log} + + def process(self, event: ProtocolEvent) -> bool: + if event["method"] == "values": + self._log.push("got_it") + return True + + graph = _build_simple_graph() + handler = StreamingHandler(graph) + t = MinimalTransformer() + run = handler.stream({"value": "x", "items": []}, transformers=[t]) + _ = run.output + items = list(run.extensions["minimal"]) + assert len(items) > 0