From 0457b5d0cdf4d86c7690cfe408b08d9b1a8b0f0c Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 27 Aug 2026 13:23:27 +0200 Subject: [PATCH] display correct provider --- routstr/upstream/base.py | 14 +++--- tests/unit/test_x_cashu_provider_path.py | 54 ++++++++++++++++++++++++ 2 files changed, 60 insertions(+), 8 deletions(-) create mode 100644 tests/unit/test_x_cashu_provider_path.py diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 24f4d165..f8b7fe3e 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -3928,10 +3928,9 @@ class BaseUpstreamProvider: 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 + provider_before = data_json.get("provider") + self._apply_provider_field(data_json) + 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) changed = True @@ -4937,10 +4936,9 @@ class BaseUpstreamProvider: continue if not isinstance(data_json, dict): continue - changed = False - if "provider" not in data_json: - self._apply_provider_field(data_json) - changed = True + provider_before = data_json.get("provider") + self._apply_provider_field(data_json) + changed = data_json.get("provider") != provider_before payload = _responses_usage_payload(data_json) if cost_data and isinstance(payload.get("usage"), dict): _inject_cost_into_usage(payload, cost_data) diff --git a/tests/unit/test_x_cashu_provider_path.py b/tests/unit/test_x_cashu_provider_path.py new file mode 100644 index 00000000..11b47d6f --- /dev/null +++ b/tests/unit/test_x_cashu_provider_path.py @@ -0,0 +1,54 @@ +import json +from typing import Any +from unittest.mock import AsyncMock, patch + +import httpx +import pytest + +from routstr.upstream.openrouter import OpenRouterUpstreamProvider + + +async def _body(response: Any) -> bytes: + chunks: list[bytes] = [] + async for chunk in response.body_iterator: + chunks.append(chunk if isinstance(chunk, bytes) else chunk.encode()) + return b"".join(chunks) + + +@pytest.mark.asyncio +async def test_x_cashu_chat_stream_reports_complete_provider_path() -> None: + provider = OpenRouterUpstreamProvider(api_key="test-key") + payload = {"model": "glm-4.5", "provider": "z.ai"} + content = f"data: {json.dumps(payload)}\n" + + 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, + ) + + event = json.loads((await _body(response)).decode().removeprefix("data: ")) + assert event["provider"] == "openrouter:z.ai" + + +@pytest.mark.asyncio +async def test_x_cashu_responses_stream_reports_complete_provider_path() -> None: + provider = OpenRouterUpstreamProvider(api_key="test-key") + event = {"type": "response.created", "provider": "z.ai"} + content = f"data: {json.dumps(event)}\n\n" + + 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, + ) + + payload = json.loads((await _body(response)).decode().removeprefix("data: ")) + assert payload["provider"] == "openrouter:z.ai"