From 06c8a071a580776824a9162b108f6980bcf5eba3 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 7 Jun 2026 13:00:19 +0200 Subject: [PATCH] feat: reactive request-correction retry for recoverable upstream 400s --- routstr/proxy.py | 111 ++++++++++++------ routstr/upstream/request_correction.py | 142 ++++++++++++++++++++++ tests/unit/test_request_correction.py | 156 +++++++++++++++++++++++++ 3 files changed, 372 insertions(+), 37 deletions(-) create mode 100644 routstr/upstream/request_correction.py create mode 100644 tests/unit/test_request_correction.py diff --git a/routstr/proxy.py b/routstr/proxy.py index dcb32409..373b80b2 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -28,6 +28,7 @@ from .payment.helpers import ( from .payment.models import Model from .upstream import BaseUpstreamProvider from .upstream.helpers import init_upstreams +from .upstream.request_correction import correct_request, extract_error_message logger = get_logger(__name__) proxy_router = APIRouter() @@ -352,50 +353,86 @@ async def proxy( if request_body_dict: await pay_for_request(key, max_cost_for_model, session) + # Tracks request params already removed in response to upstream rejections, + # shared across providers so a stripped param stays stripped on failover and + # the reactive retry can never loop unboundedly. + already_stripped: set[str] = set() + for i, upstream in enumerate(upstreams): headers = upstream.prepare_headers(dict(request.headers)) try: - try: - if is_responses_api: - response = await upstream.forward_responses_request( - request, - path, - headers, - request_body, - key, - max_cost_for_model, - session, - model_obj, + while True: + try: + if is_responses_api: + response = await upstream.forward_responses_request( + request, + path, + headers, + request_body, + key, + max_cost_for_model, + session, + model_obj, + ) + else: + response = await upstream.forward_request( + request, + path, + headers, + request_body, + key, + max_cost_for_model, + session, + model_obj, + ) + except UpstreamError: + # Let the outer UpstreamError handler manage retry/revert + raise + except Exception as e: + # Unexpected error (not an upstream failure) — revert and propagate + logger.error( + "Unexpected error in upstream request, reverting payment", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "path": path, + "key_hash": key.hashed_key[:8] + "...", + "max_cost_for_model": max_cost_for_model, + }, ) - else: - response = await upstream.forward_request( - request, - path, - headers, + await revert_pay_for_request(key, session, max_cost_for_model) + raise + + # Reactive recovery: some models reject one specific request + # param (e.g. newer Anthropic models deprecating `temperature`). + # When the upstream 400s naming such a param, strip it from the + # body and retry the SAME upstream. ``already_stripped`` bounds + # this to one retry per distinct param so it always terminates. + if response.status_code == 400: + correction = correct_request( request_body, - key, - max_cost_for_model, - session, - model_obj, + extract_error_message(response), + already_stripped, ) - except UpstreamError: - # Let the outer UpstreamError handler manage retry/revert - raise - except Exception as e: - # Unexpected error (not an upstream failure) — revert and propagate - logger.error( - "Unexpected error in upstream request, reverting payment", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "path": path, - "key_hash": key.hashed_key[:8] + "...", - "max_cost_for_model": max_cost_for_model, - }, - ) - await revert_pay_for_request(key, session, max_cost_for_model) - raise + if correction is not None: + request_body, bad_param = correction.body, correction.label + already_stripped.add(bad_param) + logger.warning( + "Upstream %s rejected param '%s' for model=%s; " + "stripping and retrying same upstream", + upstream.provider_type, + bad_param, + model_id, + extra={ + "provider": upstream.provider_type, + "model": model_id, + "stripped_param": bad_param, + "path": path, + }, + ) + continue + break if response.status_code != 200: # Check if we should retry (502 Upstream Error or 429 Rate Limit) diff --git a/routstr/upstream/request_correction.py b/routstr/upstream/request_correction.py new file mode 100644 index 00000000..c2ea5b1d --- /dev/null +++ b/routstr/upstream/request_correction.py @@ -0,0 +1,142 @@ +"""Reactive request-correction layer. + +When an upstream rejects a request with a recoverable 4xx error, this layer +tries to *fix* the request body and let the caller retry the same upstream +instead of failing outright. It is provider-agnostic: correctors key off the +upstream's own error wording, so the same recovery works across every provider. + +The layer is a small pipeline of :data:`Corrector` callables. Each corrector +inspects the parsed request body and the upstream error message and either +returns a corrected body (plus a short label identifying the fix) or declines +by returning ``None``. Adding a new reactive fix means writing one corrector +and adding it to :data:`DEFAULT_CORRECTORS` — no changes to the proxy loop. + +All corrections are immutable: a corrector never mutates the body it is given, +it returns a new ``dict``. The proxy threads an ``applied`` set of fix labels +through retries so each distinct fix is applied at most once, guaranteeing the +retry loop always terminates. +""" + +from __future__ import annotations + +import json +import re +from collections.abc import Callable, Sequence +from dataclasses import dataclass + +from fastapi.responses import Response + +from ..core import get_logger + +logger = get_logger(__name__) + + +# Matches upstream error text that names a single rejected request parameter, +# e.g. "`temperature` is deprecated for this model." or +# "parameter 'top_p' is not supported". Keys off the upstream's own wording so +# a 400 about an unsupported sampling/option field can be recovered by stripping +# that field and retrying the same upstream. +_UNSUPPORTED_PARAM_RE = re.compile( + r"[`'\"]?(?P[a-zA-Z_][a-zA-Z0-9_]*)[`'\"]?\s+is\s+" + r"(?:deprecated|not\s+supported|unsupported|no\s+longer\s+supported)", + re.IGNORECASE, +) + + +# A corrector inspects the parsed request body and the upstream error message +# and returns ``(new_body_dict, label)`` for a fix it can apply, or ``None`` to +# decline. ``label`` identifies the fix so it is applied at most once per request. +Corrector = Callable[[dict, str], "tuple[dict, str] | None"] + + +@dataclass(frozen=True) +class Correction: + """A successful request correction ready to retry. + + ``body`` is the corrected JSON body (encoded), ``label`` identifies the fix + that was applied (e.g. the stripped param name) so the caller can guard + against applying the same fix twice. + """ + + body: bytes + label: str + + +def extract_error_message(response: Response) -> str: + """Best-effort extraction of an error message string from a proxy Response.""" + body_bytes = getattr(response, "body", None) + if not body_bytes: + return "" + try: + data = json.loads(body_bytes) + except Exception: + return body_bytes.decode("utf-8", errors="ignore")[:500] + if isinstance(data, dict): + err = data.get("error") + if isinstance(err, dict): + msg = err.get("message") or err.get("detail") + if isinstance(msg, str): + return msg + elif isinstance(err, str): + return err + if isinstance(data.get("message"), str): + return data["message"] + return "" + + +def strip_unsupported_param( + body: dict, error_message: str +) -> tuple[dict, str] | None: + """Drop a top-level param the upstream named as unsupported/deprecated. + + Returns ``(new_body, param)`` (a new dict, original untouched) when the + error names a top-level param present in the body, otherwise ``None``. + """ + match = _UNSUPPORTED_PARAM_RE.search(error_message) + if not match: + return None + param = match.group("param") + if param not in body: + return None + new_body = {k: v for k, v in body.items() if k != param} + return new_body, param + + +# Ordered pipeline of correctors tried on each recoverable rejection. +DEFAULT_CORRECTORS: tuple[Corrector, ...] = (strip_unsupported_param,) + + +def correct_request( + request_body: bytes, + error_message: str, + applied: set[str], + correctors: Sequence[Corrector] = DEFAULT_CORRECTORS, +) -> Correction | None: + """Try to correct a rejected request body so it can be retried. + + Runs each corrector in order against the parsed body and ``error_message``. + The first corrector that proposes a fix whose ``label`` is not already in + ``applied`` wins; its result is returned as a :class:`Correction`. Returns + ``None`` when nothing parses, nothing matches, or every proposed fix was + already applied — the caller then treats the response as a normal failure. + + ``applied`` is read-only here; the caller records the returned ``label`` to + bound retries and guarantee forward progress. + """ + if not request_body or not error_message: + return None + try: + data = json.loads(request_body) + except Exception: + return None + if not isinstance(data, dict): + return None + for corrector in correctors: + result = corrector(data, error_message) + if result is None: + continue + new_body, label = result + if label in applied: + continue + return Correction(body=json.dumps(new_body).encode(), label=label) + return None diff --git a/tests/unit/test_request_correction.py b/tests/unit/test_request_correction.py new file mode 100644 index 00000000..903fe16e --- /dev/null +++ b/tests/unit/test_request_correction.py @@ -0,0 +1,156 @@ +"""Unit tests for the reactive request-correction layer. + +Covers the recovery path that lets a request survive a 400 where the upstream +names a single unsupported request param (e.g. newer Anthropic models +deprecating ``temperature``): the param is stripped from the JSON body and the +same upstream is retried, provider-agnostically, keyed off the error text. +""" + +from __future__ import annotations + +import json + +from fastapi.responses import Response + +from routstr.upstream.request_correction import ( + Correction, + correct_request, + extract_error_message, + strip_unsupported_param, +) + + +def _body(**kwargs: object) -> bytes: + return json.dumps(kwargs).encode() + + +class TestCorrectRequest: + def test_strips_deprecated_temperature(self) -> None: + body = _body(model="claude-opus-4-8", temperature=1, messages=[]) + result = correct_request( + body, "`temperature` is deprecated for this model.", set() + ) + assert isinstance(result, Correction) + assert result.label == "temperature" + decoded = json.loads(result.body) + assert "temperature" not in decoded + assert decoded["model"] == "claude-opus-4-8" + + def test_strips_not_supported_param(self) -> None: + body = _body(model="m", top_p=0.9, messages=[]) + result = correct_request(body, "Parameter 'top_p' is not supported", set()) + assert result is not None + assert result.label == "top_p" + assert "top_p" not in json.loads(result.body) + + def test_returns_none_when_label_already_applied(self) -> None: + body = _body(model="m", temperature=1) + assert ( + correct_request(body, "`temperature` is deprecated", {"temperature"}) + is None + ) + + def test_returns_none_when_param_absent_from_body(self) -> None: + body = _body(model="m", messages=[]) + assert correct_request(body, "`temperature` is deprecated", set()) is None + + def test_returns_none_when_message_does_not_match(self) -> None: + body = _body(model="m", temperature=1) + assert correct_request(body, "Insufficient balance", set()) is None + + def test_returns_none_on_empty_inputs(self) -> None: + assert correct_request(b"", "`temperature` is deprecated", set()) is None + assert correct_request(_body(temperature=1), "", set()) is None + + def test_returns_none_on_non_object_body(self) -> None: + assert correct_request(b"[1, 2, 3]", "`temperature` is deprecated", set()) is None + + def test_deprecated_model_name_is_not_stripped_as_param(self) -> None: + """A 'model is deprecated' error must not strip an unrelated body field. + + The regex matches the `` is deprecated`` wording, but the + ``param not in body`` guard means a deprecated *model* name (not a + request param) yields no correction rather than a false strip. + """ + body = _body(model="gpt-3", temperature=1, messages=[]) + assert correct_request(body, "`gpt-3` is deprecated, use gpt-4", set()) is None + + def test_streaming_400_buffered_error_is_correctable(self) -> None: + """Streaming 400s funnel through a buffered JSON Response, so the same + correction path applies as for non-streaming requests.""" + # Mirrors forward_upstream_error_response's buffered JSON envelope. + resp = Response( + content=json.dumps( + {"error": {"message": "`temperature` is deprecated for this model"}} + ).encode(), + status_code=400, + ) + body = _body(model="claude-opus-4-8", temperature=1, messages=[]) + result = correct_request(body, extract_error_message(resp), set()) + assert isinstance(result, Correction) + assert result.label == "temperature" + assert "temperature" not in json.loads(result.body) + + +class TestStripUnsupportedParam: + def test_does_not_mutate_input(self) -> None: + body = {"model": "m", "temperature": 1} + result = strip_unsupported_param(body, "`temperature` is deprecated") + assert result is not None + new_body, param = result + assert param == "temperature" + assert "temperature" not in new_body + # original untouched (immutability) + assert body == {"model": "m", "temperature": 1} + + def test_declines_when_no_match(self) -> None: + assert strip_unsupported_param({"temperature": 1}, "nope") is None + + +class TestExtractErrorMessage: + def test_extracts_nested_error_message(self) -> None: + resp = Response( + content=json.dumps( + {"error": {"message": "`temperature` is deprecated", "type": "x"}} + ).encode(), + status_code=400, + ) + assert extract_error_message(resp) == "`temperature` is deprecated" + + def test_extracts_string_error(self) -> None: + resp = Response( + content=json.dumps({"error": "bad request"}).encode(), status_code=400 + ) + assert extract_error_message(resp) == "bad request" + + def test_extracts_top_level_message(self) -> None: + resp = Response( + content=json.dumps({"message": "nope"}).encode(), status_code=400 + ) + assert extract_error_message(resp) == "nope" + + def test_empty_body_returns_empty_string(self) -> None: + assert extract_error_message(Response(status_code=400)) == "" + + def test_non_json_body_returns_preview(self) -> None: + resp = Response(content=b"plain text error", status_code=400) + assert extract_error_message(resp) == "plain text error" + + +class TestEndToEndChaining: + def test_two_distinct_params_corrected_sequentially(self) -> None: + """Simulates the proxy loop: each 400 fixes one param, set guards reuse.""" + body = _body(model="m", temperature=1, top_p=0.5, messages=[]) + applied: set[str] = set() + + first = correct_request(body, "`temperature` is deprecated", applied) + assert first is not None + body, applied = first.body, applied | {first.label} + + second = correct_request(body, "`top_p` is not supported", applied) + assert second is not None + body, applied = second.body, applied | {second.label} + + decoded = json.loads(body) + assert "temperature" not in decoded and "top_p" not in decoded + assert applied == {"temperature", "top_p"}