diff --git a/routstr/core/exceptions.py b/routstr/core/exceptions.py index b8dcdc38..88e0d370 100644 --- a/routstr/core/exceptions.py +++ b/routstr/core/exceptions.py @@ -18,6 +18,11 @@ class UpstreamError(Exception): string-matching the message. ``details`` holds optional structured, redaction-safe context. Both default to ``None`` for backwards compatibility. + + ``from_upstream_response`` is True only when ``status_code`` is the status + the upstream itself answered with, as opposed to a status this proxy chose + for a transport failure, timeout or internal fault. Callers use it to + decide whether a status is safe to retry. """ def __init__( @@ -26,11 +31,13 @@ class UpstreamError(Exception): status_code: int = 502, code: str | None = None, details: dict[str, object] | None = None, + from_upstream_response: bool = False, ): self.message = message self.status_code = status_code self.code = code self.details = details + self.from_upstream_response = from_upstream_response super().__init__(message) diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 00d47faa..77d623e7 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -36,6 +36,10 @@ class Settings(BaseSettings): # Core upstream_base_url: str = Field(default="", env="UPSTREAM_BASE_URL") upstream_api_key: str = Field(default="", env="UPSTREAM_API_KEY") + # Extra attempts against the same upstream on a transient 5xx. 0 disables. + upstream_5xx_retry_attempts: int = Field( + default=1, ge=0, env="UPSTREAM_5XX_RETRY_ATTEMPTS" + ) # Node info name: str = Field(default="ARoutstrNode", env="NAME") diff --git a/routstr/proxy.py b/routstr/proxy.py index 2ada5d3a..887c81f2 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -395,6 +395,12 @@ def _forwarding_allowed(path: str, method: str) -> bool: return method in _allowed_methods_for(_canonical_api_path(path)) +# Gateway conditions a retry usually clears. 500 is excluded: as likely to be a +# deterministic rejection that fails identically on the next attempt. +_RETRYABLE_UPSTREAM_5XX = frozenset({502, 503, 504}) +_UPSTREAM_5XX_RETRY_BACKOFF_SECONDS = 0.5 + + @proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) async def proxy( request: Request, path: str, session: AsyncSession = Depends(get_session) @@ -782,6 +788,8 @@ async def _proxy( await _finish_read_transaction(session) max_cost_for_model = candidate_max + retries_left = settings.upstream_5xx_retry_attempts + retry_index = 0 headers = upstream.prepare_headers(dict(request.headers)) try: @@ -834,8 +842,39 @@ async def _proxy( model_obj, reservation_snapshot, ) - except UpstreamError: - # Let the outer UpstreamError handler manage retry/revert + except UpstreamError as e: + # Only a gateway status the upstream itself answered with: + # re-sending the buffered body cannot double-bill. A 502 this + # proxy invented for a transport error or timeout is not + # retried — that request may already be running upstream. + if ( + e.from_upstream_response + and e.status_code in _RETRYABLE_UPSTREAM_5XX + and retries_left > 0 + ): + retries_left -= 1 + retry_index += 1 + logger.warning( + "Upstream %s returned %s for model=%s; retrying same " + "upstream (attempt %s, %s retries left)", + upstream.provider_type, + e.status_code, + model_id, + retry_index + 1, + retries_left, + extra={ + "provider": upstream.provider_type, + "model": model_id, + "status_code": e.status_code, + "path": path, + "retries_left": retries_left, + }, + ) + await asyncio.sleep( + _UPSTREAM_5XX_RETRY_BACKOFF_SECONDS * retry_index + ) + continue + # Let the outer UpstreamError handler manage failover/revert raise except Exception as e: # Unexpected error (not an upstream failure) — revert and propagate diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 488d07b1..f270442d 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -3151,6 +3151,7 @@ class BaseUpstreamProvider: status_code=response.status_code, code=rate_limit.code if rate_limit else None, details=rate_limit.as_details() if rate_limit else None, + from_upstream_response=True, ) try: @@ -3526,6 +3527,7 @@ class BaseUpstreamProvider: status_code=response.status_code, code=rate_limit.code if rate_limit else None, details=rate_limit.as_details() if rate_limit else None, + from_upstream_response=True, ) try: diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index 90b92719..0b077709 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -879,6 +879,7 @@ async def forward_ehbp_request( f"EHBP upstream {provider_type} returned {resp.status_code} " f"for model {model_obj.id}: {body_preview[:200] or ''}", status_code=resp.status_code, + from_upstream_response=True, ) # Check for usage metrics in response headers (non-streaming) or diff --git a/routstr/upstream/gemini_messages.py b/routstr/upstream/gemini_messages.py index 11440b91..d822f70f 100644 --- a/routstr/upstream/gemini_messages.py +++ b/routstr/upstream/gemini_messages.py @@ -348,6 +348,7 @@ async def _post_and_stream( raise UpstreamError( f"Upstream error via gemini compat: {body_text}", status_code=response.status_code, + from_upstream_response=True, ) return client, response diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py index c7e698e3..2856256c 100644 --- a/routstr/upstream/messages_dispatch.py +++ b/routstr/upstream/messages_dispatch.py @@ -567,6 +567,7 @@ async def dispatch_anthropic_messages( status_code=status_for_classify, code=rate_limit.code if rate_limit else None, details=rate_limit.as_details() if rate_limit else None, + from_upstream_response=True, ) from exc if not client_stream and hasattr(result, "__aiter__"): diff --git a/tests/integration/test_failover_billing.py b/tests/integration/test_failover_billing.py index d67dc32b..f82b3248 100644 --- a/tests/integration/test_failover_billing.py +++ b/tests/integration/test_failover_billing.py @@ -27,6 +27,12 @@ EXPENSIVE_BASE_URL = "https://expensive.example.com/v1" THIRD_BASE_URL = "https://third.example.com/v1" +@pytest.fixture(autouse=True) +def _no_upstream_5xx_retry_backoff(monkeypatch: pytest.MonkeyPatch) -> None: + """Keep the same-upstream retry backoff out of the test runtime.""" + monkeypatch.setattr("routstr.proxy._UPSTREAM_5XX_RETRY_BACKOFF_SECONDS", 0) + + def _make_model( model_id: str, prompt_sats: float, @@ -191,14 +197,15 @@ async def test_failover_serve_billed_at_serving_providers_rate( assert response.status_code == 200 payload = response.json() - # Both providers were attempted, cheapest first. + # The winner is tried twice: its 502 is retried in place before failover. assert [r.url.host for r in sent_requests] == [ + "cheap.example.com", "cheap.example.com", "expensive.example.com", ] # The fallback must be asked for ITS OWN model spelling, not the winner's. - forwarded_body = json.loads(sent_requests[1].content) + forwarded_body = json.loads(sent_requests[2].content) assert forwarded_body["model"] == "provb/dual-model" # The response echo names the model that actually served. @@ -319,7 +326,9 @@ async def test_same_id_failover_settles_at_serving_price( ) assert response.status_code == 200 + # The winner's 502 is retried in place before failover. assert [r.url.host for r in sent_requests] == [ + "cheap.example.com", "cheap.example.com", "expensive.example.com", ] @@ -459,7 +468,9 @@ async def test_usd_cost_serve_carries_serving_providers_fee( ) assert response.status_code == 200 + # The winner's 502 is retried in place before failover. assert [r.url.host for r in sent_requests] == [ + "cheap.example.com", "cheap.example.com", "expensive.example.com", ] @@ -530,7 +541,11 @@ async def test_failover_beyond_balance_envelope_is_rejected( # The 20_000-sat envelope exceeds the key's 10_000-sat balance: the # fallback must be rejected before its upstream is ever contacted. assert response.status_code == 402 - assert [r.url.host for r in sent_requests] == ["cheap.example.com"] + # The winner is retried in place; the fallback is still never contacted. + assert [r.url.host for r in sent_requests] == [ + "cheap.example.com", + "cheap.example.com", + ] @pytest.fixture async def raised_envelope_provider_maps( patched_db_engine: None, @@ -593,7 +608,9 @@ async def test_failover_reserves_serving_candidates_envelope( ) assert response.status_code == 200 + # The winner's 502 is retried in place before failover. assert [r.url.host for r in sent_requests] == [ + "cheap.example.com", "cheap.example.com", "expensive.example.com", ] @@ -610,3 +627,109 @@ async def test_failover_reserves_serving_candidates_envelope( charged = next(record for record in records if record.status == "charged") assert charged.reserved_msats > released.reserved_msats assert all(record.status != "active" for record in records) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_transient_502_retries_same_upstream_before_failing_over( + authenticated_client: AsyncClient, + dual_provider_maps: tuple[_StaticProvider, _StaticProvider], +) -> None: + """A transient 502 is retried on the same upstream, not failed over at once. + + The winner answers the first attempt with a 502 and the second with a + completion, so the pricier fallback is never contacted and the request is + billed at the winner's rate (2_000 msats, not the fallback's 10_000). + """ + sent_requests: list[httpx.Request] = [] + cheap_attempts = 0 + + async def fake_transport( + request: httpx.Request, *args: Any, **kwargs: Any + ) -> httpx.Response: + nonlocal cheap_attempts + sent_requests.append(request) + if request.url.host == "cheap.example.com": + cheap_attempts += 1 + if cheap_attempts == 1: + return httpx.Response( + 502, + content=json.dumps({"error": {"message": "bad gateway"}}).encode(), + headers={"content-type": "application/json"}, + ) + return _successful_upstream_response() + + with ( + patch( + "httpx.AsyncHTTPTransport.handle_async_request", + side_effect=fake_transport, + ), + patch( + "routstr.payment.cost_calculation.sats_usd_price", + return_value=0.0005, + ), + ): + response = await authenticated_client.post( + "/v1/chat/completions", + json={ + "model": "dual-model", + "messages": [{"role": "user", "content": "hello"}], + }, + ) + + assert response.status_code == 200 + # Retried in place: same host twice, fallback never consulted. + assert [r.url.host for r in sent_requests] == [ + "cheap.example.com", + "cheap.example.com", + ] + payload = response.json() + assert payload["model"] == "prova/dual-model" + # Billed at the winner's rate, not the fallback's 10_000. + assert payload["cost"]["total_msats"] == 2_000 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_transport_failure_is_not_retried_in_place( + authenticated_client: AsyncClient, + dual_provider_maps: tuple[_StaticProvider, _StaticProvider], +) -> None: + """A transport failure fails over at once instead of retrying in place. + + The proxy maps a connect/timeout error to a 502 of its own, so the upstream + may already have accepted and billed the request: re-sending it is not safe. + """ + sent_requests: list[httpx.Request] = [] + + async def fake_transport( + request: httpx.Request, *args: Any, **kwargs: Any + ) -> httpx.Response: + sent_requests.append(request) + if request.url.host == "cheap.example.com": + raise httpx.ConnectError("connection refused", request=request) + return _successful_upstream_response() + + with ( + patch( + "httpx.AsyncHTTPTransport.handle_async_request", + side_effect=fake_transport, + ), + patch( + "routstr.payment.cost_calculation.sats_usd_price", + return_value=0.0005, + ), + ): + response = await authenticated_client.post( + "/v1/chat/completions", + json={ + "model": "dual-model", + "messages": [{"role": "user", "content": "hello"}], + }, + ) + + assert response.status_code == 200 + assert [r.url.host for r in sent_requests] == [ + "cheap.example.com", + "expensive.example.com", + ]