diff --git a/strix/core/execution.py b/strix/core/execution.py index 508eccfe..5997e3db 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -184,9 +184,7 @@ def _is_transient_model_error(exc: BaseException) -> bool: return True code = _model_error_status_code(exc) if code is not None: - import litellm - - return bool(litellm._should_retry(code)) + return code >= 400 and code not in (401, 402, 403, 404) return isinstance(exc, APIError) diff --git a/tests/test_execution_transient_retry.py b/tests/test_execution_transient_retry.py index 57b26db3..f0276932 100644 --- a/tests/test_execution_transient_retry.py +++ b/tests/test_execution_transient_retry.py @@ -79,12 +79,13 @@ def test_content_guardrail_is_not_retried() -> None: assert execution._is_transient_model_error(guardrail) is False -def test_client_errors_are_not_transient() -> None: +def test_client_errors_are_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(bad_request) is True + for status in (401, 402, 403, 404): + assert execution._is_transient_model_error(_status_error(status)) is False assert execution._is_transient_model_error(ValueError("nope")) is False @@ -167,11 +168,8 @@ async def test_run_cycle_gives_up_after_max_retries( 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): + streams = [_FakeStream(exc=ValueError("nope")), _FakeStream()] + with pytest.raises(ValueError, match="nope"): await _run_once(monkeypatch, streams)