diff --git a/routstr/payment/models.py b/routstr/payment/models.py index d6eccf5d..06799705 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -93,12 +93,32 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis try: async with httpx.AsyncClient() as client: - response = await client.get(f"{base_url}/models", timeout=30) - response.raise_for_status() - data = response.json() + models_response, embeddings_response = await asyncio.gather( + client.get(f"{base_url}/models", timeout=30), + client.get(f"{base_url}/embeddings/models", timeout=30), + return_exceptions=True, + ) + + def process_models_response( + response: httpx.Response | BaseException, + ) -> list[dict]: + if not isinstance(response, BaseException): + response.raise_for_status() + data = response.json() + return [ + model + for model in data.get("data", []) + if ":free" not in model.get("id", "").lower() + ] + return [] models_data: list[dict] = [] - for model in data.get("data", []): + models_data.extend(process_models_response(models_response)) + models_data.extend(process_models_response(embeddings_response)) + + # Apply source filter and exclusions + filtered_models = [] + for model in models_data: model_id = model.get("id", "") if source_filter: @@ -116,9 +136,9 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis if not _has_valid_pricing(model): continue - models_data.append(model) + filtered_models.append(model) - return models_data + return filtered_models except Exception as e: logger.error(f"Error (async) fetching models from OpenRouter API: {e}") return [] diff --git a/routstr/proxy.py b/routstr/proxy.py index ce558fc5..edceaaa6 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -65,7 +65,7 @@ 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) + return _model_instances.get(model_id.lower()) def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None: diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 7726b5f6..ca7505cd 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -734,51 +734,53 @@ class BaseUpstreamProvider: await client.aclose() return mapped_error - if path.endswith("chat/completions"): - client_wants_streaming = False - if request_body: - try: - request_data = json.loads(request_body) - client_wants_streaming = request_data.get("stream", False) - logger.debug( - "Chat completion request analysis", - extra={ - "client_wants_streaming": client_wants_streaming, - "model": request_data.get("model", "unknown"), - "key_hash": key.hashed_key[:8] + "...", - }, - ) - except json.JSONDecodeError: - logger.warning( - "Failed to parse request body JSON for streaming detection" - ) + if path.endswith("chat/completions") or path.endswith("embeddings"): + if path.endswith("chat/completions"): + client_wants_streaming = False + if request_body: + try: + request_data = json.loads(request_body) + client_wants_streaming = request_data.get("stream", False) + logger.debug( + "Chat completion request analysis", + extra={ + "client_wants_streaming": client_wants_streaming, + "model": request_data.get("model", "unknown"), + "key_hash": key.hashed_key[:8] + "...", + }, + ) + except json.JSONDecodeError: + logger.warning( + "Failed to parse request body JSON for streaming detection" + ) - content_type = response.headers.get("content-type", "") - upstream_is_streaming = "text/event-stream" in content_type - is_streaming = client_wants_streaming and upstream_is_streaming + content_type = response.headers.get("content-type", "") + upstream_is_streaming = "text/event-stream" in content_type + is_streaming = client_wants_streaming and upstream_is_streaming - logger.debug( - "Response type analysis", - extra={ - "is_streaming": is_streaming, - "client_wants_streaming": client_wants_streaming, - "upstream_is_streaming": upstream_is_streaming, - "content_type": content_type, - "key_hash": key.hashed_key[:8] + "...", - }, - ) - - if is_streaming and response.status_code == 200: - result = await self.handle_streaming_chat_completion( - response, key, max_cost_for_model + logger.debug( + "Response type analysis", + extra={ + "is_streaming": is_streaming, + "client_wants_streaming": client_wants_streaming, + "upstream_is_streaming": upstream_is_streaming, + "content_type": content_type, + "key_hash": key.hashed_key[:8] + "...", + }, ) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - result.background = background_tasks - return result - elif response.status_code == 200: + if is_streaming and response.status_code == 200: + result = await self.handle_streaming_chat_completion( + response, key, max_cost_for_model + ) + background_tasks = BackgroundTasks() + background_tasks.add_task(response.aclose) + background_tasks.add_task(client.aclose) + result.background = background_tasks + return result + + # Handle both non-streaming chat completions and embeddings + if response.status_code == 200: try: return await self.handle_non_streaming_chat_completion( response, key, session, max_cost_for_model @@ -1519,9 +1521,9 @@ class BaseUpstreamProvider: error_response.headers["X-Cashu"] = refund_token return error_response - if path.endswith("chat/completions"): + if path.endswith("chat/completions") or path.endswith("embeddings"): logger.debug( - "Processing chat completion response", + "Processing completion/embeddings response", extra={"path": path, "amount": amount, "unit": unit}, ) @@ -1770,15 +1772,32 @@ class BaseUpstreamProvider: async def _fetch_openrouter_models(self) -> list[dict]: """Fetch models from OpenRouter API.""" url = "https://openrouter.ai/api/v1/models" + embeddings_url = "https://openrouter.ai/api/v1/embeddings/models" + async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.get(url) - response.raise_for_status() - models = response.json() - return [ - model - for model in models.get("data", []) - if ":free" not in model.get("id", "").lower() - ] + models_response, embeddings_response = await asyncio.gather( + client.get(url), client.get(embeddings_url), return_exceptions=True + ) + + all_models = [] + + def process_models_response( + response: httpx.Response | BaseException, + ) -> list[dict]: + if not isinstance(response, BaseException): + response.raise_for_status() + data = response.json() + return [ + model + for model in data.get("data", []) + if ":free" not in model.get("id", "").lower() + ] + return [] + + all_models.extend(process_models_response(models_response)) + all_models.extend(process_models_response(embeddings_response)) + + return all_models async def _fetch_provider_models(self) -> dict: """Fetch models from provider's API.""" diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index 3ee58125..1736f1d7 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -50,7 +50,13 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): async def fetch_models(self) -> list[Model]: """Fetch all OpenRouter models.""" models_data = await async_fetch_openrouter_models() - return [Model(**model) for model in models_data] # type: ignore + models = [Model(**model) for model in models_data] # type: ignore + # manual alias for openai/text-embedding-ada-002 due to openrouter api bug + for model in models: + if model.id == "openai/text-embedding-ada-002": + model.alias_ids = ["text-embedding-ada-002-v2"] + break + return models async def get_balance(self) -> float | None: """Get the current account balance from OpenRouter. diff --git a/tests/integration/test_embeddings.py b/tests/integration/test_embeddings.py new file mode 100644 index 00000000..fe6ec922 --- /dev/null +++ b/tests/integration/test_embeddings.py @@ -0,0 +1,97 @@ +import json +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from httpx import AsyncClient + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_proxy_embeddings_endpoint(authenticated_client: AsyncClient) -> None: + """Test the embeddings endpoint proxy functionality""" + + test_payload = { + "model": "text-embedding-ada-002", + "input": "The quick brown fox", + } + + mock_response_data = { + "object": "list", + "data": [ + {"object": "embedding", "embedding": [0.0023, -0.0012, 0.0045], "index": 0} + ], + "model": "text-embedding-ada-002", + "usage": {"prompt_tokens": 5, "total_tokens": 5}, + } + + with patch("httpx.AsyncClient.send") as mock_send: + # Create a proper async generator for iter_bytes + async def mock_iter_bytes(*args: Any, **kwargs: Any) -> Any: + yield json.dumps(mock_response_data).encode() + + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.text = json.dumps(mock_response_data) + # Use MagicMock for synchronous .json() method + mock_response.json = MagicMock(return_value=mock_response_data) + mock_response.iter_bytes = mock_iter_bytes + mock_response.aiter_bytes = mock_iter_bytes + mock_send.return_value = mock_response + + # Make POST request to embeddings endpoint + response = await authenticated_client.post("/v1/embeddings", json=test_payload) + + assert response.status_code == 200 + response_data = response.json() + assert response_data["object"] == "list" + assert len(response_data["data"]) == 1 + assert response_data["data"][0]["object"] == "embedding" + + # Verify request was forwarded + mock_send.assert_called_once() + forwarded_request = mock_send.call_args[0][0] + # Verify the path ends with embeddings + # Note: forwarded path might be full URL + assert str(forwarded_request.url).endswith("embeddings") + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_model_case_insensitivity(authenticated_client: AsyncClient) -> None: + """Test that model lookups are case insensitive""" + + # We'll use a mixed-case model ID that should match the lowercase one in the system + # We assume 'gpt-3.5-turbo' is available in the mock env/database + + test_payload = { + "model": "GPT-3.5-TURBO", + "messages": [{"role": "user", "content": "Hello"}], + } + + with patch("httpx.AsyncClient.send") as mock_send: + mock_response_data = { + "id": "chatcmpl-123", + "object": "chat.completion", + "choices": [{"message": {"content": "Hi"}}], + "usage": {"total_tokens": 10}, + } + + async def mock_iter_bytes(*args: Any, **kwargs: Any) -> Any: + yield json.dumps(mock_response_data).encode() + + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.text = json.dumps(mock_response_data) + mock_response.json = MagicMock(return_value=mock_response_data) + mock_response.iter_bytes = mock_iter_bytes + mock_response.aiter_bytes = mock_iter_bytes + mock_send.return_value = mock_response + + response = await authenticated_client.post( + "/v1/chat/completions", json=test_payload + ) + + assert response.status_code == 200