mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
Merge pull request #251 from Routstr/embeddings-integration
Embeddings integration
This commit is contained in:
@@ -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
@@ -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
@@ -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."""
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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
|
||||||
Reference in New Issue
Block a user