feat: use OpenRouter for embeddings with text-embedding-3-small (768 dims)
- Ollama Cloud has no embedding endpoint, OpenRouter does - text-embedding-3-small with dimensions=768 matches DB vector(768) column - API_KEY_OPENROUTER env var is primary, DB provider is fallback - Added dimensions parameter for text-embedding-3 models
This commit is contained in:
@@ -1,4 +1,4 @@
|
||||
"""Embedding pipeline using LiteLLM with configurable Ollama Cloud provider."""
|
||||
"""Embedding pipeline using LiteLLM with OpenRouter for embeddings."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -16,17 +16,28 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
MAX_INPUT_CHARS = 8000
|
||||
|
||||
DEFAULT_EMBEDDING_MODEL = os.environ.get('SEARCH_EMBEDDING_MODEL', 'ollama/nomic-embed-text')
|
||||
# OpenRouter for embeddings (Ollama Cloud has no embedding endpoint)
|
||||
OPENROUTER_API_KEY = os.environ.get('API_KEY_OPENROUTER', '')
|
||||
OPENROUTER_BASE_URL = 'https://openrouter.ai/api/v1'
|
||||
OPENROUTER_EMBEDDING_MODEL = os.environ.get('SEARCH_EMBEDDING_MODEL', 'openai/text-embedding-3-small')
|
||||
EMBEDDING_DIMENSIONS = 768 # Must match DB column vector(768)
|
||||
|
||||
|
||||
async def _get_api_credentials(
|
||||
db: "AsyncSession | None", tenant_id: "uuid.UUID | None"
|
||||
) -> tuple[str | None, str | None, str | None]:
|
||||
"""Get API key, base_url and provider_type from the default AI provider in DB.
|
||||
"""Get API key, base_url and provider_type for embeddings.
|
||||
|
||||
Falls back to API_KEY_OLLAMA_CLOUD env var.
|
||||
Returns (api_key, base_url, provider_type).
|
||||
Priority:
|
||||
1. OpenRouter env var (API_KEY_OPENROUTER) – dedicated embedding provider
|
||||
2. Default AI provider from DB (fallback)
|
||||
3. API_KEY_OLLAMA_CLOUD env var (last resort)
|
||||
"""
|
||||
# OpenRouter is the primary embedding provider
|
||||
if OPENROUTER_API_KEY:
|
||||
return OPENROUTER_API_KEY, OPENROUTER_BASE_URL, 'openai'
|
||||
|
||||
# Fallback to DB provider
|
||||
if db and tenant_id:
|
||||
try:
|
||||
from app.plugins.builtins.ai_assistant.services import get_default_provider
|
||||
@@ -55,16 +66,18 @@ async def generate_embedding(
|
||||
) -> list[float]:
|
||||
"""Generate a single embedding via LiteLLM.
|
||||
|
||||
Uses OpenRouter with text-embedding-3-small (768 dimensions).
|
||||
|
||||
Args:
|
||||
text: Input text (truncated to 8000 chars).
|
||||
model: Embedding model name (default: ollama/nomic-embed-text).
|
||||
model: Embedding model name (default: openai/text-embedding-3-small).
|
||||
db: Optional DB session for API key lookup.
|
||||
tenant_id: Optional tenant ID for API key lookup.
|
||||
|
||||
Returns:
|
||||
Embedding vector as list of floats.
|
||||
"""
|
||||
model = model or DEFAULT_EMBEDDING_MODEL
|
||||
model = model or OPENROUTER_EMBEDDING_MODEL
|
||||
truncated = text[:MAX_INPUT_CHARS]
|
||||
try:
|
||||
api_key, api_base, provider_type = await _get_api_credentials(db, tenant_id)
|
||||
@@ -79,10 +92,14 @@ async def generate_embedding(
|
||||
if api_base:
|
||||
litellm_kwargs["api_base"] = api_base
|
||||
|
||||
# Request 768 dimensions to match DB vector(768) column
|
||||
if 'text-embedding-3' in litellm_model:
|
||||
litellm_kwargs['dimensions'] = EMBEDDING_DIMENSIONS
|
||||
|
||||
response = await litellm.aembedding(**litellm_kwargs)
|
||||
return response.data[0]["embedding"]
|
||||
except Exception:
|
||||
logger.warning("Failed to generate embedding (Ollama Cloud may not support embeddings)", exc_info=True)
|
||||
logger.warning("Failed to generate embedding", exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
@@ -96,14 +113,14 @@ async def generate_embeddings_batch(
|
||||
|
||||
Args:
|
||||
texts: List of input texts.
|
||||
model: Embedding model name (default: ollama/nomic-embed-text).
|
||||
model: Embedding model name (default: openai/text-embedding-3-small).
|
||||
db: Optional DB session for API key lookup.
|
||||
tenant_id: Optional tenant ID for API key lookup.
|
||||
|
||||
Returns:
|
||||
List of embedding vectors.
|
||||
"""
|
||||
model = model or DEFAULT_EMBEDDING_MODEL
|
||||
model = model or OPENROUTER_EMBEDDING_MODEL
|
||||
truncated = [t[:MAX_INPUT_CHARS] for t in texts]
|
||||
try:
|
||||
api_key, api_base, provider_type = await _get_api_credentials(db, tenant_id)
|
||||
@@ -118,6 +135,9 @@ async def generate_embeddings_batch(
|
||||
if api_base:
|
||||
litellm_kwargs["api_base"] = api_base
|
||||
|
||||
if 'text-embedding-3' in litellm_model:
|
||||
litellm_kwargs['dimensions'] = EMBEDDING_DIMENSIONS
|
||||
|
||||
response = await litellm.aembedding(**litellm_kwargs)
|
||||
return [d["embedding"] for d in response.data]
|
||||
except Exception:
|
||||
|
||||
Reference in New Issue
Block a user