From fa0cc8af13830737771eb16c2734e813db5e22c2 Mon Sep 17 00:00:00 2001 From: TheArtificialQ Date: Fri, 26 Jun 2026 15:09:19 +0200 Subject: [PATCH] feat: implement retry logic for restartable stream errors in agent execution --- strix/core/execution.py | 30 ++++++++++- tests/test_execution.py | 111 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 140 insertions(+), 1 deletion(-) diff --git a/strix/core/execution.py b/strix/core/execution.py index 06dc3ddf..330861f3 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -10,9 +10,10 @@ from collections.abc import Callable from typing import TYPE_CHECKING, Any, cast from agents import RunConfig, Runner -from agents.exceptions import AgentsException, MaxTurnsExceeded, UserError +from agents.exceptions import AgentsException, MaxTurnsExceeded, ModelBehaviorError, UserError from agents.sandbox.errors import ExecTransportError from docker import errors as docker_errors # type: ignore[import-untyped, unused-ignore] +from litellm.exceptions import BadRequestError, MidStreamFallbackError from openai import APIError from strix.core.hooks import BudgetExceededError @@ -36,6 +37,12 @@ logger = logging.getLogger(__name__) StreamEventSink = Callable[[str, Any], None] _INPUT_REJECTION_CODES = frozenset({400, 404, 422}) +_STREAM_RESTART_LIMIT = 3 +_STREAM_RESTARTABLE_EXCEPTIONS = ( + ModelBehaviorError, + MidStreamFallbackError, + BadRequestError, +) async def run_agent_loop( @@ -346,7 +353,9 @@ async def _run_cycle( # noqa: PLR0912, PLR0915 hooks: RunHooks[dict[str, Any]] | None, ) -> RunResultBase | None: image_strips = 0 + stream_restarts = 0 while True: + restart_stream = False try: await coordinator.mark_running(agent_id) stream = Runner.run_streamed( @@ -388,8 +397,27 @@ async def _run_cycle( # noqa: PLR0912, PLR0915 agent_id, exc_info=True, ) + except _STREAM_RESTARTABLE_EXCEPTIONS as exc: + if stream_restarts >= _STREAM_RESTART_LIMIT: + raise + stream_restarts += 1 + restart_stream = True + logger.warning( + "Restartable stream exception for %s", + agent_id, + exc_info=True, + ) + logger.warning( + "Restarting stream for %s after %s (%d/%d)", + agent_id, + type(exc).__name__, + stream_restarts, + _STREAM_RESTART_LIMIT, + ) finally: await coordinator.detach_stream(agent_id, stream) + if restart_stream: + continue except BudgetExceededError as exc: logger.info( "agent %s reached the scan budget limit; stopping the scan: %s", agent_id, exc diff --git a/tests/test_execution.py b/tests/test_execution.py index 59a37e65..1066ad7b 100644 --- a/tests/test_execution.py +++ b/tests/test_execution.py @@ -3,10 +3,14 @@ from __future__ import annotations import asyncio +import logging +from typing import Any import pytest +from strix.core import execution from strix.core.agents import AgentCoordinator +from strix.core.execution import _run_cycle @pytest.mark.asyncio @@ -42,3 +46,110 @@ async def test_wait_for_message_returns_immediately_after_budget_stop() -> None: # No pending messages, but the stop flag short-circuits the wait. await asyncio.wait_for(coordinator.wait_for_message("agent"), timeout=1.0) + + +class _RestartableStreamError(Exception): + pass + + +class _FakeStream: + def __init__( + self, + *, + events: tuple[Any, ...] | list[Any] = (), + exc: Exception | None = None, + ) -> None: + self._events = list(events) + self._exc = exc + self.run_loop_exception: Exception | None = None + + async def stream_events(self): + for event in self._events: + yield event + if self._exc is not None: + raise self._exc + + +@pytest.mark.asyncio +async def test_run_cycle_retries_restartable_stream_errors( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + + streams = [ + _FakeStream(exc=_RestartableStreamError("first")), + _FakeStream(exc=_RestartableStreamError("second")), + _FakeStream(events=[{"type": "done"}]), + ] + emitted: list[tuple[str, Any]] = [] + + def run_streamed(*_args: object, **_kwargs: object) -> _FakeStream: + return streams.pop(0) + + monkeypatch.setattr( + execution, + "_STREAM_RESTARTABLE_EXCEPTIONS", + (_RestartableStreamError,), + ) + monkeypatch.setattr(execution.Runner, "run_streamed", run_streamed) + + with caplog.at_level(logging.WARNING): + result = await _run_cycle( + object(), + coordinator, + "root", + input_data=[{"role": "user", "content": "hello"}], + run_config=object(), + context={}, + max_turns=4, + session=None, + interactive=False, + event_sink=lambda agent_id, event: emitted.append((agent_id, event)), + hooks=None, + ) + + assert result is not None + assert emitted == [("root", {"type": "done"})] + assert coordinator.runtimes["root"].stream is None + assert sum("Restarting stream for root" in message for message in caplog.messages) == 2 + + +@pytest.mark.asyncio +async def test_run_cycle_stops_retrying_after_restart_limit( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + coordinator = AgentCoordinator() + await coordinator.register("root", "strix", parent_id=None) + + streams = [_FakeStream(exc=_RestartableStreamError(f"boom-{index}")) for index in range(4)] + + def run_streamed(*_args: object, **_kwargs: object) -> _FakeStream: + return streams.pop(0) + + monkeypatch.setattr( + execution, + "_STREAM_RESTARTABLE_EXCEPTIONS", + (_RestartableStreamError,), + ) + monkeypatch.setattr(execution.Runner, "run_streamed", run_streamed) + + with caplog.at_level(logging.WARNING), pytest.raises(_RestartableStreamError, match="boom-3"): + await _run_cycle( + object(), + coordinator, + "root", + input_data=[], + run_config=object(), + context={}, + max_turns=4, + session=None, + interactive=False, + event_sink=None, + hooks=None, + ) + + assert coordinator.runtimes["root"].stream is None + assert sum("Restarting stream for root" in message for message in caplog.messages) == 3