From 195da0c9da5633e29f239ee726f609d06f8a1b91 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 3 Dec 2025 21:58:31 +0100 Subject: [PATCH 1/7] embedding integration --- routstr/payment/models.py | 30 ++++++++++++++++++++++++------ routstr/upstream/base.py | 33 +++++++++++++++++++++++++-------- 2 files changed, 49 insertions(+), 14 deletions(-) diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 849cf103..e05635d4 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -111,12 +111,30 @@ 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): + if not isinstance(response, Exception): + 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: @@ -138,9 +156,9 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis ): 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/upstream/base.py b/routstr/upstream/base.py index 6b908b8e..205589ba 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -1767,15 +1767,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): + if not isinstance(response, Exception): + 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.""" From 5d2219880d90df948bbb3edb5d099002e9897eb1 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 3 Dec 2025 22:20:31 +0100 Subject: [PATCH 2/7] embedding integration --- routstr/payment/models.py | 8 ++-- routstr/proxy.py | 2 +- routstr/upstream/base.py | 99 ++++++++++++++++++++------------------- 3 files changed, 57 insertions(+), 52 deletions(-) diff --git a/routstr/payment/models.py b/routstr/payment/models.py index e05635d4..9c385feb 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -114,11 +114,13 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis 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 + return_exceptions=True, ) - def process_models_response(response): - if not isinstance(response, Exception): + def process_models_response( + response: httpx.Response | BaseException, + ) -> list[dict]: + if not isinstance(response, BaseException): response.raise_for_status() data = response.json() 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 205589ba..4815251b 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -721,51 +721,54 @@ 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" - ) + # Handle endpoints that require cost calculation and payment adjustment + 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 @@ -1506,9 +1509,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}, ) @@ -1771,15 +1774,15 @@ class BaseUpstreamProvider: async with httpx.AsyncClient(timeout=30.0) as client: models_response, embeddings_response = await asyncio.gather( - client.get(url), - client.get(embeddings_url), - return_exceptions=True + client.get(url), client.get(embeddings_url), return_exceptions=True ) all_models = [] - def process_models_response(response): - if not isinstance(response, Exception): + def process_models_response( + response: httpx.Response | BaseException, + ) -> list[dict]: + if not isinstance(response, BaseException): response.raise_for_status() data = response.json() return [ From 9438bc957fabb9aa2cc521f59db407d58d1ed129 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 3 Dec 2025 22:24:43 +0100 Subject: [PATCH 3/7] clean up --- routstr/upstream/base.py | 1 - 1 file changed, 1 deletion(-) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 4815251b..954a82cd 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -721,7 +721,6 @@ class BaseUpstreamProvider: await client.aclose() return mapped_error - # Handle endpoints that require cost calculation and payment adjustment if path.endswith("chat/completions") or path.endswith("embeddings"): if path.endswith("chat/completions"): client_wants_streaming = False From 547365894d3fc76ab6206d2c99b1fb47d9d84c02 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Thu, 11 Dec 2025 13:58:57 +0800 Subject: [PATCH 4/7] add simple test --- tests/integration/test_embeddings.py | 97 ++++++++++++++++++++++++++++ 1 file changed, 97 insertions(+) create mode 100644 tests/integration/test_embeddings.py 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 From 43e97326e035ff416a7536e603ae16050ff78c32 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 12 Dec 2025 20:14:43 +0100 Subject: [PATCH 5/7] fix model naming issue in response --- routstr/proxy.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/routstr/proxy.py b/routstr/proxy.py index edceaaa6..38224ea8 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -65,7 +65,11 @@ 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()) + return next( + (model for alias, model in _model_instances.items() + if alias.lower() in model_id.lower()), + None + ) def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None: From cc42534a9711033cd80ed4fa0d59c4f3bbaa2612 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Mon, 22 Dec 2025 11:30:51 +0100 Subject: [PATCH 6/7] manual alias due to openrouter api bug --- routstr/upstream/openrouter.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) 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. From dd4ed7541f6fabc2ba9380618392ae81d7c86b30 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Mon, 22 Dec 2025 11:31:19 +0100 Subject: [PATCH 7/7] undo 43e9732 --- routstr/proxy.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/routstr/proxy.py b/routstr/proxy.py index 38224ea8..edceaaa6 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -65,11 +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 next( - (model for alias, model in _model_instances.items() - if alias.lower() in model_id.lower()), - None - ) + return _model_instances.get(model_id.lower()) def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None: