From 4eae739b227295663c14bbf180294e1edfb7e397 Mon Sep 17 00:00:00 2001 From: Bit Wei Date: Thu, 3 Sep 2026 11:17:13 -0500 Subject: [PATCH] fix(execution): cancel agent streams during teardown --- strix/config/models.py | 89 ++++++++++---------- strix/core/execution.py | 39 ++++++++- tests/test_execution.py | 132 ++++++++++++++++++++++++++++++ tests/test_stream_idle_timeout.py | 85 ++++++++++++++++++- 4 files changed, 296 insertions(+), 49 deletions(-) diff --git a/strix/config/models.py b/strix/config/models.py index babb643f6..c7fe71096 100644 --- a/strix/config/models.py +++ b/strix/config/models.py @@ -8,7 +8,6 @@ import inspect import logging import os import time -from collections.abc import AsyncGenerator from typing import TYPE_CHECKING, Any, cast from agents import ( @@ -43,7 +42,7 @@ from strix.config.tool_call_limits import TurnToolCallLimiter if TYPE_CHECKING: - from collections.abc import AsyncIterator + from collections.abc import AsyncGenerator, AsyncIterator from agents.agent_output import AgentOutputSchemaBase from agents.handoffs import Handoff @@ -78,6 +77,23 @@ def _retry_statusless_provider_errors(context: RetryPolicyContext) -> bool: return normalized.status_code is None +async def _close_stream(stream: Any) -> None: + """Best-effort close for async and synchronous provider streams.""" + aclose = getattr(stream, "aclose", None) + if callable(aclose): + with contextlib.suppress(Exception): + result = aclose() + if inspect.isawaitable(result): + await result + return + close = getattr(stream, "close", None) + if callable(close): + with contextlib.suppress(Exception): + result = close() + if inspect.isawaitable(result): + await result + + class _CodexResponsesModel(OpenAIResponsesModel): """Responses model for the ChatGPT subscription backend (always streamed, stateless).""" @@ -146,21 +162,7 @@ class _CodexResponsesModel(OpenAIResponsesModel): raise guardrail from exc raise finally: - await self._aclose(events) - - @staticmethod - async def _aclose(events: Any) -> None: - aclose = getattr(events, "aclose", None) - if callable(aclose): - with contextlib.suppress(Exception): - await aclose() - return - close = getattr(events, "close", None) - if callable(close): - with contextlib.suppress(Exception): - result = close() - if inspect.isawaitable(result): - await result + await _close_stream(events) class _NonStreamingModel(Model): @@ -351,39 +353,38 @@ class _TurnGuardModel(Model): conversation_id=conversation_id, prompt=prompt, ) - async for event in _with_idle_timeout(stream, self._stream_idle_timeout): - guarded = _guard_event(event, rewriter, limiter) - if guarded is not None: - yield guarded + async with contextlib.aclosing( + _with_idle_timeout(stream, self._stream_idle_timeout) + ) as events: + async for event in events: + guarded = _guard_event(event, rewriter, limiter) + if guarded is not None: + yield guarded self._log_dropped(limiter) -async def _aclose(stream: AsyncIterator[TResponseStreamEvent]) -> None: - if isinstance(stream, AsyncGenerator): - with contextlib.suppress(Exception): - await stream.aclose() - - async def _with_idle_timeout( stream: AsyncIterator[TResponseStreamEvent], timeout: float -) -> AsyncIterator[TResponseStreamEvent]: - if timeout <= 0: - async for event in stream: - yield event - return - - iterator = stream.__aiter__() - while True: - try: - event = await asyncio.wait_for(iterator.__anext__(), timeout) - except StopAsyncIteration: +) -> AsyncGenerator[TResponseStreamEvent]: + try: + if timeout <= 0: + async for event in stream: + yield event return - except TimeoutError: - await _aclose(stream) - message = f"model stream produced no event for {timeout:.0f}s" - logger.warning("%s; abandoning the turn", message) - raise TimeoutError(message) from None - yield event + + iterator = stream.__aiter__() + while True: + try: + event = await asyncio.wait_for(iterator.__anext__(), timeout) + except StopAsyncIteration: + return + except TimeoutError: + message = f"model stream produced no event for {timeout:.0f}s" + logger.warning("%s; abandoning the turn", message) + raise TimeoutError(message) from None + yield event + finally: + await _close_stream(stream) def _guard_event( diff --git a/strix/core/execution.py b/strix/core/execution.py index dfcd39fa2..a152d1e92 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -53,6 +53,32 @@ StreamEventSink = Callable[[str, Any], None] _INPUT_REJECTION_CODES = frozenset({400, 404, 422}) _MAX_COMPACTIONS_PER_CYCLE = 2 +_STREAM_CANCEL_SETTLE_TIMEOUT_S = 1.0 + + +async def _cancel_and_settle_stream(stream: Any) -> None: + """Cancel an unfinished SDK stream and briefly wait for its run loop to unwind.""" + if stream is None: + return + if not getattr(stream, "is_complete", False): + try: + stream.cancel(mode="immediate") + except Exception: + logger.exception("failed to cancel unfinished agent stream") + + run_loop_task = getattr(stream, "run_loop_task", None) + if not isinstance(run_loop_task, asyncio.Task) or run_loop_task.done(): + return + task = cast("asyncio.Task[Any]", run_loop_task) # type: ignore[redundant-cast] + done, _pending = await asyncio.wait( + {task}, + timeout=_STREAM_CANCEL_SETTLE_TIMEOUT_S, + ) + if not done: + logger.warning( + "agent stream run loop did not settle within %.1fs after cancellation", + _STREAM_CANCEL_SETTLE_TIMEOUT_S, + ) @cache @@ -685,8 +711,8 @@ async def _run_cycle( # noqa: PLR0912, PLR0915 session=session, hooks=hooks, ) - await coordinator.attach_stream(agent_id, stream) try: + await coordinator.attach_stream(agent_id, stream) try: async for event in stream.stream_events(): if event_sink is not None: @@ -698,7 +724,11 @@ async def _run_cycle( # noqa: PLR0912, PLR0915 raise stream.run_loop_exception if refusal := _structured_provider_refusal(stream): raise ProviderRefusalError(refusal) - except (BudgetExceededError, BudgetPausedError, SubagentBudgetReservedError): + except ( + BudgetExceededError, + BudgetPausedError, + SubagentBudgetReservedError, + ): raise except RuntimeError as stream_exc: if "after shutdown" not in str(stream_exc): @@ -716,7 +746,10 @@ async def _run_cycle( # noqa: PLR0912, PLR0915 exc_info=True, ) finally: - await coordinator.detach_stream(agent_id, stream) + try: + await _cancel_and_settle_stream(stream) + finally: + await coordinator.detach_stream(agent_id, stream) except BudgetPausedError as exc: logger.info("agent %s paused at the scan budget limit: %s", agent_id, exc) await coordinator.pause_for_budget(agent_id) diff --git a/tests/test_execution.py b/tests/test_execution.py index 8fbf18ff5..417e8f636 100644 --- a/tests/test_execution.py +++ b/tests/test_execution.py @@ -53,6 +53,47 @@ class _StructuredRefusalStream: return +class _BlockingRunStream: + def __init__(self) -> None: + self.is_complete = False + self.run_loop_exception: BaseException | None = None + self.new_items: list[Any] = [] + self.cancel_mode: str | None = None + self.cancel_calls = 0 + self.stream_events_started = asyncio.Event() + self.run_loop_settled = asyncio.Event() + self.cleanup_order: list[str] = [] + self.run_loop_task = asyncio.create_task(self._run_loop()) + + async def _run_loop(self) -> None: + try: + await asyncio.Event().wait() + finally: + # Make ordering observable: cancellation must be joined before the + # coordinator forgets which stream belonged to this agent. + await asyncio.sleep(0) + self.cleanup_order.append("run-loop-settled") + self.run_loop_settled.set() + + async def stream_events(self) -> Any: + try: + self.stream_events_started.set() + await asyncio.Event().wait() + except asyncio.CancelledError: + # Match openai-agents 0.19.0: cancellation while waiting for an + # event cancels the run result before propagating. + self.cancel() + raise + for event in _NO_STREAM_EVENTS: + yield event + + def cancel(self, mode: str = "immediate") -> None: + self.cancel_calls += 1 + self.cancel_mode = mode + self.is_complete = True + self.run_loop_task.cancel() + + async def _call_finish_scan( coordinator: AgentCoordinator, agent_id: str, parent_id: str | None ) -> dict[str, Any]: @@ -861,6 +902,97 @@ async def test_structured_provider_refusal_fails_interactive_agent( assert coordinator.errors["root"] == refusal +@pytest.mark.asyncio +async def test_cancelled_while_attaching_cancels_settles_and_detaches_stream( + monkeypatch: pytest.MonkeyPatch, +) -> None: + stream = _BlockingRunStream() + monkeypatch.setattr( + "strix.core.execution.Runner.run_streamed", lambda *_args, **_kwargs: stream + ) + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + attach_started = asyncio.Event() + + async def _blocked_attach(_agent_id: str, _stream: Any) -> None: + attach_started.set() + await asyncio.Event().wait() + + async def _record_detach(agent_id: str, attached_stream: Any) -> None: + stream.cleanup_order.append("detach") + await AgentCoordinator.detach_stream(coordinator, agent_id, attached_stream) + + monkeypatch.setattr(coordinator, "attach_stream", _blocked_attach) + monkeypatch.setattr(coordinator, "detach_stream", _record_detach) + + task = asyncio.create_task( + execution._run_cycle( + MagicMock(), + coordinator, + "root", + input_data="task", + run_config=MagicMock(), + context={}, + max_turns=5, + session=None, + interactive=False, + event_sink=None, + hooks=None, + ) + ) + await asyncio.wait_for(attach_started.wait(), timeout=1.0) + + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert stream.cancel_mode == "immediate" + assert stream.cancel_calls == 1 + assert stream.run_loop_task.done() + assert stream.run_loop_settled.is_set() + assert stream.cleanup_order == ["run-loop-settled", "detach"] + assert coordinator.runtimes["root"].stream is None + + +@pytest.mark.asyncio +async def test_sdk_self_cancel_is_settled_and_not_cancelled_twice( + monkeypatch: pytest.MonkeyPatch, +) -> None: + stream = _BlockingRunStream() + monkeypatch.setattr( + "strix.core.execution.Runner.run_streamed", lambda *_args, **_kwargs: stream + ) + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + + task = asyncio.create_task( + execution._run_cycle( + MagicMock(), + coordinator, + "root", + input_data="task", + run_config=MagicMock(), + context={}, + max_turns=5, + session=None, + interactive=False, + event_sink=None, + hooks=None, + ) + ) + await asyncio.wait_for(stream.stream_events_started.wait(), timeout=1.0) + + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert stream.cancel_mode == "immediate" + assert stream.cancel_calls == 1 + assert stream.run_loop_task.done() + assert stream.run_loop_settled.is_set() + assert coordinator.runtimes["root"].stream is None + + @pytest.mark.asyncio async def test_structured_provider_refusal_fails_noninteractive_child( tmp_path: Any, diff --git a/tests/test_stream_idle_timeout.py b/tests/test_stream_idle_timeout.py index 9d618978f..54bac051f 100644 --- a/tests/test_stream_idle_timeout.py +++ b/tests/test_stream_idle_timeout.py @@ -13,7 +13,7 @@ import json import threading import time from http.server import BaseHTTPRequestHandler, HTTPServer -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, cast import pytest from agents.model_settings import ModelSettings @@ -27,12 +27,52 @@ from strix.config.models import StrixProvider, _TurnGuardModel, _with_idle_timeo if TYPE_CHECKING: - from collections.abc import AsyncIterator, Iterator + from collections.abc import AsyncGenerator, AsyncIterator, Iterator _STALL_SECONDS = 30.0 +class _ClosableBlockingStream: + def __init__(self) -> None: + self.started = asyncio.Event() + self.closed = False + + def __aiter__(self) -> _ClosableBlockingStream: + return self + + async def __anext__(self) -> Any: + self.started.set() + await asyncio.Event().wait() + raise StopAsyncIteration + + async def aclose(self) -> None: + self.closed = True + + +class _OneEventClosableStream(_ClosableBlockingStream): + def __init__(self) -> None: + super().__init__() + self._sent = False + + async def __anext__(self) -> Any: + if not self._sent: + self._sent = True + return "event" + return await super().__anext__() + + +class _StreamModel(Model): + def __init__(self, stream: _OneEventClosableStream) -> None: + self.stream = stream + + async def get_response(self, *args: Any, **kwargs: Any) -> Any: + raise NotImplementedError + + def stream_response(self, *_args: Any, **_kwargs: Any) -> Any: + return self.stream + + def _chunk(text: str) -> bytes: payload = { "id": "chatcmpl-1", @@ -129,6 +169,47 @@ async def test_events_keep_flowing_while_the_stream_is_alive() -> None: assert seen == [f"event-{i}" for i in range(5)] +@pytest.mark.asyncio +async def test_consumer_cancellation_closes_the_inner_stream() -> None: + stream = _ClosableBlockingStream() + task = asyncio.create_task( + anext(_with_idle_timeout(stream, 60.0)), + ) + await asyncio.wait_for(stream.started.wait(), timeout=1.0) + + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert stream.closed is True + + +@pytest.mark.asyncio +async def test_closing_guarded_model_stream_closes_provider_stream_immediately() -> None: + provider_stream = _OneEventClosableStream() + model = _TurnGuardModel(_StreamModel(provider_stream), stream_idle_timeout=60.0) + guarded_stream = cast( + "AsyncGenerator[Any]", + model.stream_response( + None, + "go", + ModelSettings(), + [], + None, + [], + ModelTracing.DISABLED, + previous_response_id=None, + conversation_id=None, + prompt=None, + ), + ) + + assert await anext(guarded_stream) == "event" + await guarded_stream.aclose() + + assert provider_stream.closed is True + + @pytest.fixture def _reset_settings(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: for key in ("STRIX_LLM", "LLM_DISABLE_STREAMING", "LLM_STREAM_IDLE_TIMEOUT"):