diff --git a/src/agentex/lib/core/harness/auto_send.py b/src/agentex/lib/core/harness/auto_send.py index b645a4aae..5a0f7d7b8 100644 --- a/src/agentex/lib/core/harness/auto_send.py +++ b/src/agentex/lib/core/harness/auto_send.py @@ -2,6 +2,7 @@ from __future__ import annotations +import contextlib from typing import Any, AsyncIterator from datetime import datetime @@ -46,7 +47,13 @@ async def auto_send( Index-keyed routing: each Start(index=i) opens a context stored in ctx_map[i]; Delta(index=i) routes to ctx_map.get(i); Done(index=i) closes and removes ctx_map[i]. Events with index is None are skipped. The finally - block closes all remaining open contexts. + block closes all remaining open contexts, then closes `events` itself (when + it exposes `aclose`) so delivery stopping early — cancelled while awaiting + the backend rather than the source — tears the tap down instead of leaving it + suspended at a yield: the turn object pins its event generator, so GC does + not rescue it and the harness subprocess leaks. Closing an exhausted + generator is a no-op, and a failure there is suppressed so it cannot mask the + original exception. final_text last-segment semantics: a new Start(TextContent) resets final_text_parts so that multi-step turns return the LAST text segment. @@ -148,9 +155,15 @@ async def _close_all() -> None: pass finally: - await _close_all() - if deriver is not None and tracer is not None: - for signal in deriver.flush(): - await tracer.handle(signal) + try: + await _close_all() + if deriver is not None and tracer is not None: + for signal in deriver.flush(): + await tracer.handle(signal) + finally: + aclose = getattr(events, "aclose", None) + if aclose is not None: + with contextlib.suppress(Exception): + await aclose() return TurnResult(final_text="".join(final_text_parts), usage=usage or TurnUsage()) diff --git a/src/agentex/lib/core/harness/emitter.py b/src/agentex/lib/core/harness/emitter.py index 5b56793bf..f39b5f575 100644 --- a/src/agentex/lib/core/harness/emitter.py +++ b/src/agentex/lib/core/harness/emitter.py @@ -53,9 +53,19 @@ def __init__( self.tracer = None async def yield_turn(self, turn: HarnessTurn) -> AsyncGenerator[StreamTaskMessage, None]: - """Sync HTTP ACP delivery: forward events, trace as side effect.""" - async for event in yield_events(turn.events, tracer=self.tracer): - yield event + """Sync HTTP ACP delivery: forward events, trace as side effect. + + The finally closes the delivery generator, which closes the turn's event + source in turn, so a consumer that stops early (client disconnect) tears + the tap down here rather than leaving it to async-generator finalization, + which never runs promptly while the turn still pins its generator. + """ + delivery = yield_events(turn.events, tracer=self.tracer) + try: + async for event in delivery: + yield event + finally: + await delivery.aclose() async def auto_send_turn(self, turn: HarnessTurn, created_at: datetime | None = None) -> TurnResult: """Async/temporal delivery: push to the task stream, return TurnResult. diff --git a/src/agentex/lib/core/harness/yield_delivery.py b/src/agentex/lib/core/harness/yield_delivery.py index 69b39f152..c06f67961 100644 --- a/src/agentex/lib/core/harness/yield_delivery.py +++ b/src/agentex/lib/core/harness/yield_delivery.py @@ -2,6 +2,7 @@ from __future__ import annotations +import contextlib from typing import AsyncIterator, AsyncGenerator from agentex.lib.core.harness.types import StreamTaskMessage @@ -17,6 +18,13 @@ async def yield_events( For sync HTTP ACP agents that yield events back over the response. When `tracer` is None, this is a pure passthrough. + + The finally also closes `events` (when it exposes `aclose`), so a consumer + that stops early — a client disconnect closes this generator — tears the tap + down instead of leaving it suspended at a yield: the turn object pins its + event generator, so GC does not rescue it and the harness subprocess leaks. + Closing an exhausted generator is a no-op, and a failure there is suppressed + so it cannot mask the original exception. """ deriver = SpanDeriver() if tracer is not None else None try: @@ -26,6 +34,12 @@ async def yield_events( await tracer.handle(signal) yield event finally: - if deriver is not None and tracer is not None: - for signal in deriver.flush(): - await tracer.handle(signal) + try: + if deriver is not None and tracer is not None: + for signal in deriver.flush(): + await tracer.handle(signal) + finally: + aclose = getattr(events, "aclose", None) + if aclose is not None: + with contextlib.suppress(Exception): + await aclose() diff --git a/src/agentex/lib/sdk/fastacp/base/base_acp_server.py b/src/agentex/lib/sdk/fastacp/base/base_acp_server.py index 50c304c92..3755a0bd7 100644 --- a/src/agentex/lib/sdk/fastacp/base/base_acp_server.py +++ b/src/agentex/lib/sdk/fastacp/base/base_acp_server.py @@ -3,6 +3,7 @@ import uuid import asyncio import inspect +import contextlib from typing import Any from datetime import datetime from contextlib import asynccontextmanager @@ -434,6 +435,11 @@ async def generate_json_rpc_stream(): error=JSONRPCError(code=-32603, message=str(e)).model_dump(), ) yield f"{error_response.model_dump_json()}\n" + finally: + aclose = getattr(async_gen, "aclose", None) + if aclose is not None: + with contextlib.suppress(Exception): + await aclose() return StreamingResponse( generate_json_rpc_stream(), diff --git a/tests/lib/core/harness/test_auto_send.py b/tests/lib/core/harness/test_auto_send.py index 8133a488c..a949966c6 100644 --- a/tests/lib/core/harness/test_auto_send.py +++ b/tests/lib/core/harness/test_auto_send.py @@ -9,6 +9,8 @@ This mirrors _langgraph_async.py lines 62-78 and 100-127. """ +import asyncio +from typing import override from datetime import datetime import pytest @@ -478,3 +480,116 @@ async def test_auto_send_created_at_forwarded(): await auto_send(_gen(events), task_id="task1", tracer=None, streaming=streaming, created_at=dt) assert all(ts == dt for ts in streaming.recorded_created_at) + + +class _BlockingCtx(_FakeCtx): + """A context whose stream_update never returns (a stalled backend). + + Sets `blocked` once stream_update is awaited so the test can cancel exactly + while delivery is suspended on the backend rather than on the source. + """ + + def __init__(self, sink, content_type, initial_content, blocked): + super().__init__(sink, content_type, initial_content) + self.blocked = blocked + + @override + async def stream_update(self, update): + self.sink.append(("update", update)) + self.blocked.set() + await asyncio.Event().wait() + + +class _BlockingStreaming(_FakeStreaming): + """_FakeStreaming whose contexts block forever inside stream_update.""" + + def __init__(self): + super().__init__() + self.blocked = asyncio.Event() + + @override + def streaming_task_message_context(self, task_id, initial_content, streaming_mode="coalesced", created_at=None): + ctype = getattr(initial_content, "type", None) + self.sink.append(("ctx", ctype)) + self.recorded_created_at.append(created_at) + return _BlockingCtx(self.sink, ctype, initial_content, self.blocked) + + +class _PlainAsyncIterator: + """An async iterator with no aclose (the AsyncIterator contract minimum).""" + + def __init__(self, events): + self._events = iter(events) + + def __aiter__(self): + return self + + async def __anext__(self): + try: + return next(self._events) + except StopIteration: + raise StopAsyncIteration from None + + +@pytest.mark.asyncio +async def test_auto_send_closes_source_when_cancelled_mid_delivery(): + """Cancelling delivery while the backend blocks must close the event source. + + The turn object pins its event generator, so GC cannot rescue it: when + auto_send returns with the source still suspended at a yield, the tap's + finally never runs and the harness CLI subprocess leaks. The source is held + by a local here for the whole test, and nothing calls gc.collect(). + """ + streaming = _BlockingStreaming() + closed: list[bool] = [] + + async def _recording_source(): + try: + yield StreamTaskMessageStart( + type="start", + index=0, + content=TextContent(type="text", author="agent", content=""), + ) + yield StreamTaskMessageDelta( + type="delta", + index=0, + delta=TextDelta(type="text", text_delta="hi"), + ) + yield StreamTaskMessageDone(type="done", index=0) + finally: + closed.append(True) + + source = _recording_source() + task = asyncio.create_task(auto_send(source, task_id="task1", tracer=None, streaming=streaming)) + await asyncio.wait_for(streaming.blocked.wait(), timeout=5) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert closed == [True] + assert ("close", "text") in [(s[0], s[1]) for s in streaming.sink] + + +@pytest.mark.asyncio +async def test_auto_send_accepts_source_without_aclose(): + """A plain async iterator (no aclose) must deliver exactly as before.""" + streaming = _FakeStreaming() + events = [ + StreamTaskMessageStart( + type="start", + index=0, + content=TextContent(type="text", author="agent", content=""), + ), + StreamTaskMessageDelta( + type="delta", + index=0, + delta=TextDelta(type="text", text_delta="Hi"), + ), + StreamTaskMessageDone(type="done", index=0), + ] + result = await auto_send(_PlainAsyncIterator(events), task_id="task1", tracer=None, streaming=streaming) + + assert result.final_text == "Hi" + kinds = [s[0] for s in streaming.sink] + assert kinds.count("open") == 1 + assert kinds.count("close") == 1 diff --git a/tests/lib/core/harness/test_emitter.py b/tests/lib/core/harness/test_emitter.py index 3f70660ec..a9ec846b0 100644 --- a/tests/lib/core/harness/test_emitter.py +++ b/tests/lib/core/harness/test_emitter.py @@ -140,3 +140,54 @@ async def test_emitter_auto_send_turn_reads_usage_after_exhaustion(): result = await emitter.auto_send_turn(turn) assert result.usage == real_usage assert result.usage.input_tokens == 11 and result.usage.total_tokens == 33 + + +class _PinnedTurn: + """A turn that pins its event generator, as the real CLI taps do. + + `closed` records that the generator's finally ran (what terminates the CLI + subprocess in the scaffolds). + """ + + def __init__(self, events_list): + self._events_list = events_list + self.closed: list[bool] = [] + self._gen = None + + @property + def events(self): + if self._gen is None: + self._gen = self._stream() + return self._gen + + async def _stream(self): + try: + for e in self._events_list: + yield e + finally: + self.closed.append(True) + + def usage(self): + return TurnUsage() + + +@pytest.mark.asyncio +async def test_emitter_yield_turn_closes_source_on_early_close(): + """Closing the delivery generator must close the turn's event source. + + No gc.collect() and no sleep: the turn keeps the generator referenced, and + async-generator finalization is too late for a disconnected client anyway. + """ + events = [ + StreamTaskMessageStart(type="start", index=0, content=TextContent(type="text", author="agent", content="")), + StreamTaskMessageDelta(type="delta", index=0, delta=TextDelta(type="text", text_delta="hi")), + StreamTaskMessageDone(type="done", index=0), + ] + turn = _PinnedTurn(events) + emitter = UnifiedEmitter(task_id="t", trace_id=None, parent_span_id=None) + gen = emitter.yield_turn(turn) + first = await gen.__anext__() + await gen.aclose() + + assert first.index == 0 + assert turn.closed == [True] diff --git a/tests/lib/core/harness/test_yield_delivery.py b/tests/lib/core/harness/test_yield_delivery.py index 21c93a95c..3267b2a0b 100644 --- a/tests/lib/core/harness/test_yield_delivery.py +++ b/tests/lib/core/harness/test_yield_delivery.py @@ -76,3 +76,30 @@ async def test_flush_runs_on_early_close(): await gen.aclose() # triggers the finally -> flush() assert fake.started_names == ["Bash"] assert fake.ended_outputs == [None] # flush closed the unpaired span (incomplete, no output) + + +@pytest.mark.asyncio +async def test_source_closed_when_consumer_closes_early(): + """Closing the delivery generator must close the upstream event source. + + A client disconnect (or any early break) closes the generator handed to the + caller. The source is pinned by the turn object, so if it is left suspended + at a yield the tap's finally never runs and the CLI subprocess leaks. No + gc.collect() here: the source stays referenced for the whole test. + """ + closed: list[bool] = [] + + async def _recording_source(): + try: + yield StreamTaskMessageDone(type="done", index=0) + yield StreamTaskMessageDone(type="done", index=1) + finally: + closed.append(True) + + source = _recording_source() + gen = yield_events(source, tracer=None) + first = await gen.__anext__() + await gen.aclose() + + assert first.index == 0 + assert closed == [True] diff --git a/tests/test_acp_streaming_disconnect.py b/tests/test_acp_streaming_disconnect.py new file mode 100644 index 000000000..9fe2720ed --- /dev/null +++ b/tests/test_acp_streaming_disconnect.py @@ -0,0 +1,39 @@ +"""A client that disconnects mid-stream must close the handler's generator. + +Starlette closes the response body iterator on disconnect. If that iterator +does not close the handler generator it loops over, the handler (and any +harness turn and CLI subprocess under it) stays suspended until garbage +collection, which never happens while something still references it. +""" + +from __future__ import annotations + +from typing import Any, AsyncGenerator, cast + +import pytest + +from agentex.types.text_delta import TextDelta +from agentex.types.task_message_update import StreamTaskMessageDelta +from agentex.lib.sdk.fastacp.impl.sync_acp import SyncACP + + +@pytest.mark.asyncio +async def test_closing_the_stream_closes_the_handler_generator() -> None: + closed: list[bool] = [] + + async def handler_stream(): + try: + for text in ("a", "b", "c"): + yield StreamTaskMessageDelta(type="delta", index=0, delta=TextDelta(type="text", text_delta=text)) + finally: + closed.append(True) + + source = handler_stream() + response = await SyncACP()._handle_streaming_response("req-1", source) + body = cast(AsyncGenerator[Any, None], response.body_iterator) + + first = await body.__anext__() + await body.aclose() + + assert '"text_delta":"a"' in str(first) + assert closed == [True]