mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-13 21:27:52 +02:00
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:
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user