mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
add x-cashu support
This commit is contained in:
+128
-3
@@ -2,6 +2,7 @@ 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
|
||||
@@ -38,9 +39,17 @@ 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,
|
||||
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)
|
||||
@@ -424,6 +452,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(
|
||||
@@ -464,6 +499,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:
|
||||
@@ -471,6 +541,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 urlsplit(pinned[1].base_url).hostname != "openrouter.ai"
|
||||
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)
|
||||
@@ -521,11 +636,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(
|
||||
@@ -728,7 +853,7 @@ async def _proxy(
|
||||
# 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:
|
||||
if response.status_code == 400 and not is_ehbp and selector is None:
|
||||
correction = correct_request(
|
||||
request_body,
|
||||
extract_error_message(response),
|
||||
|
||||
@@ -533,6 +533,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)
|
||||
@@ -4244,6 +4245,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.
|
||||
|
||||
@@ -4263,7 +4266,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")
|
||||
@@ -4459,6 +4463,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.
|
||||
|
||||
@@ -4523,6 +4529,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
mint,
|
||||
request_body=request_body,
|
||||
)
|
||||
except Exception as e:
|
||||
error_message = str(e)
|
||||
@@ -4579,6 +4586,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.
|
||||
|
||||
@@ -4600,7 +4609,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(
|
||||
@@ -5162,6 +5172,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.
|
||||
|
||||
@@ -5241,6 +5253,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,6 +129,57 @@ 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()
|
||||
|
||||
@@ -0,0 +1,537 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user