mirror of
https://github.com/usestrix/strix.git
synced 2026-10-11 03:37:54 +00:00
feat: implement retry logic for restartable stream errors in agent execution
This commit is contained in:
parent
7141ccff62
commit
fa0cc8af13
2 changed files with 140 additions and 1 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue