Merge pull request #765 from Routstr/fix/upstream-5xx-same-upstream-retry

fix(upstream): retry a transient 5xx on the same upstream before failing over
This commit is contained in:
9qeklajc
2026-09-23 21:57:28 +02:00
committed by GitHub
8 changed files with 183 additions and 5 deletions
+7
View File
@@ -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)
+4
View File
@@ -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")
+41 -2
View File
@@ -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
+2
View File
@@ -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:
+1
View File
@@ -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 '<empty>'}",
status_code=resp.status_code,
from_upstream_response=True,
)
# Check for usage metrics in response headers (non-streaming) or
+1
View File
@@ -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
+1
View File
@@ -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__"):
+126 -3
View File
@@ -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",
]