diff --git a/routstr/proxy.py b/routstr/proxy.py index 38e2dfdc..8107868f 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -46,6 +46,7 @@ from .upstream.model_paths import ( ModelPathSelector, decode_model_path, is_openrouter_base_url, + pinned_endpoint_context, public_model_id, public_provider_url, ) @@ -625,6 +626,9 @@ async def _proxy( ) model_id = selector.model_id + # Set for every request so an unpinned one never inherits a stale pin. + pinned_endpoint_context.set(selector.endpoint_tag if selector else None) + candidates = get_candidates(model_id) if not candidates: diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index a53ec0f0..a5ef1f77 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -23,6 +23,7 @@ import json import random import time from collections.abc import Iterable +from contextvars import ContextVar from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Callable from urllib.parse import parse_qsl, urlencode, urlsplit @@ -146,6 +147,14 @@ class ModelPathSelector: provider_id: int | None = None +# Endpoint tag the current request is pinned to, set by the proxy once the +# selector is resolved. Response stamping reads it to name the serving +# provider when the upstream omits it; ``None`` for unpinned requests. +pinned_endpoint_context: ContextVar[str | None] = ContextVar( + "pinned_endpoint_tag", default=None +) + + def decode_model_path(path: str) -> ModelPathSelector | None: """Inverse of ``encode_model_path``; ``None`` when the selector is malformed.""" try: diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index 3ff13d90..33516c9b 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -5,7 +5,7 @@ import httpx from ..core.logging import get_logger from ..payment.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider, _reported_provider -from .model_paths import public_provider_url +from .model_paths import pinned_endpoint_context, public_provider_url if TYPE_CHECKING: from ..core.db import UpstreamProviderRow @@ -41,6 +41,8 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): - Real upstream sub-provider (e.g. ``"GMICloud"``) -> ``"openrouter:GMICloud"``. - Missing sub-provider, or one that merely echoes ``"openrouter"`` -> + the endpoint the request was pinned to via the model path + (``"openrouter:deepinfra/fp8"``) when there is one, else ``"openrouter:unknown"``: the router is still known even when the serving provider is not (e.g. the Responses API never reports it). - Idempotent: re-stamping never produces ``"openrouter:openrouter:..."``; @@ -62,6 +64,9 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): return # No real sub-provider, or it just echoes our own router name. if not sub or sub.lower() == provider_type.lower(): + # A pinned endpoint is the only provider OpenRouter may route to + # (allow_fallbacks=False), so it names the serving provider. + pinned = pinned_endpoint_context.get() # Warn only on the billed payload, not on every stream chunk. if _carries_usage(response_json): logger.warning( @@ -69,9 +74,12 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): extra={ "model": response_json.get("model"), "response_id": response_json.get("id"), + "pinned_endpoint": pinned, }, ) - response_json["provider"] = f"{provider_type}:{_UNKNOWN_SUB_PROVIDER}" + response_json["provider"] = ( + f"{provider_type}:{pinned or _UNKNOWN_SUB_PROVIDER}" + ) return response_json["provider"] = f"{provider_type}:{sub}" diff --git a/tests/unit/test_model_path_routing.py b/tests/unit/test_model_path_routing.py index e2c3e52f..61eef377 100644 --- a/tests/unit/test_model_path_routing.py +++ b/tests/unit/test_model_path_routing.py @@ -15,7 +15,11 @@ from routstr.core.error_scope import ( ERROR_SCOPE_UPSTREAM, UPSTREAM_UNAVAILABLE, ) -from routstr.upstream.model_paths import decode_model_path, encode_model_path +from routstr.upstream.model_paths import ( + decode_model_path, + encode_model_path, + pinned_endpoint_context, +) from .proxy_test_utils import mock_request_stream, patch_proxy_session @@ -289,6 +293,39 @@ async def test_endpoint_tag_pins_the_upstream_subprovider() -> None: } +@pytest.mark.asyncio +async def test_endpoint_tag_is_exposed_to_response_stamping() -> None: + """The pin is visible while the upstream handles the request, so response + stamping can name the endpoint; the next unpinned request sees none.""" + upstream = _make_upstream(1) + upstream.base_url = "https://openrouter.ai/api/v1" + seen: list[str | None] = [] + + async def forward(*args: Any, **kwargs: Any) -> Any: + seen.append(pinned_endpoint_context.get()) + return MagicMock(status_code=200, body=b"{}") + + upstream.forward_request = AsyncMock(side_effect=forward) + pinned = _make_request( + { + "authorization": "Bearer sk-mpkey", + "x-routstr-model-path": encode_model_path( + upstream.base_url, MODEL_ID, "deepinfra/fp8" + ), + }, + json.dumps({"model": MODEL_ID}).encode(), + ) + unpinned = _make_request( + {"authorization": "Bearer sk-mpkey"}, + json.dumps({"model": MODEL_ID}).encode(), + ) + + await _run_proxy(pinned, [(MagicMock(), upstream)]) + await _run_proxy(unpinned, [(MagicMock(), upstream)]) + + assert seen == ["deepinfra/fp8", None] + + @pytest.mark.asyncio @pytest.mark.parametrize( "raw", diff --git a/tests/unit/test_provider_field_injection.py b/tests/unit/test_provider_field_injection.py index 6caea063..06c7e12f 100644 --- a/tests/unit/test_provider_field_injection.py +++ b/tests/unit/test_provider_field_injection.py @@ -3,6 +3,7 @@ from unittest.mock import patch from routstr.upstream.anthropic import AnthropicUpstreamProvider from routstr.upstream.base import BaseUpstreamProvider from routstr.upstream.generic import GenericUpstreamProvider +from routstr.upstream.model_paths import pinned_endpoint_context from routstr.upstream.openrouter import OpenRouterUpstreamProvider @@ -79,6 +80,31 @@ def test_apply_provider_field_openrouter_warns_once_on_billed_payload() -> None: assert chunk["provider"] == completed["provider"] == "openrouter:unknown" +def test_apply_provider_field_openrouter_falls_back_to_pinned_endpoint() -> None: + """With the request pinned to one endpoint, an unreported provider is that + endpoint, still logged; a reported one keeps winning.""" + p = _make_provider(OpenRouterUpstreamProvider, "openrouter") + token = pinned_endpoint_context.set("deepinfra/fp8") + try: + missing: dict = {"id": "gen-abc", "usage": {"prompt_tokens": 1}} + with patch("routstr.upstream.openrouter.logger.warning") as warning: + p._apply_provider_field(missing) + p._apply_provider_field(missing) + warning.assert_called_once() + assert warning.call_args.kwargs["extra"]["pinned_endpoint"] == "deepinfra/fp8" + assert missing["provider"] == "openrouter:deepinfra/fp8" + + reported: dict = {"provider": "Fireworks"} + p._apply_provider_field(reported) + assert reported["provider"] == "openrouter:Fireworks" + finally: + pinned_endpoint_context.reset(token) + + unpinned: dict = {"id": "gen-def"} + p._apply_provider_field(unpinned) + assert unpinned["provider"] == "openrouter:unknown" + + def test_apply_provider_field_openrouter_idempotent_no_double_prefix() -> None: """Re-stamping must never nest the prefix: openrouter only once.""" p = _make_provider(OpenRouterUpstreamProvider, "openrouter")