From 8df0c17bc3fc8a4c943f2c8e33930d7355974d1b Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 3 Dec 2025 22:20:31 +0100 Subject: [PATCH] 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 d14b6c58..7a76ff5f 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -119,11 +119,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 c39b4e91..a8f63185 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -734,51 +734,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 @@ -1519,9 +1522,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}, ) @@ -1784,15 +1787,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 [