diff --git a/routstr/proxy.py b/routstr/proxy.py index 5feff279..c21f71b9 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -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), diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 67a10531..4f19ff92 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -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) diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 95b6edf2..f5fdd9c4 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -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() diff --git a/tests/unit/test_model_path_routing.py b/tests/unit/test_model_path_routing.py new file mode 100644 index 00000000..fdd60e76 --- /dev/null +++ b/tests/unit/test_model_path_routing.py @@ -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()