mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
Merge pull request #720 from Routstr/x-routstr-model-path
X routstr model path
This commit is contained in:
+135
-7
@@ -38,9 +38,18 @@ from .payment.models import Model
|
||||
from .upstream import BaseUpstreamProvider
|
||||
from .upstream.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request
|
||||
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,
|
||||
)
|
||||
from .upstream.request_correction import correct_request, extract_error_message
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
MODEL_PATH_HEADER = "x-routstr-model-path"
|
||||
proxy_router = APIRouter()
|
||||
|
||||
_upstreams: list[BaseUpstreamProvider] = []
|
||||
@@ -113,6 +122,25 @@ def get_candidates(
|
||||
return None
|
||||
|
||||
|
||||
def _model_ids_match(requested: str, selected: str) -> bool:
|
||||
if requested.lower() == selected.lower():
|
||||
return True
|
||||
return public_model_id(requested).lower() == public_model_id(selected).lower()
|
||||
|
||||
|
||||
def _candidate_for_selector(
|
||||
selector: ModelPathSelector,
|
||||
candidates: list[tuple[Model, BaseUpstreamProvider]],
|
||||
) -> tuple[Model, BaseUpstreamProvider] | None:
|
||||
for model_obj, upstream in candidates:
|
||||
if (
|
||||
upstream.db_id == selector.provider_id
|
||||
and public_provider_url(upstream.base_url) == selector.base_url
|
||||
):
|
||||
return model_obj, upstream
|
||||
return None
|
||||
|
||||
|
||||
def get_model_instance(model_id: str) -> Model | None:
|
||||
"""Get the best-ranked Model instance for a model ID."""
|
||||
candidates = get_candidates(model_id)
|
||||
@@ -416,6 +444,13 @@ async def _proxy(
|
||||
# without model/cost/auth lookups. Do not prefix-match here: paths such as
|
||||
# /attestationjunk must continue through normal authentication.
|
||||
if request.method == "GET" and _is_tinfoil_attestation_path(path):
|
||||
if MODEL_PATH_HEADER in headers:
|
||||
return create_error_response(
|
||||
"unsupported_request",
|
||||
"Model paths do not apply to attestation",
|
||||
400,
|
||||
request=request,
|
||||
)
|
||||
selected_upstreams = _select_unauthenticated_get_upstreams(path, _upstreams)
|
||||
if not selected_upstreams:
|
||||
return create_error_response(
|
||||
@@ -456,6 +491,41 @@ async def _proxy(
|
||||
"upstream_error", "All upstreams failed", 502, request=request
|
||||
)
|
||||
|
||||
selector: ModelPathSelector | None = None
|
||||
if MODEL_PATH_HEADER in headers:
|
||||
selector = decode_model_path(headers[MODEL_PATH_HEADER])
|
||||
if (
|
||||
selector is None
|
||||
or sum(
|
||||
name.lower() == MODEL_PATH_HEADER for name, _ in request.headers.items()
|
||||
)
|
||||
!= 1
|
||||
):
|
||||
return create_error_response(
|
||||
"invalid_request",
|
||||
f"Malformed {MODEL_PATH_HEADER} header",
|
||||
400,
|
||||
request=request,
|
||||
)
|
||||
if not isinstance(model_id, str) or not _model_ids_match(
|
||||
model_id, selector.model_id
|
||||
):
|
||||
return create_error_response(
|
||||
"invalid_request",
|
||||
f"{MODEL_PATH_HEADER} selects model '{selector.model_id}' but the "
|
||||
f"request asks for '{model_id}'",
|
||||
400,
|
||||
request=request,
|
||||
)
|
||||
if "models" in request_body_dict:
|
||||
return create_error_response(
|
||||
"invalid_request",
|
||||
"Model paths cannot be combined with model fallbacks",
|
||||
400,
|
||||
request=request,
|
||||
)
|
||||
model_id = selector.model_id
|
||||
|
||||
candidates = get_candidates(model_id)
|
||||
|
||||
if not candidates:
|
||||
@@ -463,6 +533,51 @@ async def _proxy(
|
||||
"invalid_model", f"Model '{model_id}' not found", 400, request=request
|
||||
)
|
||||
|
||||
if selector is not None:
|
||||
pinned = _candidate_for_selector(selector, candidates)
|
||||
if pinned is None:
|
||||
return create_error_response(
|
||||
"invalid_model_path",
|
||||
f"Model '{selector.model_id}' is not routable through provider "
|
||||
f"{selector.provider_id}",
|
||||
404,
|
||||
request=request,
|
||||
)
|
||||
# Explicit routes must never enter cross-provider failover.
|
||||
candidates = [pinned]
|
||||
|
||||
if selector.endpoint_tag:
|
||||
if (
|
||||
is_ehbp
|
||||
or not request_body_dict
|
||||
or not is_openrouter_base_url(pinned[1].base_url)
|
||||
or _canonical_api_path(path)
|
||||
not in {"chat/completions", "completions", "responses"}
|
||||
):
|
||||
return create_error_response(
|
||||
"unsupported_request",
|
||||
"Endpoint pinning requires an OpenRouter completion or Responses JSON request",
|
||||
400,
|
||||
request=request,
|
||||
)
|
||||
provider_options = request_body_dict.get("provider", {})
|
||||
if not isinstance(provider_options, dict):
|
||||
return create_error_response(
|
||||
"invalid_request",
|
||||
"provider must be an object",
|
||||
400,
|
||||
request=request,
|
||||
)
|
||||
request_body_dict = {
|
||||
**request_body_dict,
|
||||
"provider": {
|
||||
**provider_options,
|
||||
"order": [selector.endpoint_tag],
|
||||
"allow_fallbacks": False,
|
||||
},
|
||||
}
|
||||
request_body = json.dumps(request_body_dict).encode()
|
||||
|
||||
if is_ehbp:
|
||||
candidates = [
|
||||
(model, upstream)
|
||||
@@ -513,11 +628,21 @@ async def _proxy(
|
||||
)
|
||||
elif is_responses_api:
|
||||
return await upstream.handle_x_cashu_responses(
|
||||
request, x_cashu, path, max_cost_for_model, model_obj
|
||||
request,
|
||||
x_cashu,
|
||||
path,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
request_body=request_body,
|
||||
)
|
||||
else:
|
||||
return await upstream.handle_x_cashu(
|
||||
request, x_cashu, path, max_cost_for_model, model_obj
|
||||
request,
|
||||
x_cashu,
|
||||
path,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
request_body=request_body,
|
||||
)
|
||||
except UpstreamError as e:
|
||||
logger.warning(
|
||||
@@ -715,17 +840,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.
|
||||
# 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)
|
||||
|
||||
@@ -591,6 +591,7 @@ class BaseUpstreamProvider:
|
||||
"refund-lnurl",
|
||||
"key-expiry-time",
|
||||
"x-cashu",
|
||||
"x-routstr-model-path",
|
||||
]:
|
||||
if headers.pop(header, None) is not None:
|
||||
removed_headers.append(header)
|
||||
@@ -4336,6 +4337,8 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model: int,
|
||||
model_obj: Model,
|
||||
mint: str | None = None,
|
||||
*,
|
||||
request_body: bytes | None = None,
|
||||
) -> Response | StreamingResponse:
|
||||
"""Forward request paid with X-Cashu token to upstream service.
|
||||
|
||||
@@ -4355,7 +4358,8 @@ class BaseUpstreamProvider:
|
||||
if path.startswith("v1/"):
|
||||
path = path.replace("v1/", "")
|
||||
|
||||
request_body = await request.body()
|
||||
if request_body is None:
|
||||
request_body = await request.body()
|
||||
|
||||
if (
|
||||
path.endswith("messages/count_tokens")
|
||||
@@ -4552,6 +4556,8 @@ class BaseUpstreamProvider:
|
||||
path: str,
|
||||
max_cost_for_model: int,
|
||||
model_obj: Model,
|
||||
*,
|
||||
request_body: bytes | None = None,
|
||||
) -> Response | StreamingResponse:
|
||||
"""Handle X-Cashu payment for Responses API requests.
|
||||
|
||||
@@ -4616,6 +4622,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
mint,
|
||||
request_body=request_body,
|
||||
)
|
||||
except Exception as e:
|
||||
error_message = str(e)
|
||||
@@ -4672,6 +4679,8 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model: int,
|
||||
model_obj: Model,
|
||||
mint: str | None = None,
|
||||
*,
|
||||
request_body: bytes | None = None,
|
||||
) -> Response | StreamingResponse:
|
||||
"""Forward Responses API request paid with X-Cashu token to upstream service.
|
||||
|
||||
@@ -4693,7 +4702,8 @@ class BaseUpstreamProvider:
|
||||
|
||||
url = f"{self.base_url}/{path}"
|
||||
|
||||
request_body = await request.body()
|
||||
if request_body is None:
|
||||
request_body = await request.body()
|
||||
transformed_body = self.prepare_responses_request_body(request_body, model_obj)
|
||||
|
||||
logger.debug(
|
||||
@@ -5281,6 +5291,8 @@ class BaseUpstreamProvider:
|
||||
path: str,
|
||||
max_cost_for_model: int,
|
||||
model_obj: Model,
|
||||
*,
|
||||
request_body: bytes | None = None,
|
||||
) -> Response | StreamingResponse:
|
||||
"""Handle request with X-Cashu token payment, redeeming token and forwarding request.
|
||||
|
||||
@@ -5360,6 +5372,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
mint,
|
||||
request_body=request_body,
|
||||
)
|
||||
except Exception as e:
|
||||
error_message = str(e)
|
||||
|
||||
@@ -1,9 +1,6 @@
|
||||
"""Model-path discovery service.
|
||||
|
||||
Exposes every selectable upstream route a Routstr model is reachable through.
|
||||
This PR remains discovery-only: request-side routing will consume the opaque
|
||||
selectors in a follow-up.
|
||||
|
||||
A path is a standard percent-encoded query string containing the configured
|
||||
upstream URL, provider ID, client-visible model ID and, for an exact OpenRouter
|
||||
endpoint, its machine-readable tag::
|
||||
@@ -20,7 +17,7 @@ import random
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Callable
|
||||
from urllib.parse import urlencode, urlsplit
|
||||
from urllib.parse import parse_qsl, urlencode, urlsplit
|
||||
|
||||
import httpx
|
||||
from sqlalchemy.dialects.sqlite import insert
|
||||
@@ -132,21 +129,68 @@ def encode_model_path(
|
||||
return urlencode(components)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelPathSelector:
|
||||
"""Decoded client-supplied route selector."""
|
||||
|
||||
base_url: str
|
||||
provider_id: int
|
||||
model_id: str
|
||||
endpoint_tag: str | None = None
|
||||
|
||||
|
||||
def decode_model_path(path: str) -> ModelPathSelector | None:
|
||||
"""Inverse of ``encode_model_path``; ``None`` when the selector is malformed."""
|
||||
try:
|
||||
pairs = parse_qsl(
|
||||
path,
|
||||
keep_blank_values=True,
|
||||
strict_parsing=True,
|
||||
max_num_fields=4,
|
||||
errors="strict",
|
||||
)
|
||||
except ValueError:
|
||||
return None
|
||||
params = dict(pairs)
|
||||
if len(params) != len(pairs) or params.keys() - {
|
||||
"url",
|
||||
"provider-id",
|
||||
"model-id",
|
||||
"endpoint",
|
||||
}:
|
||||
return None
|
||||
if any(not value.strip() for value in params.values()):
|
||||
return None
|
||||
base_url = params.get("url", "")
|
||||
model_id = params.get("model-id", "")
|
||||
raw_provider_id = params.get("provider-id", "")
|
||||
if not base_url or not model_id or not raw_provider_id:
|
||||
return None
|
||||
try:
|
||||
provider_id = int(raw_provider_id)
|
||||
except ValueError:
|
||||
return None
|
||||
if provider_id <= 0:
|
||||
return None
|
||||
return ModelPathSelector(
|
||||
base_url=base_url,
|
||||
provider_id=provider_id,
|
||||
model_id=model_id,
|
||||
endpoint_tag=params.get("endpoint") or None,
|
||||
)
|
||||
|
||||
|
||||
def _make_http_client() -> httpx.AsyncClient:
|
||||
"""Client factory, separated so tests can substitute a mock transport."""
|
||||
return 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:
|
||||
|
||||
@@ -0,0 +1,625 @@
|
||||
"""Model-path routing and fail-closed behavior."""
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr import proxy as proxy_module
|
||||
from routstr.auth import ReservationSnapshot
|
||||
from routstr.core.db import ApiKey
|
||||
from routstr.upstream.model_paths import decode_model_path, encode_model_path
|
||||
|
||||
MODEL_ID = "test-model"
|
||||
|
||||
|
||||
def _make_upstream(db_id: int, status_code: int = 200) -> MagicMock:
|
||||
upstream = MagicMock()
|
||||
upstream.db_id = db_id
|
||||
upstream.base_url = "http://localhost"
|
||||
upstream.provider_type = f"provider-{db_id}"
|
||||
upstream.supports_ehbp = False
|
||||
upstream.prepare_headers = MagicMock(side_effect=lambda h: h)
|
||||
upstream.on_upstream_error_redirect = AsyncMock()
|
||||
upstream.forward_request = AsyncMock(
|
||||
return_value=MagicMock(status_code=status_code, body=b"{}")
|
||||
)
|
||||
return upstream
|
||||
|
||||
|
||||
def _make_request(headers: dict[str, str], body: bytes) -> MagicMock:
|
||||
request = MagicMock()
|
||||
request.method = "POST"
|
||||
request.headers = headers
|
||||
request.body = AsyncMock(return_value=body)
|
||||
request.state = MagicMock()
|
||||
request.state.request_id = "req-model-path"
|
||||
return request
|
||||
|
||||
|
||||
async def _run_proxy(
|
||||
request: MagicMock,
|
||||
candidates: list[tuple[Any, Any]],
|
||||
path: str = "v1/chat/completions",
|
||||
) -> Any:
|
||||
key = ApiKey(hashed_key="mpkey", balance=10_000)
|
||||
reservation = ReservationSnapshot(
|
||||
release_id="model-path-release",
|
||||
key_hash=key.hashed_key,
|
||||
billing_key_hash=key.hashed_key,
|
||||
reserved_msats=1_000,
|
||||
)
|
||||
with (
|
||||
patch.object(proxy_module, "get_candidates", return_value=candidates),
|
||||
patch.object(
|
||||
proxy_module, "get_max_cost_for_model", AsyncMock(return_value=1_000)
|
||||
),
|
||||
patch.object(
|
||||
proxy_module,
|
||||
"calculate_discounted_max_cost",
|
||||
AsyncMock(return_value=1_000),
|
||||
),
|
||||
patch.object(proxy_module, "check_token_balance", MagicMock()),
|
||||
patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)),
|
||||
patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)),
|
||||
patch.object(
|
||||
proxy_module,
|
||||
"get_reservation_snapshot",
|
||||
AsyncMock(return_value=reservation),
|
||||
),
|
||||
patch.object(proxy_module, "revert_pay_for_request", AsyncMock()),
|
||||
):
|
||||
return await proxy_module.proxy(request, path, session=MagicMock())
|
||||
|
||||
|
||||
def test_decode_model_path_round_trips_encode() -> None:
|
||||
selector = decode_model_path(
|
||||
encode_model_path("https://openrouter.ai/api/v1", 7, MODEL_ID, "deepinfra/fp8")
|
||||
)
|
||||
assert selector is not None
|
||||
assert selector.base_url == "https://openrouter.ai/api/v1"
|
||||
assert selector.provider_id == 7
|
||||
assert selector.model_id == MODEL_ID
|
||||
assert selector.endpoint_tag == "deepinfra/fp8"
|
||||
|
||||
|
||||
def test_decode_model_path_without_endpoint_has_no_tag() -> None:
|
||||
selector = decode_model_path(encode_model_path("http://localhost", 1, MODEL_ID))
|
||||
assert selector is not None
|
||||
assert selector.endpoint_tag is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"path",
|
||||
[
|
||||
"",
|
||||
"url=http://localhost&model-id=test-model",
|
||||
"url=http://localhost&provider-id=abc&model-id=test-model",
|
||||
"provider-id=1&model-id=test-model",
|
||||
"url=http://localhost&provider-id=1",
|
||||
],
|
||||
)
|
||||
def test_decode_model_path_rejects_malformed_selectors(path: str) -> None:
|
||||
assert decode_model_path(path) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_path_routes_to_the_selected_provider() -> None:
|
||||
first, selected = _make_upstream(1), _make_upstream(2)
|
||||
request = _make_request(
|
||||
{
|
||||
"authorization": "Bearer sk-mpkey",
|
||||
"x-routstr-model-path": encode_model_path("http://localhost", 2, MODEL_ID),
|
||||
},
|
||||
json.dumps({"model": MODEL_ID}).encode(),
|
||||
)
|
||||
|
||||
await _run_proxy(request, [(MagicMock(), first), (MagicMock(), selected)])
|
||||
|
||||
selected.forward_request.assert_awaited_once()
|
||||
first.forward_request.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("status_code", [400, 429, 500, 502, 503])
|
||||
async def test_model_path_failure_is_returned_without_falling_back(
|
||||
status_code: int,
|
||||
) -> None:
|
||||
selected = _make_upstream(1, status_code=status_code)
|
||||
fallback = _make_upstream(2)
|
||||
request = _make_request(
|
||||
{
|
||||
"authorization": "Bearer sk-mpkey",
|
||||
"x-routstr-model-path": encode_model_path("http://localhost", 1, MODEL_ID),
|
||||
},
|
||||
json.dumps({"model": MODEL_ID}).encode(),
|
||||
)
|
||||
|
||||
response = await _run_proxy(
|
||||
request, [(MagicMock(), selected), (MagicMock(), fallback)]
|
||||
)
|
||||
|
||||
assert response.status_code == status_code
|
||||
selected.forward_request.assert_awaited_once()
|
||||
fallback.forward_request.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unknown_provider_in_model_path_is_rejected() -> None:
|
||||
upstream = _make_upstream(1)
|
||||
request = _make_request(
|
||||
{
|
||||
"authorization": "Bearer sk-mpkey",
|
||||
"x-routstr-model-path": encode_model_path("http://localhost", 99, MODEL_ID),
|
||||
},
|
||||
json.dumps({"model": MODEL_ID}).encode(),
|
||||
)
|
||||
|
||||
response = await _run_proxy(request, [(MagicMock(), upstream)])
|
||||
|
||||
assert response.status_code == 404
|
||||
assert json.loads(bytes(response.body))["error"]["type"] == "invalid_model_path"
|
||||
upstream.forward_request.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_path_disagreeing_with_the_body_model_is_rejected() -> None:
|
||||
upstream = _make_upstream(1)
|
||||
request = _make_request(
|
||||
{
|
||||
"authorization": "Bearer sk-mpkey",
|
||||
"x-routstr-model-path": encode_model_path(
|
||||
"http://localhost", 1, "other-model"
|
||||
),
|
||||
},
|
||||
json.dumps({"model": MODEL_ID}).encode(),
|
||||
)
|
||||
|
||||
response = await _run_proxy(request, [(MagicMock(), upstream)])
|
||||
|
||||
assert response.status_code == 400
|
||||
upstream.forward_request.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_malformed_model_path_header_is_rejected() -> None:
|
||||
upstream = _make_upstream(1)
|
||||
request = _make_request(
|
||||
{
|
||||
"authorization": "Bearer sk-mpkey",
|
||||
"x-routstr-model-path": "not-a-model-path",
|
||||
},
|
||||
json.dumps({"model": MODEL_ID}).encode(),
|
||||
)
|
||||
|
||||
response = await _run_proxy(request, [(MagicMock(), upstream)])
|
||||
|
||||
assert response.status_code == 400
|
||||
upstream.forward_request.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_endpoint_tag_pins_the_upstream_subprovider() -> None:
|
||||
upstream = _make_upstream(1)
|
||||
upstream.base_url = "https://openrouter.ai/api/v1"
|
||||
request = _make_request(
|
||||
{
|
||||
"authorization": "Bearer sk-mpkey",
|
||||
"x-routstr-model-path": encode_model_path(
|
||||
"https://openrouter.ai/api/v1", 1, MODEL_ID, "deepinfra/fp8"
|
||||
),
|
||||
},
|
||||
json.dumps({"model": MODEL_ID}).encode(),
|
||||
)
|
||||
|
||||
await _run_proxy(request, [(MagicMock(), upstream)])
|
||||
|
||||
forwarded_body = upstream.forward_request.await_args.args[3]
|
||||
assert json.loads(forwarded_body)["provider"] == {
|
||||
"order": ["deepinfra/fp8"],
|
||||
"allow_fallbacks": False,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"raw",
|
||||
["", " ", "url=http://localhost&provider-id=1&model-id=test-model&provider-id=2"],
|
||||
)
|
||||
async def test_ambiguous_headers_do_not_route(raw: str) -> None:
|
||||
upstream = _make_upstream(1)
|
||||
request = _make_request(
|
||||
{"authorization": "Bearer key", "x-routstr-model-path": raw},
|
||||
json.dumps({"model": MODEL_ID}).encode(),
|
||||
)
|
||||
response = await _run_proxy(request, [(MagicMock(), upstream)])
|
||||
assert response.status_code == 400
|
||||
upstream.forward_request.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_selector_url_must_match_configured_provider() -> None:
|
||||
upstream = _make_upstream(1)
|
||||
request = _make_request(
|
||||
{
|
||||
"authorization": "Bearer key",
|
||||
"x-routstr-model-path": encode_model_path(
|
||||
"http://169.254.169.254", 1, MODEL_ID
|
||||
),
|
||||
},
|
||||
json.dumps({"model": MODEL_ID}).encode(),
|
||||
)
|
||||
response = await _run_proxy(request, [(MagicMock(), upstream)])
|
||||
assert response.status_code == 404
|
||||
upstream.forward_request.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_endpoint_pin_cannot_be_stripped_on_retry() -> None:
|
||||
upstream = _make_upstream(1, 400)
|
||||
upstream.base_url = "https://openrouter.ai/api/v1"
|
||||
upstream.forward_request.return_value.body = (
|
||||
b'{"error":{"message":"provider is not supported"}}'
|
||||
)
|
||||
request = _make_request(
|
||||
{
|
||||
"authorization": "Bearer key",
|
||||
"x-routstr-model-path": encode_model_path(
|
||||
upstream.base_url, 1, MODEL_ID, "deepinfra/fp8"
|
||||
),
|
||||
},
|
||||
json.dumps({"model": MODEL_ID}).encode(),
|
||||
)
|
||||
response = await _run_proxy(request, [(MagicMock(), upstream)])
|
||||
assert response.status_code == 400
|
||||
upstream.forward_request.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"path,handler",
|
||||
[
|
||||
("v1/chat/completions", "handle_x_cashu"),
|
||||
("v1/responses", "handle_x_cashu_responses"),
|
||||
],
|
||||
)
|
||||
async def test_cashu_receives_endpoint_pinned_body(path: str, handler: str) -> None:
|
||||
upstream = _make_upstream(1)
|
||||
upstream.base_url = "https://openrouter.ai/api/v1"
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
async def handle(request: Any, *args: Any, **kwargs: Any) -> Any:
|
||||
captured.update(json.loads(kwargs.get("request_body") or await request.body()))
|
||||
return MagicMock(status_code=200)
|
||||
|
||||
setattr(upstream, handler, AsyncMock(side_effect=handle))
|
||||
request = _make_request(
|
||||
{
|
||||
"x-cashu": "test-token",
|
||||
"x-routstr-model-path": encode_model_path(
|
||||
upstream.base_url, 1, MODEL_ID, "deepinfra/fp8"
|
||||
),
|
||||
},
|
||||
json.dumps(
|
||||
{
|
||||
"model": MODEL_ID,
|
||||
"provider": {
|
||||
"order": ["other"],
|
||||
"allow_fallbacks": True,
|
||||
"data_collection": "deny",
|
||||
},
|
||||
}
|
||||
).encode(),
|
||||
)
|
||||
await _run_proxy(request, [(MagicMock(), upstream)], path)
|
||||
assert captured["provider"] == {
|
||||
"order": ["deepinfra/fp8"],
|
||||
"allow_fallbacks": False,
|
||||
"data_collection": "deny",
|
||||
}
|
||||
|
||||
|
||||
def test_model_path_header_is_not_forwarded() -> None:
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
upstream = BaseUpstreamProvider("http://localhost", "upstream-key")
|
||||
headers = upstream.prepare_headers({"x-routstr-model-path": "private-routing-data"})
|
||||
assert "x-routstr-model-path" not in headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", ["v1/chat/completions", "v1/responses"])
|
||||
@pytest.mark.parametrize("status_code", [200, 429, 502])
|
||||
async def test_cashu_pin_reaches_http_transport(path: str, status_code: int) -> None:
|
||||
import httpx
|
||||
from fastapi.responses import Response
|
||||
|
||||
from routstr.upstream.openrouter import OpenRouterUpstreamProvider
|
||||
|
||||
upstream = OpenRouterUpstreamProvider(api_key="upstream-key")
|
||||
upstream.db_id = 1
|
||||
fallback = _make_upstream(2)
|
||||
fallback.handle_x_cashu = AsyncMock()
|
||||
fallback.handle_x_cashu_responses = AsyncMock()
|
||||
sent: list[httpx.Request] = []
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
sent.append(request)
|
||||
return httpx.Response(status_code, json={"error": "test"})
|
||||
|
||||
client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
|
||||
model = MagicMock(id=MODEL_ID, forwarded_model_id=None, canonical_slug=None)
|
||||
request = _make_request(
|
||||
{
|
||||
"x-cashu": "test-token",
|
||||
"x-routstr-model-path": encode_model_path(
|
||||
upstream.base_url, 1, MODEL_ID, "deepinfra/fp8"
|
||||
),
|
||||
},
|
||||
json.dumps({"model": MODEL_ID, "provider": {"allow_fallbacks": True}}).encode(),
|
||||
)
|
||||
request.query_params = {}
|
||||
with (
|
||||
patch("routstr.upstream.base.httpx.AsyncClient", return_value=client),
|
||||
patch(
|
||||
"routstr.upstream.base.recieve_token",
|
||||
AsyncMock(return_value=(1000, "msat", "https://mint.test")),
|
||||
) as redeem,
|
||||
patch("routstr.upstream.base.store_cashu_transaction", AsyncMock()),
|
||||
patch.object(upstream, "send_refund", AsyncMock(return_value="refund")),
|
||||
patch.object(
|
||||
upstream,
|
||||
"handle_x_cashu_chat_completion",
|
||||
AsyncMock(return_value=Response(status_code=200)),
|
||||
),
|
||||
patch.object(
|
||||
upstream,
|
||||
"handle_x_cashu_responses_completion",
|
||||
AsyncMock(return_value=Response(status_code=200)),
|
||||
),
|
||||
):
|
||||
response = await _run_proxy(
|
||||
request, [(model, upstream), (model, fallback)], path
|
||||
)
|
||||
assert response.status_code == status_code
|
||||
redeem.assert_awaited_once()
|
||||
assert len(sent) == 1
|
||||
assert sent[0].url.host == "openrouter.ai"
|
||||
assert json.loads(sent[0].content)["provider"] == {
|
||||
"order": ["deepinfra/fp8"],
|
||||
"allow_fallbacks": False,
|
||||
}
|
||||
assert "x-routstr-model-path" not in sent[0].headers
|
||||
fallback.handle_x_cashu.assert_not_awaited()
|
||||
fallback.handle_x_cashu_responses.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unpinned_requests_still_fall_back() -> None:
|
||||
first, fallback = _make_upstream(1, 502), _make_upstream(2)
|
||||
request = _make_request(
|
||||
{"authorization": "Bearer key"}, json.dumps({"model": MODEL_ID}).encode()
|
||||
)
|
||||
response = await _run_proxy(
|
||||
request, [(MagicMock(), first), (MagicMock(), fallback)]
|
||||
)
|
||||
assert response.status_code == 200
|
||||
first.forward_request.assert_awaited_once()
|
||||
fallback.forward_request.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pinned_exception_does_not_fall_back() -> None:
|
||||
from routstr.core.exceptions import UpstreamError
|
||||
|
||||
first, fallback = _make_upstream(1), _make_upstream(2)
|
||||
first.forward_request.side_effect = UpstreamError("unavailable", status_code=503)
|
||||
request = _make_request(
|
||||
{
|
||||
"authorization": "Bearer key",
|
||||
"x-routstr-model-path": encode_model_path(first.base_url, 1, MODEL_ID),
|
||||
},
|
||||
json.dumps({"model": MODEL_ID}).encode(),
|
||||
)
|
||||
response = await _run_proxy(
|
||||
request, [(MagicMock(), first), (MagicMock(), fallback)]
|
||||
)
|
||||
assert response.status_code == 503
|
||||
first.forward_request.assert_awaited_once()
|
||||
fallback.forward_request.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"path,is_ehbp", [("v1/messages", False), ("v1/chat/completions", True)]
|
||||
)
|
||||
async def test_unsupported_endpoint_pins_fail_before_payment(
|
||||
path: str, is_ehbp: bool
|
||||
) -> None:
|
||||
upstream = _make_upstream(1)
|
||||
upstream.base_url = "https://openrouter.ai/api/v1"
|
||||
headers = {
|
||||
"x-cashu": "test-token",
|
||||
"x-routstr-model-path": encode_model_path(
|
||||
upstream.base_url, 1, MODEL_ID, "deepinfra/fp8"
|
||||
),
|
||||
}
|
||||
if is_ehbp:
|
||||
headers.update({"ehbp-encapsulated-key": "sealed", "x-routstr-model": MODEL_ID})
|
||||
request = _make_request(headers, json.dumps({"model": MODEL_ID}).encode())
|
||||
with (
|
||||
patch.object(proxy_module, "check_token_balance") as payment,
|
||||
patch.object(
|
||||
proxy_module, "get_candidates", return_value=[(MagicMock(), upstream)]
|
||||
),
|
||||
):
|
||||
response = await proxy_module.proxy(request, path, MagicMock())
|
||||
assert response.status_code == 400
|
||||
assert json.loads(response.body)["error"]["type"] == "unsupported_request"
|
||||
payment.assert_not_called()
|
||||
upstream.forward_request.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("cashu", [False, True])
|
||||
async def test_ehbp_pin_does_not_fall_back(cashu: bool) -> None:
|
||||
from routstr.core.exceptions import UpstreamError
|
||||
|
||||
selected, fallback = _make_upstream(1), _make_upstream(2)
|
||||
selected.supports_ehbp = fallback.supports_ehbp = True
|
||||
headers = {
|
||||
"ehbp-encapsulated-key": "sealed",
|
||||
"x-routstr-model": MODEL_ID,
|
||||
"x-routstr-model-path": encode_model_path(selected.base_url, 1, MODEL_ID),
|
||||
}
|
||||
headers.update({"x-cashu": "token"} if cashu else {"authorization": "Bearer key"})
|
||||
request = _make_request(headers, b"encrypted-body")
|
||||
handler = "forward_ehbp_x_cashu_request" if cashu else "forward_ehbp_request"
|
||||
with patch.object(
|
||||
proxy_module,
|
||||
handler,
|
||||
AsyncMock(side_effect=UpstreamError("unavailable", status_code=503)),
|
||||
) as forward:
|
||||
response = await _run_proxy(
|
||||
request, [(MagicMock(), selected), (MagicMock(), fallback)]
|
||||
)
|
||||
assert response.status_code == 503
|
||||
forward.assert_awaited_once()
|
||||
assert forward.await_args is not None
|
||||
assert forward.await_args.kwargs["upstream"] is selected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_header_fields_are_rejected() -> None:
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
selected = _make_upstream(1)
|
||||
route = encode_model_path(selected.base_url, 1, MODEL_ID).encode()
|
||||
request = _make_request({}, json.dumps({"model": MODEL_ID}).encode())
|
||||
request.headers = Headers(
|
||||
raw=[
|
||||
(b"authorization", b"Bearer key"),
|
||||
(b"x-routstr-model-path", route),
|
||||
(b"x-routstr-model-path", route),
|
||||
]
|
||||
)
|
||||
response = await _run_proxy(request, [(MagicMock(), selected)])
|
||||
assert response.status_code == 400
|
||||
selected.forward_request.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attestation_does_not_ignore_model_path() -> None:
|
||||
request = _make_request(
|
||||
{"x-routstr-model-path": encode_model_path("http://localhost", 1, MODEL_ID)},
|
||||
b"",
|
||||
)
|
||||
request.method = "GET"
|
||||
with patch.object(proxy_module, "_select_unauthenticated_get_upstreams") as select:
|
||||
response = await _run_proxy(request, [], "attestation")
|
||||
assert response.status_code == 400
|
||||
select.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_fallback_list_is_rejected_when_pinned() -> None:
|
||||
selected = _make_upstream(1)
|
||||
request = _make_request(
|
||||
{
|
||||
"authorization": "Bearer key",
|
||||
"x-routstr-model-path": encode_model_path(selected.base_url, 1, MODEL_ID),
|
||||
},
|
||||
json.dumps({"model": MODEL_ID, "models": ["other-model"]}).encode(),
|
||||
)
|
||||
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()
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user