diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 217cf9fc..0d01fccd 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -63,6 +63,7 @@ from .cache_breakpoints import ( ) from .count_tokens import MissingUsageEstimator, count_tokens_locally from .litellm_routing import detect_litellm_prefix +from .model_paths import public_provider_url from .rate_limit import UPSTREAM_RATE_LIMIT, classify_rate_limit from .reasoning_effort import apply_reasoning_effort @@ -476,9 +477,13 @@ class BaseUpstreamProvider: Idempotent: re-stamping an already-stamped payload must not nest the prefix repeatedly (e.g. never ``"anthropic:anthropic"``). This matters because streaming paths can apply the field more than once per chunk. + + Also stamps ``provider_url`` with the upstream base URL that served + the request. """ if not isinstance(response_json, dict): return + response_json["provider_url"] = public_provider_url(self.base_url) provider_type = (self.provider_type or "").strip() existing = response_json.get("provider") existing_str = existing.strip() if isinstance(existing, str) else "" diff --git a/routstr/upstream/generic.py b/routstr/upstream/generic.py index 3faf80e9..b0cc69cc 100644 --- a/routstr/upstream/generic.py +++ b/routstr/upstream/generic.py @@ -1,10 +1,12 @@ from __future__ import annotations from typing import TYPE_CHECKING +from urllib.parse import urlparse import httpx from .base import BaseUpstreamProvider +from .model_paths import public_provider_url from .pricing_resolver import ( FallbackPricingResolver, ResolvedPricing, @@ -50,6 +52,22 @@ class GenericUpstreamProvider(BaseUpstreamProvider): provider_fee=provider_fee, ) + def _apply_provider_field(self, response_json: object) -> None: + """Stamp ``"generic:"`` unless the upstream named itself. + + A generic upstream is not a router, so nothing identifies the serving + endpoint in the payload; the base URL host fills that role. + """ + if not isinstance(response_json, dict): + return + existing = response_json.get("provider") + if not (isinstance(existing, str) and existing.strip()): + response_json["provider"] = ( + urlparse(public_provider_url(self.base_url)).hostname + or self.upstream_name + ) + super()._apply_provider_field(response_json) + @classmethod def _build_from_row( cls, provider_row: "UpstreamProviderRow" diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index 1caeaa5c..9ca190ce 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -4,6 +4,7 @@ import httpx from ..payment.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider +from .model_paths import public_provider_url if TYPE_CHECKING: from ..core.db import UpstreamProviderRow @@ -32,6 +33,7 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): """ if not isinstance(response_json, dict): return + response_json["provider_url"] = public_provider_url(self.base_url) provider_type = (self.provider_type or "").strip() existing = response_json.get("provider") sub = existing.strip() if isinstance(existing, str) else "" diff --git a/tests/unit/test_provider_field_injection.py b/tests/unit/test_provider_field_injection.py index bf2813e0..86942ae6 100644 --- a/tests/unit/test_provider_field_injection.py +++ b/tests/unit/test_provider_field_injection.py @@ -1,5 +1,6 @@ from routstr.upstream.anthropic import AnthropicUpstreamProvider from routstr.upstream.base import BaseUpstreamProvider +from routstr.upstream.generic import GenericUpstreamProvider from routstr.upstream.openrouter import OpenRouterUpstreamProvider @@ -127,3 +128,53 @@ def test_inject_cost_metadata_sets_provider() -> None: p.inject_cost_metadata(response_json, cost_data, key) assert response_json["provider"] == "openrouter:Anthropic" + + +def test_apply_provider_field_generic_uses_upstream_host() -> None: + """A generic upstream has no router-reported provider; the serving host + identifies it, mirroring ``openrouter:``.""" + p = GenericUpstreamProvider(base_url="https://api.deepseek.com/v1", api_key="k") + data: dict = {"id": "chatcmpl-1", "model": "deepseek-chat"} + p._apply_provider_field(data) + assert data["provider"] == "generic:api.deepseek.com" + + +def test_apply_provider_field_generic_keeps_upstream_reported_provider() -> None: + p = GenericUpstreamProvider(base_url="https://api.deepseek.com/v1", api_key="k") + data: dict = {"provider": "Fireworks"} + p._apply_provider_field(data) + assert data["provider"] == "generic:Fireworks" + + +def test_apply_provider_field_generic_idempotent() -> None: + p = GenericUpstreamProvider(base_url="https://api.deepseek.com/v1", api_key="k") + data: dict = {} + p._apply_provider_field(data) + p._apply_provider_field(data) + assert data["provider"] == "generic:api.deepseek.com" + + +def test_apply_provider_field_sets_provider_url() -> None: + """Every provider exposes the upstream base URL it served from.""" + generic = GenericUpstreamProvider( + base_url="https://api.deepseek.com/v1", api_key="k" + ) + data: dict = {} + generic._apply_provider_field(data) + assert data["provider_url"] == "https://api.deepseek.com/v1" + + openrouter = _make_provider(OpenRouterUpstreamProvider, "openrouter") + data = {"provider": "Anthropic"} + openrouter._apply_provider_field(data) + assert data["provider_url"] == "https://openrouter.ai/api/v1" + + +def test_apply_provider_field_masks_private_upstream() -> None: + """Private or port-bearing upstream URLs are masked the same way model + paths mask them, so neither ``provider`` nor ``provider_url`` leaks a + local address.""" + p = GenericUpstreamProvider(base_url="http://10.0.0.5:11434/v1", api_key="k") + data: dict = {} + p._apply_provider_field(data) + assert data["provider"] == "generic:localhost" + assert data["provider_url"] == "http://localhost"