From a1ae6e94e91deecafb6cb1d19290de14db7da4c0 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 15 Apr 2026 19:20:50 +0200 Subject: [PATCH] match model when versioned --- routstr/proxy.py | 20 ++- tests/integration/test_usage_injection.py | 145 ---------------------- 2 files changed, 19 insertions(+), 146 deletions(-) delete mode 100644 tests/integration/test_usage_injection.py diff --git a/routstr/proxy.py b/routstr/proxy.py index ddab1c31..9d2dbf09 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -69,7 +69,25 @@ def get_upstreams() -> list[BaseUpstreamProvider]: def get_model_instance(model_id: str) -> Model | None: """Get Model instance by ID from global cache.""" - return _model_instances.get(model_id.lower()) + if not model_id: + return None + + model_id_lower = model_id.lower() + # Try exact match first + if model := _model_instances.get(model_id_lower): + return model + + # Try stripping common version suffixes (e.g., -20251222) + # This handles cases where upstream returns a specific version + # but we only track the base model name. + import re + + base_model_id = re.sub(r"-\d{8}$", "", model_id_lower) + if base_model_id != model_id_lower: + if model := _model_instances.get(base_model_id): + return model + + return None def get_provider_for_model(model_id: str) -> list[BaseUpstreamProvider] | None: diff --git a/tests/integration/test_usage_injection.py b/tests/integration/test_usage_injection.py deleted file mode 100644 index 59b672fd..00000000 --- a/tests/integration/test_usage_injection.py +++ /dev/null @@ -1,145 +0,0 @@ -import json -from collections.abc import AsyncGenerator -from unittest.mock import AsyncMock, patch - -import pytest -from httpx import AsyncClient - - -@pytest.mark.integration -@pytest.mark.asyncio -async def test_usage_and_metadata_injection_all_endpoints( - authenticated_client: AsyncClient, -) -> None: - """Test that usage, cost, and metadata are injected into all completion endpoints.""" - - endpoints = [ - ( - "/v1/chat/completions", - { - "model": "gpt-3.5-turbo", - "messages": [{"role": "user", "content": "Hello"}], - }, - { - "id": "chatcmpl-123", - "object": "chat.completion", - "model": "gpt-3.5-turbo", - "choices": [{"message": {"content": "Hi"}, "finish_reason": "stop"}], - "usage": { - "prompt_tokens": 10, - "completion_tokens": 5, - "total_tokens": 15, - }, - }, - ), - ( - "/v1/messages", - { - "model": "claude-3-opus-20240229", - "messages": [{"role": "user", "content": "Hello"}], - "max_tokens": 100, - }, - { - "id": "msg_123", - "type": "message", - "role": "assistant", - "model": "claude-3-opus-20240229", - "content": [{"type": "text", "text": "Hi"}], - "usage": {"input_tokens": 10, "output_tokens": 5}, - }, - ), - ] - - for path, payload, mock_data in endpoints: - with patch("httpx.AsyncClient.send") as mock_send: - - async def mock_iter_bytes( - *args: object, **kwargs: object - ) -> AsyncGenerator[bytes, None]: - yield json.dumps(mock_data).encode() - - mock_response = AsyncMock() - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - mock_response.text = json.dumps(mock_data) - mock_response.json = AsyncMock(return_value=mock_data) - mock_response.aiter_bytes = mock_iter_bytes - mock_send.return_value = mock_response - - response = await authenticated_client.post(path, json=payload) - assert response.status_code == 200 - - data = response.json() - if hasattr(data, "__await__"): - data = await data - - # Check unified usage injection - assert "usage" in data - assert "cost" in data["usage"] - assert "sats_cost" in data["usage"] - assert "remaining_balance_msats" in data["usage"] - - # Check metadata injection - assert "metadata" in data - assert "routstr" in data["metadata"] - assert "cost" in data["metadata"]["routstr"] - assert "sats_cost" in data["metadata"]["routstr"] - assert "remaining_balance_msats" in data["metadata"]["routstr"] - - # Check legacy cost injection - assert "cost" in data - assert "sats_cost" in data["cost"] - - -@pytest.mark.integration -@pytest.mark.asyncio -async def test_streaming_usage_injection( - authenticated_client: AsyncClient, -) -> None: - """Test usage injection for streaming messages.""" - - path = "/v1/messages" - payload = { - "model": "claude-3-opus-20240229", - "messages": [{"role": "user", "content": "Hello"}], - "max_tokens": 100, - "stream": True, - } - - # Final usage chunk for Anthropic - streaming_chunks = [ - b'event: message_start\ndata: {"type": "message_start", "message": {"id": "msg_123", "type": "message", "role": "assistant", "model": "claude-3-opus-20240229", "usage": {"input_tokens": 10, "output_tokens": 0}}}\n\n', - b'event: message_delta\ndata: {"type": "message_delta", "delta": {"stop_reason": "end_turn", "stop_sequence": null}, "usage": {"output_tokens": 5}}\n\n', - b'event: message_stop\ndata: {"type": "message_stop"}\n\n', - ] - - with patch("httpx.AsyncClient.send") as mock_send: - - async def mock_iter_bytes( - *args: object, **kwargs: object - ) -> AsyncGenerator[bytes, None]: - for chunk in streaming_chunks: - yield chunk - - mock_response = AsyncMock() - mock_response.status_code = 200 - mock_response.headers = {"content-type": "text/event-stream"} - mock_response.aiter_bytes = mock_iter_bytes - mock_send.return_value = mock_response - - response = await authenticated_client.post(path, json=payload) - assert response.status_code == 200 - - # In streaming, we look for the final 'event: cost' which is added by Routstr - response_text = response.text - if hasattr(response_text, "__await__"): - response_text = await response_text - lines = response_text.split("\n") - - cost_event = next( - (line for line in lines if line.startswith('data: {"cost":')), None - ) - assert cost_event is not None - - cost_data = json.loads(cost_event[6:])["cost"] - assert "total_msats" in cost_data