Merge pull request #251 from Routstr/embeddings-integration

Embeddings integration
This commit is contained in:
shroominic
2025-12-23 20:53:26 +01:00
committed by GitHub
5 changed files with 201 additions and 59 deletions
+26 -6
View File
@@ -93,12 +93,32 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis
try: try:
async with httpx.AsyncClient() as client: async with httpx.AsyncClient() as client:
response = await client.get(f"{base_url}/models", timeout=30) models_response, embeddings_response = await asyncio.gather(
response.raise_for_status() client.get(f"{base_url}/models", timeout=30),
data = response.json() 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] = [] 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", "") model_id = model.get("id", "")
if source_filter: 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): if not _has_valid_pricing(model):
continue continue
models_data.append(model) filtered_models.append(model)
return models_data return filtered_models
except Exception as e: except Exception as e:
logger.error(f"Error (async) fetching models from OpenRouter API: {e}") logger.error(f"Error (async) fetching models from OpenRouter API: {e}")
return [] return []
+1 -1
View File
@@ -65,7 +65,7 @@ def get_upstreams() -> list[BaseUpstreamProvider]:
def get_model_instance(model_id: str) -> Model | None: def get_model_instance(model_id: str) -> Model | None:
"""Get Model instance by ID from global cache.""" """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: def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None:
+70 -51
View File
@@ -734,51 +734,53 @@ class BaseUpstreamProvider:
await client.aclose() await client.aclose()
return mapped_error return mapped_error
if path.endswith("chat/completions"): if path.endswith("chat/completions") or path.endswith("embeddings"):
client_wants_streaming = False if path.endswith("chat/completions"):
if request_body: client_wants_streaming = False
try: if request_body:
request_data = json.loads(request_body) try:
client_wants_streaming = request_data.get("stream", False) request_data = json.loads(request_body)
logger.debug( client_wants_streaming = request_data.get("stream", False)
"Chat completion request analysis", logger.debug(
extra={ "Chat completion request analysis",
"client_wants_streaming": client_wants_streaming, extra={
"model": request_data.get("model", "unknown"), "client_wants_streaming": client_wants_streaming,
"key_hash": key.hashed_key[:8] + "...", "model": request_data.get("model", "unknown"),
}, "key_hash": key.hashed_key[:8] + "...",
) },
except json.JSONDecodeError: )
logger.warning( except json.JSONDecodeError:
"Failed to parse request body JSON for streaming detection" logger.warning(
) "Failed to parse request body JSON for streaming detection"
)
content_type = response.headers.get("content-type", "") content_type = response.headers.get("content-type", "")
upstream_is_streaming = "text/event-stream" in content_type upstream_is_streaming = "text/event-stream" in content_type
is_streaming = client_wants_streaming and upstream_is_streaming is_streaming = client_wants_streaming and upstream_is_streaming
logger.debug( logger.debug(
"Response type analysis", "Response type analysis",
extra={ extra={
"is_streaming": is_streaming, "is_streaming": is_streaming,
"client_wants_streaming": client_wants_streaming, "client_wants_streaming": client_wants_streaming,
"upstream_is_streaming": upstream_is_streaming, "upstream_is_streaming": upstream_is_streaming,
"content_type": content_type, "content_type": content_type,
"key_hash": key.hashed_key[:8] + "...", "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
) )
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: try:
return await self.handle_non_streaming_chat_completion( return await self.handle_non_streaming_chat_completion(
response, key, session, max_cost_for_model response, key, session, max_cost_for_model
@@ -1519,9 +1521,9 @@ class BaseUpstreamProvider:
error_response.headers["X-Cashu"] = refund_token error_response.headers["X-Cashu"] = refund_token
return error_response return error_response
if path.endswith("chat/completions"): if path.endswith("chat/completions") or path.endswith("embeddings"):
logger.debug( logger.debug(
"Processing chat completion response", "Processing completion/embeddings response",
extra={"path": path, "amount": amount, "unit": unit}, extra={"path": path, "amount": amount, "unit": unit},
) )
@@ -1770,15 +1772,32 @@ class BaseUpstreamProvider:
async def _fetch_openrouter_models(self) -> list[dict]: async def _fetch_openrouter_models(self) -> list[dict]:
"""Fetch models from OpenRouter API.""" """Fetch models from OpenRouter API."""
url = "https://openrouter.ai/api/v1/models" 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: async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.get(url) models_response, embeddings_response = await asyncio.gather(
response.raise_for_status() client.get(url), client.get(embeddings_url), return_exceptions=True
models = response.json() )
return [
model all_models = []
for model in models.get("data", [])
if ":free" not in model.get("id", "").lower() 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: async def _fetch_provider_models(self) -> dict:
"""Fetch models from provider's API.""" """Fetch models from provider's API."""
+7 -1
View File
@@ -50,7 +50,13 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider):
async def fetch_models(self) -> list[Model]: async def fetch_models(self) -> list[Model]:
"""Fetch all OpenRouter models.""" """Fetch all OpenRouter models."""
models_data = await async_fetch_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: async def get_balance(self) -> float | None:
"""Get the current account balance from OpenRouter. """Get the current account balance from OpenRouter.
+97
View File
@@ -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