From 76013998f3a676eac61d841681a18a48ca54bd13 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 25 Sep 2026 23:32:41 +0200 Subject: [PATCH] fix: keep upstream status and error code consistent across 424 mapping --- routstr/core/exceptions.py | 5 +- routstr/proxy.py | 33 +++++++++-- routstr/upstream/base.py | 37 ++++++++---- routstr/upstream/ehbp.py | 4 +- tests/unit/test_upstream_error_response.py | 65 ++++++++++++++++++++++ 5 files changed, 128 insertions(+), 16 deletions(-) diff --git a/routstr/core/exceptions.py b/routstr/core/exceptions.py index 6827284e..d82bb190 100644 --- a/routstr/core/exceptions.py +++ b/routstr/core/exceptions.py @@ -28,7 +28,10 @@ class UpstreamError(Exception): ``scope`` is ``"upstream"`` for provider failures, reported to the caller as ``424`` (see :mod:`routstr.core.error_scope`), or ``"node"`` for local faults, which keep their status. ``status_code`` stays the provider's own - status; the caller-visible mapping happens at response construction. + status whenever ``from_upstream_response`` is True; the caller-visible + mapping happens at response construction. Proxy-chosen statuses (transport + failure, timeout) may already be the caller-visible one — only read + ``status_code`` as a provider status behind ``from_upstream_response``. """ def __init__( diff --git a/routstr/proxy.py b/routstr/proxy.py index 09617add..1b9a3947 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -24,6 +24,11 @@ from .core.db import ( create_session, get_session, ) +from .core.error_scope import ( + ERROR_SCOPE_UPSTREAM, + UPSTREAM_ERROR_STATUS, + UPSTREAM_UNAVAILABLE, +) from .core.exceptions import UpstreamError from .core.not_found import build_not_found_response from .core.settings import settings @@ -504,7 +509,12 @@ async def _proxy( last_error_response = create_upstream_error_response(e, request) continue return last_error_response or create_error_response( - "upstream_error", "All upstreams failed", 502, request=request + "upstream_error", + "All upstreams failed", + UPSTREAM_ERROR_STATUS, + request=request, + code=UPSTREAM_UNAVAILABLE, + error_scope=ERROR_SCOPE_UPSTREAM, ) selector: ModelPathSelector | None = None @@ -683,7 +693,12 @@ async def _proxy( if last_error is not None: return create_upstream_error_response(last_error, request) return create_error_response( - "upstream_error", "All upstreams failed", 502, request=request + "upstream_error", + "All upstreams failed", + UPSTREAM_ERROR_STATUS, + request=request, + code=UPSTREAM_UNAVAILABLE, + error_scope=ERROR_SCOPE_UPSTREAM, ) elif auth := headers.get("authorization", None): @@ -742,7 +757,12 @@ async def _proxy( last_error_response = create_upstream_error_response(e, request) continue return last_error_response or create_error_response( - "upstream_error", "All upstreams failed", 502, request=request + "upstream_error", + "All upstreams failed", + UPSTREAM_ERROR_STATUS, + request=request, + code=UPSTREAM_UNAVAILABLE, + error_scope=ERROR_SCOPE_UPSTREAM, ) reservation_snapshot: ReservationSnapshot | None = None @@ -1037,7 +1057,12 @@ async def _proxy( # Should not be reached given logic above return create_error_response( - "upstream_error", "All upstreams failed", 502, request=request + "upstream_error", + "All upstreams failed", + UPSTREAM_ERROR_STATUS, + request=request, + code=UPSTREAM_UNAVAILABLE, + error_scope=ERROR_SCOPE_UPSTREAM, ) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index efebc1c6..01c52a03 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -1040,6 +1040,21 @@ class BaseUpstreamProvider: err["code"] = client_code err["upstream_status"] = status_code redacted_body = json.dumps(parsed).encode() + elif ( + client_status != status_code + and isinstance(parsed, dict) + and "error" not in parsed + ): + # JSON body without an ``error`` mapping (e.g. FastAPI's + # ``{"detail": ...}``). Add one so a rewritten status is + # never served without its classification. + parsed["error"] = { + "message": message or "Upstream returned an error response", + "type": "upstream_error", + "code": client_code, + "upstream_status": status_code, + } + redacted_body = json.dumps(parsed).encode() except (ValueError, AttributeError): pass return Response( @@ -4536,8 +4551,10 @@ class BaseUpstreamProvider: "error": { "message": "Error forwarding request to upstream", "type": "upstream_error", + # Pass the status as the code so a provider + # 4xx keeps the legacy numeric ``code``. "code": client_code_for_upstream_error( - response.status_code, None + response.status_code, response.status_code ), "upstream_status": response.status_code, "refund_token": refund_token, @@ -4701,14 +4718,13 @@ class BaseUpstreamProvider: # be reported as a retryable redemption error (see handle_x_cashu). if redeemed: upstream_status = getattr(e, "status_code", None) + upstream_code = getattr(e, "code", None) return create_error_response( "upstream_error", "Payment succeeded but the upstream request failed", - client_status_for_upstream_error(upstream_status), + client_status_for_upstream_error(upstream_status, upstream_code), request=request, - code=client_code_for_upstream_error( - upstream_status, getattr(e, "code", None) - ), + code=client_code_for_upstream_error(upstream_status, upstream_code), details=upstream_status_details(None, upstream_status), error_scope=ERROR_SCOPE_UPSTREAM, ) @@ -4844,8 +4860,10 @@ class BaseUpstreamProvider: "error": { "message": "Error forwarding Responses API request to upstream", "type": "upstream_error", + # Pass the status as the code so a provider + # 4xx keeps the legacy numeric ``code``. "code": client_code_for_upstream_error( - response.status_code, None + response.status_code, response.status_code ), "upstream_status": response.status_code, "refund_token": refund_token, @@ -5463,14 +5481,13 @@ class BaseUpstreamProvider: # bait). Redemption classification only applies while not redeemed. if redeemed: upstream_status = getattr(e, "status_code", None) + upstream_code = getattr(e, "code", None) return create_error_response( "upstream_error", "Payment succeeded but the upstream request failed", - client_status_for_upstream_error(upstream_status), + client_status_for_upstream_error(upstream_status, upstream_code), request=request, - code=client_code_for_upstream_error( - upstream_status, getattr(e, "code", None) - ), + code=client_code_for_upstream_error(upstream_status, upstream_code), details=upstream_status_details(None, upstream_status), error_scope=ERROR_SCOPE_UPSTREAM, ) diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index bb8f97e4..673416f6 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -1140,8 +1140,10 @@ async def forward_ehbp_x_cashu_request( "error": { "message": "Error forwarding EHBP request to upstream", "type": "upstream_error", + # Pass the status as the code so a provider 4xx + # keeps the legacy numeric ``code``. "code": client_code_for_upstream_error( - resp.status_code, None + resp.status_code, resp.status_code ), "upstream_status": resp.status_code, "refund_token": refund_token, diff --git a/tests/unit/test_upstream_error_response.py b/tests/unit/test_upstream_error_response.py index c053fc90..ff2be223 100644 --- a/tests/unit/test_upstream_error_response.py +++ b/tests/unit/test_upstream_error_response.py @@ -20,6 +20,8 @@ from routstr.core.error_scope import ( ERROR_SCOPE_UPSTREAM, UPSTREAM_ERROR_STATUS, UPSTREAM_UNAVAILABLE, + client_code_for_upstream_error, + client_status_for_upstream_error, ) from routstr.core.exceptions import UpstreamError from routstr.payment.helpers import create_upstream_error_response @@ -309,3 +311,66 @@ def test_node_scoped_failure_stays_500_without_scope_header() -> None: def test_upstream_error_defaults_to_upstream_scope() -> None: assert UpstreamError("boom").scope == ERROR_SCOPE_UPSTREAM + + +@pytest.mark.asyncio +async def test_json_body_without_error_mapping_gets_classification( + provider: BaseUpstreamProvider, +) -> None: + """A rewritten status is never served without a matching ``error.code``.""" + body = json.dumps({"detail": "internal failure"}).encode() + upstream = _make_upstream_response( + body=body, status_code=503, content_type="application/json" + ) + + response = await provider.forward_upstream_error_response( + _make_request(), "v1/chat/completions", upstream + ) + + assert response.status_code == UPSTREAM_ERROR_STATUS + assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM + payload: dict[str, Any] = json.loads(bytes(response.body)) + assert payload["detail"] == "internal failure" + assert payload["error"]["code"] == UPSTREAM_UNAVAILABLE + assert payload["error"]["upstream_status"] == 503 + + +@pytest.mark.asyncio +async def test_json_body_with_non_mapping_error_is_left_alone( + provider: BaseUpstreamProvider, +) -> None: + """A provider's own ``error`` value is never clobbered by the mapping.""" + body = json.dumps({"error": "boom"}).encode() + upstream = _make_upstream_response( + body=body, status_code=503, content_type="application/json" + ) + + response = await provider.forward_upstream_error_response( + _make_request(), "v1/chat/completions", upstream + ) + + assert response.status_code == UPSTREAM_ERROR_STATUS + assert response.headers[ERROR_SCOPE_HEADER] == ERROR_SCOPE_UPSTREAM + assert json.loads(bytes(response.body)) == {"error": "boom"} + + +@pytest.mark.parametrize("upstream_status", [429, 500, 502, 503, 529]) +def test_rate_limit_status_and_code_never_disagree(upstream_status: int) -> None: + """429 and ``UPSTREAM_RATE_LIMIT`` are one classification, not two: a caller + must never see ``424`` carrying the rate-limit code.""" + assert client_status_for_upstream_error(upstream_status, UPSTREAM_RATE_LIMIT) == 429 + assert ( + client_code_for_upstream_error(upstream_status, UPSTREAM_RATE_LIMIT) + == UPSTREAM_RATE_LIMIT + ) + + +@pytest.mark.parametrize("upstream_status", [400, 401, 403, 404, 422]) +def test_provider_4xx_keeps_its_numeric_code(upstream_status: int) -> None: + """The x-cashu envelopes pass the status as the code; a 4xx must keep the + legacy numeric ``error.code`` rather than degrade to ``null``.""" + assert client_status_for_upstream_error(upstream_status) == upstream_status + assert ( + client_code_for_upstream_error(upstream_status, upstream_status) + == upstream_status + )