From efb57196799d69ea311b801245b68afb67d6d80b Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 17 May 2026 14:39:16 +0200 Subject: [PATCH] add provider field to response --- routstr/upstream/base.py | 104 ++++++++++++++----- tests/unit/test_provider_field_injection.py | 105 ++++++++++++++++++++ 2 files changed, 184 insertions(+), 25 deletions(-) create mode 100644 tests/unit/test_provider_field_injection.py diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 03355445..0823ea24 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -197,6 +197,22 @@ class BaseUpstreamProvider: except (TypeError, ValueError): pass + def _apply_provider_field(self, response_json: object) -> None: + """Stamp the routstr ``provider`` field onto an upstream response payload. + + Format is ``":"`` when the upstream + already reported its own provider (e.g. OpenRouter returns + ``"provider": "Fireworks"``), otherwise just ``""`` + for direct upstreams. + """ + if not isinstance(response_json, dict): + return + existing = response_json.get("provider") + if isinstance(existing, str) and existing.strip(): + response_json["provider"] = f"{self.provider_type}:{existing.strip()}" + else: + response_json["provider"] = self.provider_type + def inject_cost_metadata( self, response_json: dict, @@ -204,6 +220,7 @@ class BaseUpstreamProvider: key: ApiKey, ) -> None: """Unifies the injection of cost and usage metadata across all completion types.""" + self._apply_provider_field(response_json) if isinstance(cost_data, dict): total_msats = cost_data.get("total_msats", 0) total_usd = cost_data.get("total_usd", 0.0) @@ -723,6 +740,7 @@ class BaseUpstreamProvider: ): obj = json.loads(part) if isinstance(obj, dict): + self._apply_provider_field(obj) if obj.get("model"): last_model_seen = str(obj.get("model")) if requested_model: @@ -889,6 +907,7 @@ class BaseUpstreamProvider: try: content = await response.aread() response_json = json.loads(content) + self._apply_provider_field(response_json) logger.debug( "Parsed response JSON", @@ -1068,6 +1087,7 @@ class BaseUpstreamProvider: try: obj = json.loads(part) if isinstance(obj, dict): + self._apply_provider_field(obj) if obj.get("model"): last_model_seen = str(obj.get("model")) if requested_model: @@ -1261,6 +1281,7 @@ class BaseUpstreamProvider: try: content = await response.aread() response_json = json.loads(content) + self._apply_provider_field(response_json) logger.debug( "Parsed Responses API response JSON", @@ -1499,6 +1520,11 @@ class BaseUpstreamProvider: if msg and msg.get("model"): last_model_seen = str(msg.get("model")) + provider_added = ( + "provider" not in data + ) + self._apply_provider_field(data) + if requested_model: # Apply requested_model override model_updated = False @@ -1509,9 +1535,12 @@ class BaseUpstreamProvider: data["model"] = requested_model model_updated = True - if model_updated: + if model_updated or provider_added: line = "data: " + json.dumps(data) changed = True + elif provider_added: + line = "data: " + json.dumps(data) + changed = True if usage := msg.get("usage"): input_tokens += usage.get("input_tokens", 0) @@ -1833,6 +1862,7 @@ class BaseUpstreamProvider: ) response_json = messages_dispatch.coerce_litellm_payload(result) + self._apply_provider_field(response_json) if requested_model and "model" in response_json: response_json["model"] = requested_model @@ -3145,18 +3175,29 @@ class BaseUpstreamProvider: }, ) - if cost_data: - for i, line in enumerate(lines): - if line.startswith("data: "): - try: - data_json = json.loads(line[6:]) - if "usage" in data_json and data_json["usage"]: - data_json["usage"]["cost_sats"] = ( - cost_data.total_msats // 1000 - ) - lines[i] = "data: " + json.dumps(data_json) - except json.JSONDecodeError: - pass + for i, line in enumerate(lines): + if line.startswith("data: "): + try: + data_json = json.loads(line[6:]) + if not isinstance(data_json, dict): + continue + changed = False + if "provider" not in data_json: + self._apply_provider_field(data_json) + changed = True + if ( + cost_data + and "usage" in data_json + and data_json["usage"] + ): + data_json["usage"]["cost_sats"] = ( + cost_data.total_msats // 1000 + ) + changed = True + if changed: + lines[i] = "data: " + json.dumps(data_json) + except json.JSONDecodeError: + pass async def generate() -> AsyncGenerator[bytes, None]: for line in lines: @@ -3200,6 +3241,7 @@ class BaseUpstreamProvider: try: response_json = json.loads(content_str) + self._apply_provider_field(response_json) cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model) if cost_data and "usage" in response_json: @@ -4121,18 +4163,29 @@ class BaseUpstreamProvider: }, ) - if cost_data: - for i, line in enumerate(lines): - if line.startswith("data: "): - try: - data_json = json.loads(line[6:]) - if "usage" in data_json and data_json["usage"]: - data_json["usage"]["cost_sats"] = ( - cost_data.total_msats // 1000 - ) - lines[i] = "data: " + json.dumps(data_json) - except json.JSONDecodeError: - pass + for i, line in enumerate(lines): + if line.startswith("data: "): + try: + data_json = json.loads(line[6:]) + if not isinstance(data_json, dict): + continue + changed = False + if "provider" not in data_json: + self._apply_provider_field(data_json) + changed = True + if ( + cost_data + and "usage" in data_json + and data_json["usage"] + ): + data_json["usage"]["cost_sats"] = ( + cost_data.total_msats // 1000 + ) + changed = True + if changed: + lines[i] = "data: " + json.dumps(data_json) + except json.JSONDecodeError: + pass async def generate() -> AsyncGenerator[bytes, None]: for line in lines: @@ -4164,6 +4217,7 @@ class BaseUpstreamProvider: try: response_json = json.loads(content_str) + self._apply_provider_field(response_json) cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model) if cost_data and "usage" in response_json: diff --git a/tests/unit/test_provider_field_injection.py b/tests/unit/test_provider_field_injection.py new file mode 100644 index 00000000..93658a94 --- /dev/null +++ b/tests/unit/test_provider_field_injection.py @@ -0,0 +1,105 @@ +from routstr.upstream.anthropic import AnthropicUpstreamProvider +from routstr.upstream.base import BaseUpstreamProvider +from routstr.upstream.openrouter import OpenRouterUpstreamProvider + + +def _make_provider(cls: type, provider_type: str) -> BaseUpstreamProvider: + p = cls(api_key="test_key") + assert p.provider_type == provider_type + return p + + +def test_apply_provider_field_direct_upstream() -> None: + """For a direct upstream (no upstream-reported provider), the field + is just the provider_type string.""" + p = _make_provider(AnthropicUpstreamProvider, "anthropic") + data: dict = {"id": "msg_1", "model": "claude-3-5-sonnet"} + p._apply_provider_field(data) + assert data["provider"] == "anthropic" + + +def test_apply_provider_field_openrouter_passthrough() -> None: + """OpenRouter responses include an upstream ``provider`` string — + routstr should prefix with its own provider_type.""" + p = _make_provider(OpenRouterUpstreamProvider, "openrouter") + data: dict = { + "id": "gen-abc", + "model": "anthropic/claude-3.5-sonnet", + "provider": "Anthropic", + } + p._apply_provider_field(data) + assert data["provider"] == "openrouter:Anthropic" + + +def test_apply_provider_field_openrouter_no_upstream_provider() -> None: + """If OpenRouter omits the provider field, fall back to provider_type.""" + p = _make_provider(OpenRouterUpstreamProvider, "openrouter") + data: dict = {"id": "gen-abc"} + p._apply_provider_field(data) + assert data["provider"] == "openrouter" + + +def test_apply_provider_field_strips_whitespace() -> None: + p = _make_provider(OpenRouterUpstreamProvider, "openrouter") + data: dict = {"provider": " Fireworks "} + p._apply_provider_field(data) + assert data["provider"] == "openrouter:Fireworks" + + +def test_apply_provider_field_blank_upstream_treated_as_missing() -> None: + p = _make_provider(OpenRouterUpstreamProvider, "openrouter") + data: dict = {"provider": " "} + p._apply_provider_field(data) + assert data["provider"] == "openrouter" + + +def test_apply_provider_field_non_string_upstream_treated_as_missing() -> None: + p = _make_provider(OpenRouterUpstreamProvider, "openrouter") + data: dict = {"provider": 42} + p._apply_provider_field(data) + assert data["provider"] == "openrouter" + + +def test_apply_provider_field_idempotent_for_direct_upstream() -> None: + """Calling twice on a direct upstream payload should keep the same + value, not nest the prefix repeatedly.""" + p = _make_provider(AnthropicUpstreamProvider, "anthropic") + data: dict = {} + p._apply_provider_field(data) + p._apply_provider_field(data) + assert data["provider"] == "anthropic:anthropic" + # Document current (deliberate) behavior: second pass treats the + # first-pass value as an upstream-reported provider. Callers should + # only invoke this once per chunk — guarded via the + # ``"provider" not in data`` checks in streaming paths. + + +def test_apply_provider_field_ignores_non_dict() -> None: + """Lists / primitives must be skipped silently.""" + p = _make_provider(AnthropicUpstreamProvider, "anthropic") + # Should not raise. + p._apply_provider_field([1, 2, 3]) # type: ignore[arg-type] + p._apply_provider_field("hello") # type: ignore[arg-type] + p._apply_provider_field(None) # type: ignore[arg-type] + + +def test_inject_cost_metadata_sets_provider() -> None: + """``inject_cost_metadata`` is the unified injection point and must + also stamp the provider field.""" + from unittest.mock import MagicMock + + from routstr.core.db import ApiKey + + p = _make_provider(OpenRouterUpstreamProvider, "openrouter") + key = MagicMock(spec=ApiKey) + key.balance = 1000 + + response_json: dict = { + "model": "anthropic/claude-3.5-sonnet", + "provider": "Anthropic", + "usage": {"prompt_tokens": 10, "completion_tokens": 5}, + } + cost_data = {"total_msats": 2500, "total_usd": 0.0025} + p.inject_cost_metadata(response_json, cost_data, key) + + assert response_json["provider"] == "openrouter:Anthropic"