diff --git a/src/agentex/lib/core/services/adk/streaming.py b/src/agentex/lib/core/services/adk/streaming.py index 33ca7bc1c..921897669 100644 --- a/src/agentex/lib/core/services/adk/streaming.py +++ b/src/agentex/lib/core/services/adk/streaming.py @@ -395,7 +395,14 @@ async def __aenter__(self) -> "StreamingTaskMessageContext": return await self.open() async def __aexit__(self, exc_type, exc_val, exc_tb): - return await self.close() + """Close the context, then let any exception from the body propagate. + + Returning ``close()``'s truthy ``TaskMessage`` here would suppress it, + silently persisting a half-written message as DONE and hiding the + failure from the caller (and from Temporal's activity retries). + """ + await self.close() + return False async def open(self) -> "StreamingTaskMessageContext": self._is_closed = False @@ -432,6 +439,10 @@ async def _reap_buffer(self) -> None: async def close(self) -> TaskMessage: """Close the streaming context.""" + return await self._finish() + + async def _finish(self, done_index: int | None = None) -> TaskMessage: + """Drain the buffer, publish one DONE carrying ``done_index``, and persist.""" if not self.task_message: raise ValueError("Context not properly initialized - no task message") @@ -448,6 +459,7 @@ async def close(self) -> TaskMessage: done_event = StreamTaskMessageDone( parent_task_message=self.task_message, type="done", + index=done_index, ) await self._streaming_service.stream_update(done_event) @@ -484,7 +496,9 @@ async def stream_update(self, update: TaskMessageUpdate) -> TaskMessageUpdate | ``StreamTaskMessageDone`` and ``StreamTaskMessageFull`` updates always publish synchronously regardless of mode so consumers and persistence - stay in sync. + stay in sync. A Done delegates its publish to ``close()``, which drains + the buffer first and emits exactly one DONE; publishing it here as well + would put a second DONE on the stream, ahead of the buffered deltas. """ if self._is_closed: raise ValueError("Context is already done") @@ -501,6 +515,10 @@ async def stream_update(self, update: TaskMessageUpdate) -> TaskMessageUpdate | await self._buffer.add(update) return update + if isinstance(update, StreamTaskMessageDone): + await self._finish(update.index) + return update + # A Full ends the stream and supersedes buffered deltas. Drain and stop # the buffer BEFORE publishing the Full, so leftover deltas land in order # (deltas -> Full) instead of trailing the terminal Full as a stale @@ -511,10 +529,7 @@ async def stream_update(self, update: TaskMessageUpdate) -> TaskMessageUpdate | result = await self._streaming_service.stream_update(update) - if isinstance(update, StreamTaskMessageDone): - await self.close() - return update - elif isinstance(update, StreamTaskMessageFull): + if isinstance(update, StreamTaskMessageFull): await self._agentex_client.messages.update( task_id=self.task_id, message_id=update.parent_task_message.id, # type: ignore[union-attr] diff --git a/tests/lib/core/services/adk/test_streaming.py b/tests/lib/core/services/adk/test_streaming.py index a8068f307..96dc05ed6 100644 --- a/tests/lib/core/services/adk/test_streaming.py +++ b/tests/lib/core/services/adk/test_streaming.py @@ -23,6 +23,7 @@ ReasoningSummaryDelta, ) from agentex.types.task_message_update import ( + StreamTaskMessageDone, StreamTaskMessageFull, StreamTaskMessageDelta, ) @@ -62,7 +63,8 @@ def _reasoning_summary(tm: TaskMessage, idx: int, s: str) -> StreamTaskMessageDe ) -async def _make_context(streaming_mode: str) -> tuple[StreamingTaskMessageContext, MagicMock, TaskMessage]: +def _make_unopened_context(streaming_mode: str) -> tuple[StreamingTaskMessageContext, MagicMock, TaskMessage]: + """Wired-up context that has not been opened yet, for ``async with`` tests.""" tm = TaskMessage( id="m1", task_id="t1", @@ -81,6 +83,11 @@ async def _make_context(streaming_mode: str) -> tuple[StreamingTaskMessageContex streaming_service=svc, streaming_mode=streaming_mode, # type: ignore[arg-type] ) + return ctx, svc, tm + + +async def _make_context(streaming_mode: str) -> tuple[StreamingTaskMessageContext, MagicMock, TaskMessage]: + ctx, svc, tm = _make_unopened_context(streaming_mode) await ctx.open() return ctx, svc, tm @@ -597,3 +604,71 @@ async def test_full_is_terminal_publish_no_trailing_deltas(self) -> None: assert any(isinstance(u, StreamTaskMessageDelta) for u in published[:-1]), ( "expected the buffered deltas to be published before the Full" ) + + +class TestContextDoesNotSuppressErrors: + """``__aexit__`` returned ``close()``'s TaskMessage, which is truthy, so an + exception raised inside ``async with`` was swallowed: the caller carried on + past the block, the half-written message was persisted DONE, and a Temporal + activity saw a success and never retried.""" + + @pytest.mark.asyncio + async def test_exception_inside_context_propagates(self) -> None: + ctx, _svc, tm = _make_unopened_context("off") + + with pytest.raises(RuntimeError, match="model blew up"): + async with ctx as entered: + await entered.stream_update(_text(tm, "partial")) + raise RuntimeError("model blew up") + + @pytest.mark.asyncio + async def test_context_is_still_closed_when_body_raises(self) -> None: + """Not suppressing must not mean leaking: DONE is still published and + persisted so consumers and the buffer ticker are not left hanging.""" + ctx, svc, _tm = _make_unopened_context("off") + + with pytest.raises(RuntimeError): + async with ctx: + raise RuntimeError("model blew up") + + assert ctx._is_closed + published = [c.args[0] for c in svc.stream_update.await_args_list] + assert isinstance(published[-1], StreamTaskMessageDone) + update_kwargs = ctx._agentex_client.messages.update.call_args.kwargs + assert update_kwargs["streaming_status"] == "DONE" + + +class TestExplicitDonePublishesOnce: + """An explicit ``StreamTaskMessageDone`` used to be published by + ``stream_update`` and then published a second time by the ``close()`` it + triggers, putting two DONE frames on the stream. Routing the terminal + publish through ``close()`` also keeps buffered deltas ahead of the DONE.""" + + @pytest.mark.asyncio + async def test_explicit_done_publishes_exactly_one_done(self) -> None: + ctx, svc, tm = await _make_context("coalesced") + await ctx.stream_update(_text(tm, "hello")) + + await ctx.stream_update(StreamTaskMessageDone(parent_task_message=tm, type="done")) + + published = [c.args[0] for c in svc.stream_update.await_args_list] + dones = [u for u in published if isinstance(u, StreamTaskMessageDone)] + assert len(dones) == 1, f"expected exactly one DONE publish, got {len(dones)}" + assert published[-1] is dones[0], ( + f"DONE must be the terminal publish; saw trailing {type(published[-1]).__name__} after it" + ) + assert ctx._agentex_client.messages.update.call_count == 1 + update_kwargs = ctx._agentex_client.messages.update.call_args.kwargs + assert update_kwargs["content"]["content"] == "hello" + assert update_kwargs["streaming_status"] == "DONE" + + @pytest.mark.asyncio + async def test_explicit_done_keeps_the_callers_index(self) -> None: + ctx, svc, tm = await _make_context("coalesced") + await ctx.stream_update(_text(tm, "hello")) + + await ctx.stream_update(StreamTaskMessageDone(parent_task_message=tm, type="done", index=3)) + + published = [c.args[0] for c in svc.stream_update.await_args_list] + dones = [u for u in published if isinstance(u, StreamTaskMessageDone)] + assert [d.index for d in dones] == [3]