Skip to content
Merged
Show file tree
Hide file tree
Changes from 4 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
146 changes: 100 additions & 46 deletions src/agentex/lib/core/services/adk/streaming.py
Original file line number Diff line number Diff line change
Expand Up @@ -177,6 +177,12 @@ async def add(self, update: StreamTaskMessageDelta) -> None:
if self._closed:
return
async with self._lock:
# Re-check under the lock: a concurrent close() (e.g. from a racing
# Full) may have drained and shut down the ticker after the check
# above but before we acquired the lock. Appending now would strand
# the delta in a dead buffer, never published.
if self._closed:
return
self._buf.append(update)
self._buf_chars += _delta_char_len(update.delta)
if not self._first_flushed or self._buf_chars >= self.MAX_BUFFERED_CHARS:
Expand Down Expand Up @@ -380,6 +386,9 @@ def __init__(
self._agentex_client = agentex_client
self._streaming_service = streaming_service
self._is_closed = False
# Set once a terminal (Full/Done) starts processing, so a concurrent
# delta can't publish after the terminal during the close window.
self._is_closing = False
self._delta_accumulator = DeltaAccumulator()
self._streaming_mode: StreamingMode = streaming_mode
self._buffer: CoalescingBuffer | None = None
Expand All @@ -393,6 +402,7 @@ async def __aexit__(self, exc_type, exc_val, exc_tb):

async def open(self) -> "StreamingTaskMessageContext":
self._is_closed = False
self._is_closing = False

self.task_message = await self._agentex_client.messages.create(
task_id=self.task_id,
Expand All @@ -415,46 +425,65 @@ async def open(self) -> "StreamingTaskMessageContext":

return self

async def close(self) -> TaskMessage:
"""Close the streaming context."""
if not self.task_message:
raise ValueError("Context not properly initialized - no task message")
async def _reap_buffer(self) -> None:
"""Drain and stop the coalescing buffer, releasing its background ticker.

if self._is_closed:
return self.task_message # Already done

# Drain any buffered deltas before announcing DONE so consumers see the
# full sequence in order.
Idempotent: a no-op once the buffer has already been reaped.
"""
if self._buffer is not None:
await self._buffer.close()
Comment thread
eberki-scale marked this conversation as resolved.
Outdated
self._buffer = None

# Send the DONE event
done_event = StreamTaskMessageDone(
parent_task_message=self.task_message,
type="done",
)
await self._streaming_service.stream_update(done_event)
async def close(self) -> TaskMessage:
"""Close the streaming context."""
if not self.task_message:
raise ValueError("Context not properly initialized - no task message")

# Update the task message with the final content
has_deltas = (
self._delta_accumulator._accumulated_deltas
or self._delta_accumulator._reasoning_summaries
or self._delta_accumulator._reasoning_contents
)
if has_deltas:
self.task_message.content = self._delta_accumulator.convert_to_content()
# Mark closing before reaping so a concurrent coalesced delta can't slip
# past the now-None buffer and publish after DONE. Rolled back below if
# the terminal write fails, so a failed close doesn't wedge the context
# into silently dropping later deltas.
self._is_closing = True
try:
# Reap the buffer (stopping its ticker) before the _is_closed
# short-circuit, so a context already marked done by a Full update
# can't leave the ticker orphaned. Draining here also lets consumers
# see the full delta sequence in order before DONE.
await self._reap_buffer()

if self._is_closed:
return self.task_message # Already done (buffer reaped above)

# Send the DONE event
done_event = StreamTaskMessageDone(
parent_task_message=self.task_message,
type="done",
)
await self._streaming_service.stream_update(done_event)

await self._agentex_client.messages.update(
task_id=self.task_id,
message_id=self.task_message.id,
content=self.task_message.content.model_dump(),
streaming_status="DONE",
)
# Update the task message with the final content
has_deltas = (
self._delta_accumulator._accumulated_deltas
or self._delta_accumulator._reasoning_summaries
or self._delta_accumulator._reasoning_contents
)
if has_deltas:
self.task_message.content = self._delta_accumulator.convert_to_content()

# Mark the context as done
self._is_closed = True
return self.task_message
await self._agentex_client.messages.update(
task_id=self.task_id,
message_id=self.task_message.id,
content=self.task_message.content.model_dump(),
streaming_status="DONE",
)

# Mark the context as done
self._is_closed = True
return self.task_message
except Exception:
if not self._is_closed:
self._is_closing = False
raise

async def stream_update(self, update: TaskMessageUpdate) -> TaskMessageUpdate | None:
"""Stream an update to the repository.
Expand Down Expand Up @@ -485,21 +514,46 @@ async def stream_update(self, update: TaskMessageUpdate) -> TaskMessageUpdate |
if self._streaming_mode == "coalesced" and self._buffer is not None:
await self._buffer.add(update)
return update
if self._is_closing:
# A terminal (Full/Done) is in flight and the buffer is already
# reaped; drop this late delta instead of publishing it after the
# terminal event. It is still recorded in the accumulator above.
return update

result = await self._streaming_service.stream_update(update)

if isinstance(update, StreamTaskMessageDone):
await self.close()
return update
elif 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]
content=update.content.model_dump(),
streaming_status="DONE",
)
self._is_closed = True
return result
# From here the update is terminal. Mark closing so a concurrent delta
# can't fall through and publish after the terminal once the buffer is
# reaped below. Rolled back if the terminal path fails, so a failed
# terminal doesn't wedge the context into silently dropping later deltas.
is_terminal = isinstance(update, (StreamTaskMessageFull, StreamTaskMessageDone))
if is_terminal:
self._is_closing = True
try:
# 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 duplicate tail. This also stops the ticker, which would
# otherwise be orphaned when __aexit__'s close() short-circuits.
if isinstance(update, StreamTaskMessageFull):
await self._reap_buffer()

result = await self._streaming_service.stream_update(update)

if isinstance(update, StreamTaskMessageDone):
await self.close()
Comment thread
greptile-apps[bot] marked this conversation as resolved.
Outdated
return update
elif 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]
content=update.content.model_dump(),
streaming_status="DONE",
)
self._is_closed = True
return result
except Exception:
if not self._is_closed:
self._is_closing = False
Comment thread
greptile-apps[bot] marked this conversation as resolved.
Outdated
raise


class StreamingService:
Expand Down
140 changes: 139 additions & 1 deletion tests/lib/core/services/adk/test_streaming.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,10 @@
ToolResponseDelta,
ReasoningSummaryDelta,
)
from agentex.types.task_message_update import StreamTaskMessageDelta
from agentex.types.task_message_update import (
StreamTaskMessageFull,
StreamTaskMessageDelta,
)
from agentex.lib.core.services.adk.streaming import (
CoalescingBuffer,
StreamingTaskMessageContext,
Expand Down Expand Up @@ -352,6 +355,24 @@ async def on_flush(u: StreamTaskMessageDelta) -> None:
await buf.add(_text(task_message, "after"))
assert flushed == []

@pytest.mark.asyncio
async def test_add_racing_close_is_not_stranded(self, task_message: TaskMessage) -> None:
"""TOCTOU: a delta that passes add()'s pre-lock _closed check but only
acquires the lock after close() set _closed must be dropped, not appended
to a drained, ticker-less buffer where it would never be published."""
buf = CoalescingBuffer(on_flush=AsyncMock())
buf.start()
# Hold the lock so add() parks *after* its pre-lock _closed check.
await buf._lock.acquire()
add_task = asyncio.create_task(buf.add(_text(task_message, "racing")))
await asyncio.sleep(0) # add() passes the _closed check, blocks on the lock
buf._closed = True # close() wins the race
buf._lock.release()
await add_task

assert buf._buf == [], "racing delta was stranded in the closed buffer"
await buf.close() # cleanup


class TestCoalescingBufferCloseDuringFlush:
@pytest.mark.asyncio
Expand Down Expand Up @@ -520,3 +541,120 @@ async def test_open_without_created_at_passes_omit(self) -> None:

kwargs = client.messages.create.call_args.kwargs
assert kwargs["created_at"] is omit


class TestFullMessageClosesBuffer:
"""A StreamTaskMessageFull must stop the buffer ticker and drain its deltas
before the terminal Full. Marking the context done without closing the
buffer leaves close()'s _is_closed short-circuit to orphan the ticker, and
publishing buffered deltas after the Full reads as a stale duplicate tail."""

@pytest.mark.asyncio
async def test_full_message_stops_ticker(self) -> None:
ctx, _svc, tm = await _make_context("coalesced")
# A delta makes the buffer and its ticker live.
await ctx.stream_update(_text(tm, "hello"))
buf = ctx._buffer
assert buf is not None
task = buf._task
assert task is not None and not task.done()

await ctx.stream_update(
StreamTaskMessageFull(
parent_task_message=tm,
content=TextContent(author="agent", content="final", format="markdown"),
type="full",
)
)

assert ctx._buffer is None, "Full message left the buffer un-closed"
assert task.done(), "coalescing-buffer ticker still running after Full (orphaned)"

@pytest.mark.asyncio
async def test_full_is_terminal_publish_no_trailing_deltas(self) -> None:
# Buffered deltas must publish BEFORE the Full, never after (a trailing
# delta after the terminal Full reads as a stale duplicate tail).
ctx, svc, tm = await _make_context("coalesced")
# Two deltas through the buffer. Regardless of how the coalescing window
# batches them (1 or 2 publishes), the invariant under test is the same:
# every delta publishes before the terminal Full, never after it.
await ctx.stream_update(_text(tm, "alpha"))
await ctx.stream_update(_text(tm, "beta"))

full = StreamTaskMessageFull(
parent_task_message=tm,
content=TextContent(author="agent", content="alphabeta", format="markdown"),
type="full",
)
await ctx.stream_update(full)

published = [c.args[0] for c in svc.stream_update.await_args_list]
assert published, "nothing was published"
assert published[-1] is full, (
f"Full must be the terminal publish; saw trailing "
f"{type(published[-1]).__name__} after it (stale duplicate tail)"
)
assert any(isinstance(u, StreamTaskMessageDelta) for u in published[:-1]), (
"expected the buffered deltas to be published before the Full"
)

@pytest.mark.asyncio
async def test_late_delta_during_close_window_is_not_published_after_full(self) -> None:
"""Once a Full reaps the buffer, a delta racing in before _is_closed is
set (buffer already None) must be dropped, not published after the Full.
Simulates the concurrent window via the _is_closing flag directly."""
ctx, svc, tm = await _make_context("coalesced")
await ctx.stream_update(_text(tm, "hello"))

full = StreamTaskMessageFull(
parent_task_message=tm,
content=TextContent(author="agent", content="hello", format="markdown"),
type="full",
)
await ctx.stream_update(full)
# Full set _is_closed; reopen the close window to mimic a delta arriving
# after the buffer was reaped but before the terminal finished.
ctx._is_closed = False
ctx._is_closing = True
svc.stream_update.reset_mock()

await ctx.stream_update(_text(tm, "late"))

assert svc.stream_update.await_count == 0, "late delta published after the terminal Full"

@pytest.mark.asyncio
async def test_close_marks_closing_before_reaping_buffer(self) -> None:
"""close() must set _is_closing before reaping, so a concurrent delta in
the post-reap / pre-_is_closed window is dropped, not published after DONE."""
ctx, _svc, tm = await _make_context("coalesced")
await ctx.stream_update(_text(tm, "hi"))

closing_at_reap = {}
orig_reap = ctx._reap_buffer

async def spy() -> None:
closing_at_reap["value"] = ctx._is_closing
await orig_reap()

ctx._reap_buffer = spy # type: ignore[method-assign]
await ctx.close()
assert closing_at_reap.get("value") is True, "close() must mark _is_closing before reaping"

@pytest.mark.asyncio
async def test_failed_terminal_rolls_back_is_closing(self) -> None:
"""If the terminal persistence fails, _is_closing must roll back so the
still-open context doesn't silently drop later coalesced deltas."""
ctx, _svc, tm = await _make_context("coalesced")
await ctx.stream_update(_text(tm, "hi"))
ctx._agentex_client.messages.update.side_effect = RuntimeError("persist boom") # type: ignore[attr-defined]

full = StreamTaskMessageFull(
parent_task_message=tm,
content=TextContent(author="agent", content="hi", format="markdown"),
type="full",
)
with pytest.raises(RuntimeError):
await ctx.stream_update(full)

assert ctx._is_closed is False
assert ctx._is_closing is False, "failed terminal left _is_closing stuck on"
Loading