add x-cashu support

This commit is contained in:
9qeklajc
2026-09-08 00:54:52 +02:00
parent f32565e254
commit 043b082f8b
4 changed files with 732 additions and 9 deletions
+128 -3
View File
@@ -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),
+15 -2
View File
@@ -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)
+52 -4
View File
@@ -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()
+537
View File
@@ -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()