mirror of
https://github.com/usestrix/strix.git
synced 2026-08-16 09:26:39 +02:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6797c96239 | ||
|
|
979aca7a32 |
+52
-1
@@ -13,7 +13,13 @@ from agents import RunConfig, Runner
|
|||||||
from agents.exceptions import AgentsException, MaxTurnsExceeded, UserError
|
from agents.exceptions import AgentsException, MaxTurnsExceeded, UserError
|
||||||
from agents.sandbox.errors import ExecTransportError
|
from agents.sandbox.errors import ExecTransportError
|
||||||
from docker import errors as docker_errors # type: ignore[import-untyped, unused-ignore]
|
from docker import errors as docker_errors # type: ignore[import-untyped, unused-ignore]
|
||||||
from openai import APIError
|
from openai import (
|
||||||
|
APIConnectionError,
|
||||||
|
APIError,
|
||||||
|
APIStatusError,
|
||||||
|
APITimeoutError,
|
||||||
|
RateLimitError,
|
||||||
|
)
|
||||||
|
|
||||||
from strix.core.hooks import BudgetExceededError
|
from strix.core.hooks import BudgetExceededError
|
||||||
from strix.core.inputs import child_initial_input
|
from strix.core.inputs import child_initial_input
|
||||||
@@ -78,6 +84,34 @@ async def _compact_session(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_TRANSIENT_MODEL_STATUS_CODES = frozenset({408, 500, 502, 503, 504})
|
||||||
|
_MAX_TRANSIENT_MODEL_RETRIES = 4
|
||||||
|
_TRANSIENT_MODEL_RETRY_BASE_DELAY_S = 2.0
|
||||||
|
_TRANSIENT_MODEL_RETRY_MAX_DELAY_S = 30.0
|
||||||
|
|
||||||
|
|
||||||
|
def _model_error_status_code(exc: BaseException) -> int | None:
|
||||||
|
code = getattr(exc, "status_code", None)
|
||||||
|
return code if isinstance(code, int) else None
|
||||||
|
|
||||||
|
|
||||||
|
def _is_transient_model_error(exc: BaseException) -> bool:
|
||||||
|
if isinstance(exc, RateLimitError):
|
||||||
|
return False
|
||||||
|
if isinstance(exc, APITimeoutError | APIConnectionError):
|
||||||
|
return True
|
||||||
|
if isinstance(exc, APIStatusError):
|
||||||
|
return exc.status_code in _TRANSIENT_MODEL_STATUS_CODES
|
||||||
|
if isinstance(exc, APIError):
|
||||||
|
return _model_error_status_code(exc) is None
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _transient_model_retry_delay(attempt: int) -> float:
|
||||||
|
delay = _TRANSIENT_MODEL_RETRY_BASE_DELAY_S * float(2 ** (attempt - 1))
|
||||||
|
return min(delay, _TRANSIENT_MODEL_RETRY_MAX_DELAY_S)
|
||||||
|
|
||||||
|
|
||||||
async def run_agent_loop(
|
async def run_agent_loop(
|
||||||
*,
|
*,
|
||||||
agent: Any,
|
agent: Any,
|
||||||
@@ -387,6 +421,7 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
|
|||||||
) -> RunResultBase | None:
|
) -> RunResultBase | None:
|
||||||
image_strips = 0
|
image_strips = 0
|
||||||
compactions = 0
|
compactions = 0
|
||||||
|
model_retries = 0
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
await coordinator.mark_running(agent_id)
|
await coordinator.mark_running(agent_id)
|
||||||
@@ -488,6 +523,22 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
|
|||||||
)
|
)
|
||||||
input_data = []
|
input_data = []
|
||||||
continue
|
continue
|
||||||
|
if model_retries < _MAX_TRANSIENT_MODEL_RETRIES and _is_transient_model_error(exc):
|
||||||
|
model_retries += 1
|
||||||
|
delay = _transient_model_retry_delay(model_retries)
|
||||||
|
logger.warning(
|
||||||
|
"transient model/provider error for %s; replaying turn "
|
||||||
|
"(attempt %d/%d, backoff %.1fs): %r",
|
||||||
|
agent_id,
|
||||||
|
model_retries,
|
||||||
|
_MAX_TRANSIENT_MODEL_RETRIES,
|
||||||
|
delay,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
await asyncio.sleep(delay)
|
||||||
|
if session is not None:
|
||||||
|
input_data = []
|
||||||
|
continue
|
||||||
if not interactive:
|
if not interactive:
|
||||||
raise
|
raise
|
||||||
if isinstance(exc, MaxTurnsExceeded):
|
if isinstance(exc, MaxTurnsExceeded):
|
||||||
|
|||||||
@@ -0,0 +1,157 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any, cast
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
import pytest
|
||||||
|
from agents import RunConfig, Runner
|
||||||
|
from openai import (
|
||||||
|
APIConnectionError,
|
||||||
|
APIError,
|
||||||
|
APIStatusError,
|
||||||
|
APITimeoutError,
|
||||||
|
BadRequestError,
|
||||||
|
InternalServerError,
|
||||||
|
RateLimitError,
|
||||||
|
)
|
||||||
|
|
||||||
|
from strix.core import execution
|
||||||
|
from strix.core.agents import AgentCoordinator
|
||||||
|
|
||||||
|
|
||||||
|
def _request() -> httpx.Request:
|
||||||
|
return httpx.Request("POST", "https://api.openai.com/v1/responses")
|
||||||
|
|
||||||
|
|
||||||
|
def _midstream_api_error() -> APIError:
|
||||||
|
return APIError("An error occurred while processing the request.", _request(), body=None)
|
||||||
|
|
||||||
|
|
||||||
|
def _status_error(status: int) -> APIStatusError:
|
||||||
|
return APIStatusError(
|
||||||
|
f"status {status}",
|
||||||
|
response=httpx.Response(status_code=status, request=_request()),
|
||||||
|
body=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_midstream_api_error_is_transient() -> None:
|
||||||
|
assert execution._is_transient_model_error(_midstream_api_error()) is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_network_errors_are_transient() -> None:
|
||||||
|
assert execution._is_transient_model_error(APITimeoutError(_request())) is True
|
||||||
|
assert execution._is_transient_model_error(APIConnectionError(request=_request())) is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_server_errors_are_transient() -> None:
|
||||||
|
assert (
|
||||||
|
execution._is_transient_model_error(
|
||||||
|
InternalServerError("boom", response=httpx.Response(500, request=_request()), body=None)
|
||||||
|
)
|
||||||
|
is True
|
||||||
|
)
|
||||||
|
for status in (502, 503, 504, 408):
|
||||||
|
assert execution._is_transient_model_error(_status_error(status)) is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_rate_limit_is_not_retried_here() -> None:
|
||||||
|
rate_limited = RateLimitError(
|
||||||
|
"slow down", response=httpx.Response(429, request=_request()), body=None
|
||||||
|
)
|
||||||
|
assert execution._is_transient_model_error(rate_limited) is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_client_errors_are_not_transient() -> None:
|
||||||
|
bad_request = BadRequestError(
|
||||||
|
"bad", response=httpx.Response(400, request=_request()), body=None
|
||||||
|
)
|
||||||
|
assert execution._is_transient_model_error(bad_request) is False
|
||||||
|
assert execution._is_transient_model_error(_status_error(404)) is False
|
||||||
|
assert execution._is_transient_model_error(ValueError("nope")) is False
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeStream:
|
||||||
|
def __init__(self, exc: BaseException | None = None) -> None:
|
||||||
|
self._exc = exc
|
||||||
|
self._events: list[Any] = []
|
||||||
|
self.run_loop_exception: BaseException | None = None
|
||||||
|
|
||||||
|
async def stream_events(self) -> Any:
|
||||||
|
if self._exc is not None:
|
||||||
|
raise self._exc
|
||||||
|
for event in self._events:
|
||||||
|
yield event
|
||||||
|
|
||||||
|
|
||||||
|
def _patch_fast_backoff(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
monkeypatch.setattr(execution, "_TRANSIENT_MODEL_RETRY_BASE_DELAY_S", 0.0)
|
||||||
|
monkeypatch.setattr(execution, "_TRANSIENT_MODEL_RETRY_MAX_DELAY_S", 0.0)
|
||||||
|
|
||||||
|
|
||||||
|
async def _run_once(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
streams: list[_FakeStream],
|
||||||
|
) -> Any:
|
||||||
|
_patch_fast_backoff(monkeypatch)
|
||||||
|
calls = {"n": 0}
|
||||||
|
|
||||||
|
def _fake_run_streamed(*_args: Any, **_kwargs: Any) -> _FakeStream:
|
||||||
|
stream = streams[calls["n"]]
|
||||||
|
calls["n"] += 1
|
||||||
|
return stream
|
||||||
|
|
||||||
|
monkeypatch.setattr(Runner, "run_streamed", _fake_run_streamed)
|
||||||
|
|
||||||
|
coordinator = AgentCoordinator()
|
||||||
|
await coordinator.register("root", "strix", parent_id=None)
|
||||||
|
|
||||||
|
result = await execution._run_cycle(
|
||||||
|
object(),
|
||||||
|
coordinator,
|
||||||
|
"root",
|
||||||
|
input_data="task",
|
||||||
|
run_config=cast("RunConfig", object()),
|
||||||
|
context={},
|
||||||
|
max_turns=5,
|
||||||
|
session=None,
|
||||||
|
interactive=False,
|
||||||
|
event_sink=None,
|
||||||
|
hooks=None,
|
||||||
|
)
|
||||||
|
return result, calls["n"], coordinator
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_run_cycle_retries_transient_midstream_error(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
streams = [_FakeStream(exc=_midstream_api_error()), _FakeStream()]
|
||||||
|
result, attempts, _coordinator = await _run_once(monkeypatch, streams)
|
||||||
|
|
||||||
|
assert result is streams[1]
|
||||||
|
assert attempts == 2
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_run_cycle_gives_up_after_max_retries(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
streams = [
|
||||||
|
_FakeStream(exc=_midstream_api_error())
|
||||||
|
for _ in range(execution._MAX_TRANSIENT_MODEL_RETRIES + 1)
|
||||||
|
]
|
||||||
|
with pytest.raises(APIError):
|
||||||
|
await _run_once(monkeypatch, streams)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_run_cycle_does_not_retry_permanent_error(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
bad_request = BadRequestError(
|
||||||
|
"bad", response=httpx.Response(400, request=_request()), body=None
|
||||||
|
)
|
||||||
|
streams = [_FakeStream(exc=bad_request), _FakeStream()]
|
||||||
|
with pytest.raises(BadRequestError):
|
||||||
|
await _run_once(monkeypatch, streams)
|
||||||
Reference in New Issue
Block a user