diff --git a/routstr/proxy.py b/routstr/proxy.py index 38e2dfdc..c585188f 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -1015,7 +1015,7 @@ async def _proxy( already_stripped.add(bad_param) logger.warning( "Upstream %s rejected param '%s' for model=%s; " - "stripping and retrying same upstream", + "correcting and retrying same upstream", upstream.provider_type, bad_param, model_id, diff --git a/routstr/upstream/openai.py b/routstr/upstream/openai.py index f2f03cc5..b8d7c923 100644 --- a/routstr/upstream/openai.py +++ b/routstr/upstream/openai.py @@ -1,11 +1,23 @@ +import json from typing import TYPE_CHECKING +from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config +from litellm.llms.openai.chat.o_series_transformation import OpenAIOSeriesConfig + from ..payment.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider if TYPE_CHECKING: from ..core.db import UpstreamProviderRow +_O_SERIES = OpenAIOSeriesConfig() + + +def _rejects_max_tokens(model: str) -> bool: + return OpenAIGPT5Config.is_model_gpt_5_model( + model + ) or _O_SERIES.is_model_o_series_model(model) + class OpenAIUpstreamProvider(BaseUpstreamProvider): """Upstream provider specifically configured for OpenAI API.""" @@ -42,6 +54,33 @@ class OpenAIUpstreamProvider(BaseUpstreamProvider): """Strip 'openai/' prefix for OpenAI API compatibility.""" return model_id.removeprefix("openai/") + def prepare_request_body( + self, + body: bytes | None, + model_obj: Model, + include_stream_usage: bool = False, + ) -> bytes | None: + body = super().prepare_request_body(body, model_obj, include_stream_usage) + if not body: + return body + try: + data = json.loads(body) + except ValueError: + return body + # Reasoning models 400 on max_tokens; renaming up front saves the + # reject-and-retry round trip. Names litellm doesn't know yet still + # fall through to request_correction's reactive rename. + if ( + isinstance(data, dict) + and "messages" in data + and "max_tokens" in data + and "max_completion_tokens" not in data + and _rejects_max_tokens(self.transform_model_name(model_obj.id)) + ): + data["max_completion_tokens"] = data.pop("max_tokens") + return json.dumps(data).encode() + return body + async def fetch_models(self) -> list[Model]: """Fetch OpenAI models from OpenRouter API filtered by openai source.""" models_data = await async_fetch_openrouter_models(source_filter="openai") diff --git a/routstr/upstream/request_correction.py b/routstr/upstream/request_correction.py index d959a015..6bb73632 100644 --- a/routstr/upstream/request_correction.py +++ b/routstr/upstream/request_correction.py @@ -43,6 +43,18 @@ _UNSUPPORTED_PARAM_RE = re.compile( ) +# Matches upstream error text that rejects a param and names its replacement, +# e.g. OpenAI's "Unsupported parameter: 'max_tokens' is not supported with this +# model. Use 'max_completion_tokens' instead." Both names must be quoted so a +# free-form hint like "use gpt-4 instead" never reads as a rename. +_RENAMED_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)\b" + r".*?\buse\s+[`'\"](?P[a-zA-Z_][a-zA-Z0-9_]*)[`'\"]\s+instead", + re.IGNORECASE | re.DOTALL, +) + + # 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. @@ -100,6 +112,51 @@ _SPEND_SHAPING_PARAMS = frozenset( } ) +# Output caps are interchangeable spellings of the same limit, so moving the +# value from one to another keeps the priced bound intact. +_OUTPUT_CAP_PARAMS = frozenset( + { + "max_tokens", + "max_completion_tokens", + "max_output_tokens", + "max_tokens_to_sample", + } +) + + +def rename_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None: + """Move a rejected top-level param to the name the upstream asked for. + + Returns ``(new_body, label)`` with the value carried over unchanged, or + ``None`` when the error names no replacement, the param is absent, or the + replacement is already set. + + A spend-shaping field is only renamed to another output cap: that keeps the + reservation's bound, whereas renaming into or out of any other spend-shaping + field could uncap or fan out the retry. + """ + match = _RENAMED_PARAM_RE.search(error_message) + if not match: + return None + param, replacement = match.group("param"), match.group("replacement") + if param == replacement or param not in body or replacement in body: + return None + param_spend = param.lower() in _SPEND_SHAPING_PARAMS + replacement_spend = replacement.lower() in _SPEND_SHAPING_PARAMS + if (param_spend or replacement_spend) and not ( + param.lower() in _OUTPUT_CAP_PARAMS + and replacement.lower() in _OUTPUT_CAP_PARAMS + ): + logger.warning( + "Upstream asked to rename '%s' to '%s'; refusing because it would " + "change the request's spend bound — surfacing the error", + param, + replacement, + ) + return None + new_body = {(replacement if k == param else k): v for k, v in body.items()} + return new_body, f"{param}->{replacement}" + def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None: """Drop a top-level param the upstream named as unsupported/deprecated. @@ -130,7 +187,12 @@ def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] # Ordered pipeline of correctors tried on each recoverable rejection. -DEFAULT_CORRECTORS: tuple[Corrector, ...] = (strip_unsupported_param,) +# Renaming runs first so a param with a named replacement keeps its value +# instead of being dropped. +DEFAULT_CORRECTORS: tuple[Corrector, ...] = ( + rename_unsupported_param, + strip_unsupported_param, +) def correct_request( diff --git a/tests/unit/test_model_path_routing.py b/tests/unit/test_model_path_routing.py index e2c3e52f..cb764289 100644 --- a/tests/unit/test_model_path_routing.py +++ b/tests/unit/test_model_path_routing.py @@ -706,6 +706,75 @@ async def test_pinned_recovery_preserves_routing_fields( fallback.forward_request.assert_not_awaited() +_OPENAI_MAX_TOKENS_ERROR = json.dumps( + { + "error": { + "message": "Unsupported parameter: 'max_tokens' is not supported " + "with this model. Use 'max_completion_tokens' instead.", + "type": "invalid_request_error", + "param": "max_tokens", + "code": "unsupported_parameter", + } + } +).encode() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("pinned", [False, True]) +async def test_rejected_max_tokens_is_renamed_and_retried_on_same_upstream( + pinned: bool, +) -> None: + selected, fallback = _make_upstream(1), _make_upstream(2) + selected.forward_request = AsyncMock( + side_effect=[ + MagicMock(status_code=400, body=_OPENAI_MAX_TOKENS_ERROR), + MagicMock(status_code=200, body=b"{}"), + ] + ) + headers = {"authorization": "Bearer key"} + if pinned: + headers["x-routstr-model-path"] = encode_model_path(selected.base_url, MODEL_ID) + request = _make_request( + headers, + json.dumps( + {"model": MODEL_ID, "max_tokens": 300, "messages": [], "stream": True} + ).encode(), + ) + + response = await _run_proxy( + request, [(MagicMock(), selected), (MagicMock(), fallback)] + ) + + assert response.status_code == 200 + assert selected.forward_request.await_count == 2 + before, after = [ + json.loads(call.args[3]) for call in selected.forward_request.await_args_list + ] + assert before["max_tokens"] == 300 and "max_completion_tokens" not in before + assert after["max_completion_tokens"] == 300 and "max_tokens" not in after + assert {k: v for k, v in after.items() if k != "max_completion_tokens"} == { + k: v for k, v in before.items() if k != "max_tokens" + } + fallback.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_rename_that_changes_spend_bound_is_not_retried() -> None: + selected = _make_upstream(1, 400) + selected.forward_request.return_value.body = json.dumps( + {"error": {"message": "'max_tokens' is not supported. Use 'n' instead."}} + ).encode() + request = _make_request( + {"authorization": "Bearer key"}, + json.dumps({"model": MODEL_ID, "max_tokens": 300}).encode(), + ) + + response = await _run_proxy(request, [(MagicMock(), selected)]) + + assert response.status_code == 400 + selected.forward_request.assert_awaited_once() + + # --------------------------------------------------------------------------- # # Upstream 5xx -> 424 + UPSTREAM_UNAVAILABLE + scope header; node faults stay 500. # --------------------------------------------------------------------------- # diff --git a/tests/unit/test_openai_output_cap.py b/tests/unit/test_openai_output_cap.py new file mode 100644 index 00000000..84c7e234 --- /dev/null +++ b/tests/unit/test_openai_output_cap.py @@ -0,0 +1,90 @@ +"""OpenAI reasoning models get ``max_completion_tokens`` before the request is sent.""" + +from __future__ import annotations + +import json +import os + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") +os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to") + +import pytest + +from routstr.payment.models import Architecture, Model, Pricing +from routstr.upstream import GenericUpstreamProvider +from routstr.upstream.openai import OpenAIUpstreamProvider + + +def _model(model_id: str) -> Model: + return Model( + id=model_id, + name="test", + created=0, + description="", + context_length=128000, + architecture=Architecture( + modality="text->text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="x", + instruct_type=None, + ), + pricing=Pricing(prompt=0.0, completion=0.0), + ) + + +def _chat(model_id: str, **fields: object) -> bytes: + return json.dumps( + {"model": model_id, "messages": [{"role": "user", "content": "hi"}], **fields} + ).encode() + + +def _prepare(provider: object, model_id: str, body: bytes) -> dict: + out = provider.prepare_request_body(body, _model(model_id)) # type: ignore[attr-defined] + assert out is not None + return json.loads(out) + + +@pytest.mark.parametrize( + "model_id", ["gpt-5.6-sol", "openai/gpt-6-sol", "openai/gpt-5", "o3", "o4-mini"] +) +def test_reasoning_model_max_tokens_is_renamed(model_id: str) -> None: + provider = OpenAIUpstreamProvider(api_key="k") + data = _prepare(provider, model_id, _chat(model_id, max_tokens=300)) + assert data["max_completion_tokens"] == 300 + assert "max_tokens" not in data + + +@pytest.mark.parametrize("model_id", ["gpt-4o", "openai/gpt-4.1"]) +def test_non_reasoning_model_keeps_max_tokens(model_id: str) -> None: + provider = OpenAIUpstreamProvider(api_key="k") + data = _prepare(provider, model_id, _chat(model_id, max_tokens=300)) + assert data["max_tokens"] == 300 + assert "max_completion_tokens" not in data + + +def test_both_caps_set_is_left_for_upstream() -> None: + provider = OpenAIUpstreamProvider(api_key="k") + data = _prepare( + provider, + "gpt-5.6-sol", + _chat("gpt-5.6-sol", max_tokens=300, max_completion_tokens=200), + ) + assert data["max_tokens"] == 300 + assert data["max_completion_tokens"] == 200 + + +def test_non_chat_body_is_untouched() -> None: + provider = OpenAIUpstreamProvider(api_key="k") + body = json.dumps({"model": "gpt-5.6-sol", "input": "hi", "max_tokens": 5}).encode() + data = _prepare(provider, "gpt-5.6-sol", body) + assert data["max_tokens"] == 5 + assert "max_completion_tokens" not in data + + +def test_other_upstreams_keep_max_tokens() -> None: + provider = GenericUpstreamProvider(base_url="http://test", api_key="k") + data = _prepare(provider, "gpt-5.6-sol", _chat("gpt-5.6-sol", max_tokens=300)) + assert data["max_tokens"] == 300 + assert "max_completion_tokens" not in data diff --git a/tests/unit/test_request_correction.py b/tests/unit/test_request_correction.py index 3b1110a2..56a3423c 100644 --- a/tests/unit/test_request_correction.py +++ b/tests/unit/test_request_correction.py @@ -16,9 +16,15 @@ from routstr.upstream.request_correction import ( Correction, correct_request, extract_error_message, + rename_unsupported_param, strip_unsupported_param, ) +OPENAI_MAX_TOKENS_ERROR = ( + "Unsupported parameter: 'max_tokens' is not supported with this model. " + "Use 'max_completion_tokens' instead." +) + def _body(**kwargs: object) -> bytes: return json.dumps(kwargs).encode() @@ -139,6 +145,190 @@ class TestStripUnsupportedParam: assert strip_unsupported_param(body, "`Max_Tokens` is deprecated") is None +class TestRenameUnsupportedParam: + def test_renames_max_tokens_for_openai_reasoning_models(self) -> None: + body = {"model": "gpt-5.6-sol", "max_tokens": 256, "messages": []} + result = rename_unsupported_param(body, OPENAI_MAX_TOKENS_ERROR) + assert result is not None + new_body, label = result + assert label == "max_tokens->max_completion_tokens" + assert new_body == { + "model": "gpt-5.6-sol", + "max_completion_tokens": 256, + "messages": [], + } + + def test_preserves_key_order(self) -> None: + body = {"model": "m", "max_tokens": 1, "stream": True} + result = rename_unsupported_param(body, OPENAI_MAX_TOKENS_ERROR) + assert result is not None + assert list(result[0]) == ["model", "max_completion_tokens", "stream"] + + def test_does_not_mutate_input(self) -> None: + body = {"model": "m", "max_tokens": 8} + assert rename_unsupported_param(body, OPENAI_MAX_TOKENS_ERROR) is not None + assert body == {"model": "m", "max_tokens": 8} + + def test_renames_between_any_output_caps(self) -> None: + caps = ( + "max_tokens", + "max_completion_tokens", + "max_output_tokens", + "max_tokens_to_sample", + ) + for param in caps: + for replacement in caps: + if param == replacement: + continue + message = f"`{param}` is deprecated. Use `{replacement}` instead." + result = rename_unsupported_param({param: 7}, message) + assert result == ({replacement: 7}, f"{param}->{replacement}"), ( + param, + replacement, + ) + + def test_renames_non_spend_param(self) -> None: + message = "'functions' is deprecated. Use 'tools' instead." + result = rename_unsupported_param({"functions": [{"name": "f"}]}, message) + assert result == ({"tools": [{"name": "f"}]}, "functions->tools") + + def test_matches_across_quote_styles_case_and_newlines(self) -> None: + for message in ( + 'Unsupported parameter: "max_tokens" is not supported.\nUse ' + '"max_completion_tokens" instead.', + "`max_tokens` IS UNSUPPORTED here; please USE `max_completion_tokens`" + " INSTEAD", + "'max_tokens' is no longer supported, use 'max_completion_tokens' instead", + ): + result = rename_unsupported_param({"max_tokens": 3}, message) + assert result is not None, message + assert result[0] == {"max_completion_tokens": 3} + + def test_refuses_renames_that_change_the_spend_bound(self) -> None: + for param, replacement in ( + ("max_tokens", "n"), + ("n", "best_of"), + ("best_of", "n"), + ("temperature", "max_tokens"), + ("max_tokens", "temperature"), + ("n", "max_tokens"), + ): + message = f"'{param}' is not supported. Use '{replacement}' instead." + assert rename_unsupported_param({param: 2}, message) is None, ( + param, + replacement, + ) + + def test_spend_guard_is_case_insensitive(self) -> None: + ok = "'Max_Tokens' is not supported. Use 'MAX_COMPLETION_TOKENS' instead." + assert rename_unsupported_param({"Max_Tokens": 4}, ok) == ( + {"MAX_COMPLETION_TOKENS": 4}, + "Max_Tokens->MAX_COMPLETION_TOKENS", + ) + bad = "'Max_Tokens' is not supported. Use 'N' instead." + assert rename_unsupported_param({"Max_Tokens": 4}, bad) is None + + def test_declines_when_replacement_already_present(self) -> None: + body = {"max_tokens": 4, "max_completion_tokens": 8} + assert rename_unsupported_param(body, OPENAI_MAX_TOKENS_ERROR) is None + + def test_declines_when_param_absent(self) -> None: + assert rename_unsupported_param({"model": "m"}, OPENAI_MAX_TOKENS_ERROR) is None + + def test_declines_self_rename(self) -> None: + message = "'max_tokens' is deprecated. Use 'max_tokens' instead." + assert rename_unsupported_param({"max_tokens": 1}, message) is None + + def test_declines_unquoted_or_missing_replacement(self) -> None: + for message in ( + "`gpt-3` is deprecated, use gpt-4 instead", + "'max_tokens' is not supported, use max_completion_tokens instead", + "'max_tokens' is not supported with this model.", + "Use 'max_completion_tokens' instead.", + ): + assert rename_unsupported_param({"max_tokens": 1}, message) is None, message + + def test_declines_nested_only_param(self) -> None: + body = {"reasoning": {"max_tokens": 5}} + assert rename_unsupported_param(body, OPENAI_MAX_TOKENS_ERROR) is None + + +class TestCorrectRequestRename: + def test_openai_max_tokens_error_is_renamed_not_refused(self) -> None: + body = _body(model="gpt-5.6-sol", max_tokens=512, messages=[]) + result = correct_request(body, OPENAI_MAX_TOKENS_ERROR, set()) + assert isinstance(result, Correction) + assert result.label == "max_tokens->max_completion_tokens" + decoded = json.loads(result.body) + assert "max_tokens" not in decoded + assert decoded["max_completion_tokens"] == 512 + + def test_rename_wins_over_strip_for_non_spend_param(self) -> None: + body = _body(model="m", functions=[1]) + result = correct_request( + body, "'functions' is deprecated. Use 'tools' instead.", set() + ) + assert result is not None + assert json.loads(result.body) == {"model": "m", "tools": [1]} + + def test_unsafe_rename_of_cap_still_surfaces_error(self) -> None: + body = _body(model="m", max_tokens=5) + assert ( + correct_request( + body, "'max_tokens' is not supported. Use 'n' instead.", set() + ) + is None + ) + + def test_applied_rename_does_not_repeat_or_strip_cap(self) -> None: + body = _body(model="m", max_tokens=5) + applied = {"max_tokens->max_completion_tokens"} + assert correct_request(body, OPENAI_MAX_TOKENS_ERROR, applied) is None + + def test_rename_ping_pong_terminates(self) -> None: + """An upstream that flip-flops between names cannot loop forever.""" + forward = OPENAI_MAX_TOKENS_ERROR + backward = "'max_completion_tokens' is not supported. Use 'max_tokens' instead." + body = _body(model="m", max_tokens=5) + applied: set[str] = set() + for attempt in range(10): + message = forward if attempt % 2 == 0 else backward + result = correct_request(body, message, applied) + if result is None: + break + body, applied = result.body, applied | {result.label} + else: + raise AssertionError("correction loop did not terminate") + assert applied == { + "max_tokens->max_completion_tokens", + "max_completion_tokens->max_tokens", + } + assert json.loads(body) == {"model": "m", "max_tokens": 5} + + def test_buffered_openai_error_response_is_renamed(self) -> None: + resp = Response( + content=json.dumps( + { + "error": { + "message": OPENAI_MAX_TOKENS_ERROR, + "type": "invalid_request_error", + "param": "max_tokens", + "code": "unsupported_parameter", + } + } + ).encode(), + status_code=400, + ) + body = _body(model="gpt-5.6-sol", max_tokens=64, stream=True) + result = correct_request(body, extract_error_message(resp), set()) + assert result is not None + assert json.loads(result.body) == { + "model": "gpt-5.6-sol", + "max_completion_tokens": 64, + "stream": True, + } + + class TestExtractErrorMessage: def test_extracts_nested_error_message(self) -> None: resp = Response(