From 8ec5d409d520ee9fa5455b8b7a1da5c9009592d2 Mon Sep 17 00:00:00 2001 From: thefux Date: Tue, 29 Sep 2026 08:08:59 +0000 Subject: [PATCH] fix: stop reporting OpenRouter provider as unknown on stream and envelope payloads The OpenRouter stamper wrote "unknown" whenever a payload lacked a top-level provider. That hit every Anthropic /messages event, every Responses event, and the usage/cost payloads routstr synthesizes at the end of a stream. - Read the provider from the Anthropic `message` and Responses `response` envelopes as well as the top level. - Carry the provider reported earlier in a stream to later events and to the synthesized usage/cost payloads. --- routstr/upstream/base.py | 55 +++++++++++++++++---- routstr/upstream/generic.py | 5 +- routstr/upstream/openrouter.py | 5 +- tests/unit/test_provider_field_injection.py | 27 ++++++++++ tests/unit/test_x_cashu_provider_path.py | 47 ++++++++++++++++++ 5 files changed, 124 insertions(+), 15 deletions(-) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 5b0e5c0f..8f46d90b 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -214,6 +214,20 @@ def _responses_usage_payload(data_json: dict) -> dict: return nested if isinstance(nested, dict) else data_json +def _reported_provider(payload: dict) -> str | None: + """Provider named by an upstream payload, if any. + + Checked at top level first, then inside the Anthropic ``message`` and + Responses ``response`` envelopes, which is where those dialects nest it. + """ + for obj in (payload, payload.get("message"), payload.get("response")): + if isinstance(obj, dict): + value = obj.get("provider") + if isinstance(value, str) and value.strip(): + return value.strip() + return None + + def _render_sse_event(field_lines: list[str], data: str) -> str: """Re-frame one parsed event, re-prefixing every line of a multi-line data.""" body = "".join(f"{line}\n" for line in field_lines) @@ -485,8 +499,7 @@ class BaseUpstreamProvider: 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 "" + existing_str = _reported_provider(response_json) or "" if not existing_str: response_json["provider"] = provider_type return @@ -498,6 +511,17 @@ class BaseUpstreamProvider: return response_json["provider"] = f"{provider_type}:{existing_str}" + def _stamp_streamed_provider( + self, payload: dict, carried: str | None + ) -> str | None: + """Stamp a streamed payload, falling back to a provider an earlier event + reported. Returns the provider to carry forward to later payloads.""" + reported = _reported_provider(payload) + if reported is None and carried is not None: + payload["provider"] = carried + self._apply_provider_field(payload) + return reported or carried + def _log_full_refund( self, *, @@ -1169,6 +1193,7 @@ class BaseUpstreamProvider: usage_finalized = False last_model_seen: str | None = None + provider_seen: str | None = None async def finalize_db_only() -> None: nonlocal usage_finalized @@ -1243,6 +1268,7 @@ class BaseUpstreamProvider: end of stream. """ nonlocal last_model_seen, usage_chunk_data, done_seen, stream_id + nonlocal provider_seen event = raw_event.strip(b"\r\n") if not event: @@ -1282,7 +1308,7 @@ class BaseUpstreamProvider: if isinstance(obj, dict): usage_estimator.observe(obj) - self._apply_provider_field(obj) + provider_seen = self._stamp_streamed_provider(obj, provider_seen) if obj.get("model"): last_model_seen = str(obj.get("model")) if requested_model: @@ -1408,6 +1434,7 @@ class BaseUpstreamProvider: if legacy_completion else "chat.completion.chunk", "model": last_model_seen or "unknown", + "provider": provider_seen, "choices": [], "usage": { "prompt_tokens": cost_data.get("input_tokens", 0), @@ -1652,6 +1679,7 @@ class BaseUpstreamProvider: usage_finalized = False last_model_seen: str | None = None + provider_seen: str | None = None async def finalize_db_only() -> None: nonlocal usage_finalized @@ -1715,7 +1743,7 @@ class BaseUpstreamProvider: and preserves ``event:``/``id:`` fields attached to their data line so Responses API event framing stays intact. """ - nonlocal last_model_seen, usage_chunk_data, done_seen + nonlocal last_model_seen, usage_chunk_data, done_seen, provider_seen nonlocal reasoning_tokens event = raw_event.strip(b"\r\n") @@ -1751,7 +1779,7 @@ class BaseUpstreamProvider: obj = json_codec.loads(data) if isinstance(obj, dict): - self._apply_provider_field(obj) + provider_seen = self._stamp_streamed_provider(obj, provider_seen) if obj.get("model"): last_model_seen = str(obj.get("model")) if requested_model: @@ -1840,6 +1868,7 @@ class BaseUpstreamProvider: if usage_chunk_data is None: usage_chunk_data = { "type": "response.completed", + "provider": provider_seen, "response": { "model": last_model_seen or "unknown", "usage": { @@ -2195,6 +2224,7 @@ class BaseUpstreamProvider: usage_estimator = MissingUsageEstimator(request_body, model_obj) usage_finalized = False last_model_seen: str | None = None + provider_seen: str | None = None async def finalize_without_usage() -> bytes | None: nonlocal usage_finalized @@ -2244,7 +2274,7 @@ class BaseUpstreamProvider: async def stream_with_cost( max_cost_for_model: int, ) -> AsyncGenerator[bytes, None]: - nonlocal usage_finalized, last_model_seen + nonlocal usage_finalized, last_model_seen, provider_seen stored_chunks: list[bytes] = [] input_tokens: int = 0 output_tokens: int = 0 @@ -2301,7 +2331,9 @@ class BaseUpstreamProvider: last_model_seen = str(msg.get("model")) provider_added = "provider" not in data - self._apply_provider_field(data) + provider_seen = self._stamp_streamed_provider( + data, provider_seen + ) if requested_model: # Apply requested_model override @@ -2419,6 +2451,7 @@ class BaseUpstreamProvider: try: combined_data = { "model": last_model_seen or "unknown", + "provider": provider_seen, "usage": usage_data, } cost_data = await adjust_payment_for_tokens( @@ -4197,6 +4230,7 @@ class BaseUpstreamProvider: }, ) + provider_seen: str | None = None for i, line in enumerate(lines): if line.startswith("data: "): try: @@ -4204,7 +4238,9 @@ class BaseUpstreamProvider: if not isinstance(data_json, dict): continue provider_before = data_json.get("provider") - self._apply_provider_field(data_json) + provider_seen = self._stamp_streamed_provider( + data_json, provider_seen + ) changed = data_json.get("provider") != provider_before if cost_data and "usage" in data_json and data_json["usage"]: _inject_cost_into_usage(data_json, cost_data) @@ -5265,6 +5301,7 @@ class BaseUpstreamProvider: }, ) + provider_seen: str | None = None for i, (fields, data) in enumerate(events): if data.strip() == "[DONE]": continue @@ -5275,7 +5312,7 @@ class BaseUpstreamProvider: if not isinstance(data_json, dict): continue provider_before = data_json.get("provider") - self._apply_provider_field(data_json) + provider_seen = self._stamp_streamed_provider(data_json, provider_seen) changed = data_json.get("provider") != provider_before payload = _responses_usage_payload(data_json) if cost_data and isinstance(payload.get("usage"), dict): diff --git a/routstr/upstream/generic.py b/routstr/upstream/generic.py index 03bfa015..c9edf109 100644 --- a/routstr/upstream/generic.py +++ b/routstr/upstream/generic.py @@ -5,7 +5,7 @@ from urllib.parse import urlparse import httpx -from .base import BaseUpstreamProvider +from .base import BaseUpstreamProvider, _reported_provider from .model_paths import public_provider_url from .pricing_resolver import ( FallbackPricingResolver, @@ -60,8 +60,7 @@ class GenericUpstreamProvider(BaseUpstreamProvider): """ if not isinstance(response_json, dict): return - existing = response_json.get("provider") - if not (isinstance(existing, str) and existing.strip()): + if _reported_provider(response_json) is None: response_json["provider"] = ( urlparse(public_provider_url(self.base_url)).hostname or self.upstream_name diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index 9ca190ce..34394335 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -3,7 +3,7 @@ from typing import TYPE_CHECKING import httpx from ..payment.models import Model, async_fetch_openrouter_models -from .base import BaseUpstreamProvider +from .base import BaseUpstreamProvider, _reported_provider from .model_paths import public_provider_url if TYPE_CHECKING: @@ -35,8 +35,7 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): 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 "" + sub = _reported_provider(response_json) or "" # Strip any already-applied "openrouter:" prefixes (idempotency). prefix = f"{provider_type}:" while sub.lower().startswith(prefix.lower()): diff --git a/tests/unit/test_provider_field_injection.py b/tests/unit/test_provider_field_injection.py index 86942ae6..d620e054 100644 --- a/tests/unit/test_provider_field_injection.py +++ b/tests/unit/test_provider_field_injection.py @@ -89,6 +89,33 @@ def test_apply_provider_field_non_string_upstream_treated_as_missing() -> None: assert data["provider"] == "unknown" +def test_apply_provider_field_openrouter_reads_nested_envelopes() -> None: + """Anthropic ``message`` and Responses ``response`` envelopes nest the + upstream provider; it must not be reported as unknown.""" + p = _make_provider(OpenRouterUpstreamProvider, "openrouter") + message_start: dict = { + "type": "message_start", + "message": {"provider": "Anthropic"}, + } + p._apply_provider_field(message_start) + assert message_start["provider"] == "openrouter:Anthropic" + + created: dict = {"type": "response.created", "response": {"provider": "OpenAI"}} + p._apply_provider_field(created) + assert created["provider"] == "openrouter:OpenAI" + + +def test_stamp_streamed_provider_carries_earlier_provider() -> None: + """Events without their own provider inherit the one reported earlier in + the stream instead of becoming ``unknown``.""" + p = _make_provider(OpenRouterUpstreamProvider, "openrouter") + first: dict = {"provider": "Fireworks"} + carried = p._stamp_streamed_provider(first, None) + delta: dict = {"type": "content_block_delta"} + assert p._stamp_streamed_provider(delta, carried) == "Fireworks" + assert first["provider"] == delta["provider"] == "openrouter:Fireworks" + + def test_apply_provider_field_idempotent_for_direct_upstream() -> None: """Calling twice on a direct upstream payload keeps the same value and never nests the prefix (no ``anthropic:anthropic``).""" diff --git a/tests/unit/test_x_cashu_provider_path.py b/tests/unit/test_x_cashu_provider_path.py index 11b47d6f..6b4fdec2 100644 --- a/tests/unit/test_x_cashu_provider_path.py +++ b/tests/unit/test_x_cashu_provider_path.py @@ -52,3 +52,50 @@ async def test_x_cashu_responses_stream_reports_complete_provider_path() -> None payload = json.loads((await _body(response)).decode().removeprefix("data: ")) assert payload["provider"] == "openrouter:z.ai" + + +@pytest.mark.asyncio +async def test_x_cashu_messages_stream_carries_provider_to_later_events() -> None: + provider = OpenRouterUpstreamProvider(api_key="test-key") + events = [ + {"type": "message_start", "message": {"provider": "Anthropic"}}, + {"type": "content_block_delta", "delta": {"text": "hi"}}, + ] + content = "".join(f"data: {json.dumps(e)}\n" for e in events) + + response = await provider.handle_x_cashu_streaming_response( + content, + httpx.Response(200, headers={"content-type": "text/event-stream"}), + amount=1, + unit="sat", + max_cost_for_model=1, + ) + + lines = (await _body(response)).decode().splitlines() + stamped = [json.loads(line.removeprefix("data: ")) for line in lines if line] + assert [e["provider"] for e in stamped] == ["openrouter:Anthropic"] * 2 + + +@pytest.mark.asyncio +async def test_x_cashu_responses_stream_carries_nested_provider() -> None: + provider = OpenRouterUpstreamProvider(api_key="test-key") + events = [ + {"type": "response.created", "response": {"provider": "OpenAI"}}, + {"type": "response.output_text.delta", "delta": "hi"}, + ] + content = "".join(f"data: {json.dumps(e)}\n\n" for e in events) + + with patch.object( + provider, "get_x_cashu_cost", new=AsyncMock(return_value=None) + ): + response = await provider.handle_x_cashu_streaming_responses_response( + content, + httpx.Response(200, headers={"content-type": "text/event-stream"}), + amount=1, + unit="sat", + max_cost_for_model=1, + ) + + lines = (await _body(response)).decode().splitlines() + stamped = [json.loads(line.removeprefix("data: ")) for line in lines if line] + assert [e["provider"] for e in stamped] == ["openrouter:OpenAI"] * 2