diff --git a/ifixai/providers/http.py b/ifixai/providers/http.py index 1994c3e6..c199d94e 100644 --- a/ifixai/providers/http.py +++ b/ifixai/providers/http.py @@ -8,9 +8,11 @@ from ifixai.core.types import ChatMessage, ProviderConfig, RetrievedSource from ifixai.providers.base import ( + RETRYABLE_HTTP_STATUS_CODES, ChatProvider, ProviderAuthError, ProviderConnectionError, + ProviderOverloadedError, ProviderRateLimitError, ProviderResponseError, ProviderTimeoutError, @@ -137,7 +139,7 @@ async def send_message( await asyncio.sleep(2**attempt) continue raise - except (ProviderConnectionError, ProviderTimeoutError) as exc: + except (ProviderConnectionError, ProviderTimeoutError, ProviderOverloadedError) as exc: last_error = exc if attempt < config.max_retries: await asyncio.sleep(2**attempt) @@ -174,6 +176,13 @@ async def _send_request( endpoint=endpoint, details="HTTP 429: rate limited", ) + if resp.status in RETRYABLE_HTTP_STATUS_CODES: + body = await resp.text() + raise ProviderOverloadedError( + provider="http", + endpoint=endpoint, + details=f"HTTP {resp.status}: {scrub_secrets(body[:500])}", + ) if resp.status >= 400: body = await resp.text() raise ProviderResponseError( diff --git a/ifixai/tests/providers/test_http_response_shapes.py b/ifixai/tests/providers/test_http_response_shapes.py index 2bfde2ba..ddd60a48 100644 --- a/ifixai/tests/providers/test_http_response_shapes.py +++ b/ifixai/tests/providers/test_http_response_shapes.py @@ -62,3 +62,48 @@ async def retrieve(_request): finally: await provider.aclose() await runner.cleanup() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status", [408, 500, 502, 503, 504, 400]) +async def test_http_gateway_failures_follow_retry_contract(monkeypatch, status): + from ifixai.providers.base import ProviderOverloadedError, is_fatal_provider_error + + seen = [] + async def chat(request): + seen.append(await request.json()) + if len(seen) == 1: + return aiohttp.web.json_response({"error": {"message": "owned gateway unavailable"}}, status=status) + return aiohttp.web.json_response({"choices": [{"message": {"content": "recovered reply"}}]}) + + async def no_delay(_seconds): + pass + + monkeypatch.setattr("ifixai.providers.http.asyncio.sleep", no_delay) + app = aiohttp.web.Application() + app.router.add_post("/v1/chat/completions", chat) + runner = aiohttp.web.AppRunner(app) + await runner.setup() + site = aiohttp.web.TCPSite(runner, "127.0.0.1", 0) + await site.start() + url = f"http://127.0.0.1:{site._server.sockets[0].getsockname()[1]}/v1" + provider = HttpProvider() + try: + config = ProviderConfig(provider="http", endpoint=url, max_retries=1) + if status == 400: + with pytest.raises(ProviderResponseError): + await provider.send_message([ChatMessage(content="hello")], config) + assert len(seen) == 1 + else: + assert await provider.send_message([ChatMessage(content="hello")], config) == "recovered reply" + assert len(seen) == 2 + seen.clear() + config.max_retries = 0 + expected = ProviderResponseError if status == 400 else ProviderOverloadedError + with pytest.raises(expected) as caught: + await provider.send_message([ChatMessage(content="hello")], config) + assert not is_fatal_provider_error(caught.value) + assert len(seen) == 1 + finally: + await provider.aclose() + await runner.cleanup()