This commit is contained in:
9qeklajc
2026-09-08 01:02:16 +02:00
parent 043b082f8b
commit 33482b6f0b
4 changed files with 120 additions and 21 deletions
+11 -8
View File
@@ -2,7 +2,6 @@ import asyncio
import inspect
import json
from typing import Any
from urllib.parse import urlsplit
from fastapi import APIRouter, Depends, HTTPException, Request
from fastapi.responses import Response, StreamingResponse
@@ -42,6 +41,7 @@ from .upstream.helpers import init_upstreams
from .upstream.model_paths import (
ModelPathSelector,
decode_model_path,
is_openrouter_base_url,
public_model_id,
public_provider_url,
)
@@ -558,7 +558,7 @@ async def _proxy(
if (
is_ehbp
or not request_body_dict
or urlsplit(pinned[1].base_url).hostname != "openrouter.ai"
or not is_openrouter_base_url(pinned[1].base_url)
or _canonical_api_path(path)
not in {"chat/completions", "completions", "responses"}
):
@@ -848,17 +848,20 @@ async def _proxy(
)
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 and not is_ehbp and selector is None:
# Same-provider recovery must not relax an explicit route.
if response.status_code == 400 and not is_ehbp:
correction = correct_request(
request_body,
extract_error_message(response),
already_stripped,
)
if correction is not None and selector is not None:
corrected_body = json.loads(correction.body)
if any(
corrected_body.get(field) != request_body_dict.get(field)
for field in ("model", "provider")
):
correction = None
if correction is not None:
request_body, bad_param = correction.body, correction.label
already_stripped.add(bad_param)
+5 -9
View File
@@ -186,15 +186,11 @@ def _make_http_client() -> httpx.AsyncClient:
def is_openrouter_base_url(base_url: str | None) -> bool:
"""True when ``base_url`` points at OpenRouter.
Deliberately separate from ``BaseUpstreamProvider._upstream_accepts_cache_control``:
that predicate also returns True for native Anthropic (correct for
cache-control, wrong for OpenRouter endpoint discovery). This one keys only
on the URL so a ``GenericUpstreamProvider`` aimed at OpenRouter is matched
while native Anthropic is not.
"""
return "openrouter.ai" in (base_url or "")
"""Match OpenRouter itself, not compatible providers or lookalike hosts."""
try:
return urlsplit(base_url or "").hostname == "openrouter.ai"
except ValueError:
return False
def exposed_model_id(model: object) -> str:
+88
View File
@@ -535,3 +535,91 @@ async def test_model_fallback_list_is_rejected_when_pinned() -> None:
response = await _run_proxy(request, [(MagicMock(), selected)])
assert response.status_code == 400
selected.forward_request.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", [None, "deepinfra/fp8"])
@pytest.mark.parametrize("path", ["v1/chat/completions", "v1/responses"])
@pytest.mark.parametrize("final_status", [200, 429])
async def test_pinned_recovery_stays_on_selected_provider(
endpoint: str | None, path: str, final_status: int
) -> None:
selected, fallback = _make_upstream(1), _make_upstream(2)
selected.base_url = "https://openrouter.ai/api/v1"
handler = (
"forward_responses_request" if path == "v1/responses" else "forward_request"
)
forward = AsyncMock(
side_effect=[
MagicMock(
status_code=400,
body=b'{"error":{"message":"temperature is deprecated"}}',
),
MagicMock(status_code=final_status, body=b"{}"),
]
)
setattr(selected, handler, forward)
setattr(fallback, handler, AsyncMock())
request = _make_request(
{
"authorization": "Bearer key",
"x-routstr-model-path": encode_model_path(
selected.base_url, 1, MODEL_ID, endpoint
),
},
json.dumps(
{
"model": MODEL_ID,
"temperature": 0.7,
"provider": {"data_collection": "deny"},
}
).encode(),
)
response = await _run_proxy(
request, [(MagicMock(), selected), (MagicMock(), fallback)], path
)
assert response.status_code == final_status
assert forward.await_count == 2
before, after = [json.loads(call.args[3]) for call in forward.await_args_list]
assert "temperature" in before
assert "temperature" not in after
assert after["model"] == before["model"] == MODEL_ID
assert after["provider"] == before["provider"]
if endpoint:
assert after["provider"]["order"] == [endpoint]
assert after["provider"]["allow_fallbacks"] is False
getattr(fallback, handler).assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("field", ["model", "provider"])
@pytest.mark.parametrize("endpoint", [None, "deepinfra/fp8"])
async def test_pinned_recovery_preserves_routing_fields(
field: str, endpoint: str | None
) -> None:
selected, fallback = _make_upstream(1, 400), _make_upstream(2)
selected.base_url = "https://openrouter.ai/api/v1"
selected.forward_request.return_value.body = json.dumps(
{"error": {"message": f"{field} is not supported"}}
).encode()
request = _make_request(
{
"authorization": "Bearer key",
"x-routstr-model-path": encode_model_path(
selected.base_url, 1, MODEL_ID, endpoint
),
},
json.dumps(
{"model": MODEL_ID, "provider": {"data_collection": "deny"}}
).encode(),
)
response = await _run_proxy(
request, [(MagicMock(), selected), (MagicMock(), fallback)]
)
assert response.status_code == 400
selected.forward_request.assert_awaited_once()
fallback.forward_request.assert_not_awaited()
+16 -4
View File
@@ -252,10 +252,22 @@ def _path_entry(
# --------------------------------------------------------------------------- #
def test_is_openrouter_base_url() -> None:
assert mp.is_openrouter_base_url("https://openrouter.ai/api/v1") is True
assert mp.is_openrouter_base_url("https://api.anthropic.com") is False
assert mp.is_openrouter_base_url(None) is False
@pytest.mark.parametrize(
"url, expected",
[
("https://openrouter.ai/api/v1", True),
("https://OPENROUTER.AI/api/v1", True),
("https://api.anthropic.com", False),
("https://openrouter.ai.evil.test/api/v1", False),
("https://evil.test/openrouter.ai", False),
("https://openrouter.ai@evil.test/api/v1", False),
("https://[invalid", False),
("", False),
(None, False),
],
)
def test_is_openrouter_base_url(url: str | None, expected: bool) -> None:
assert mp.is_openrouter_base_url(url) is expected
def test_native_anthropic_not_openrouter() -> None: