Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 10 additions & 1 deletion ifixai/providers/http.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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(
Expand Down
45 changes: 45 additions & 0 deletions ifixai/tests/providers/test_http_response_shapes.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Loading