fix(execution): cancel agent streams during teardown

This commit is contained in:
Bit Wei 2026-09-03 11:17:13 -05:00
parent 5d015df6b1
commit 4eae739b22
4 changed files with 296 additions and 49 deletions

View file

@ -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(

View file

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

View file

@ -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,

View file

@ -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"):