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
This commit is contained in:
Nick Hollon
2026-04-16 10:24:52 -04:00
parent 28cf5ed78d
commit 986c1cc2e3
6 changed files with 189 additions and 42 deletions
+23 -9
View File
@@ -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.
+3 -1
View File
@@ -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
+26 -6
View File
@@ -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]:
@@ -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())
@@ -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)
+130 -14
View File
@@ -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