Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 21 additions & 6 deletions src/agentex/lib/core/services/adk/streaming.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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")

Expand All @@ -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)

Expand Down Expand Up @@ -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")
Expand All @@ -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
Comment thread
greptile-apps[bot] marked this conversation as resolved.

# 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
Expand All @@ -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]
Expand Down
77 changes: 76 additions & 1 deletion tests/lib/core/services/adk/test_streaming.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
ReasoningSummaryDelta,
)
from agentex.types.task_message_update import (
StreamTaskMessageDone,
StreamTaskMessageFull,
StreamTaskMessageDelta,
)
Expand Down Expand Up @@ -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",
Expand All @@ -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

Expand Down Expand Up @@ -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]