2026-07-19 19:21:49 +02:00
|
|
|
|
"""Embedding pipeline using LiteLLM with OpenRouter for embeddings."""
|
2026-07-18 11:21:51 +02:00
|
|
|
|
|
|
|
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
|
|
import os
|
|
|
|
|
|
import logging
|
|
|
|
|
|
import uuid
|
2026-07-19 02:22:25 +02:00
|
|
|
|
from typing import Any, TYPE_CHECKING
|
2026-07-18 11:21:51 +02:00
|
|
|
|
|
|
|
|
|
|
import litellm
|
|
|
|
|
|
|
|
|
|
|
|
if TYPE_CHECKING:
|
|
|
|
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
|
|
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
|
|
|
|
|
|
MAX_INPUT_CHARS = 8000
|
|
|
|
|
|
|
2026-07-19 19:21:49 +02:00
|
|
|
|
# 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)
|
2026-07-18 11:21:51 +02:00
|
|
|
|
|
|
|
|
|
|
|
2026-07-19 02:22:25 +02:00
|
|
|
|
async def _get_api_credentials(
|
|
|
|
|
|
db: "AsyncSession | None", tenant_id: "uuid.UUID | None"
|
|
|
|
|
|
) -> tuple[str | None, str | None, str | None]:
|
2026-07-19 19:21:49 +02:00
|
|
|
|
"""Get API key, base_url and provider_type for embeddings.
|
2026-07-19 02:22:25 +02:00
|
|
|
|
|
2026-07-19 19:21:49 +02:00
|
|
|
|
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)
|
2026-07-19 02:22:25 +02:00
|
|
|
|
"""
|
2026-07-19 19:21:49 +02:00
|
|
|
|
# OpenRouter is the primary embedding provider
|
|
|
|
|
|
if OPENROUTER_API_KEY:
|
|
|
|
|
|
return OPENROUTER_API_KEY, OPENROUTER_BASE_URL, 'openai'
|
|
|
|
|
|
|
|
|
|
|
|
# Fallback to DB provider
|
2026-07-19 02:22:25 +02:00
|
|
|
|
if db and tenant_id:
|
|
|
|
|
|
try:
|
|
|
|
|
|
from app.plugins.builtins.ai_assistant.services import get_default_provider
|
|
|
|
|
|
provider = await get_default_provider(db, tenant_id)
|
|
|
|
|
|
if provider and provider.api_key:
|
|
|
|
|
|
return provider.api_key, provider.base_url, provider.provider_type
|
|
|
|
|
|
except Exception:
|
|
|
|
|
|
logger.debug("Failed to get provider from DB, falling back to env")
|
|
|
|
|
|
env_key = os.environ.get('API_KEY_OLLAMA_CLOUD', '')
|
|
|
|
|
|
return (env_key if env_key else None), None, None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _build_model(model: str, provider_type: str | None) -> str:
|
|
|
|
|
|
"""Build litellm model string with provider prefix."""
|
|
|
|
|
|
if provider_type:
|
|
|
|
|
|
model_parts = model.split("/", 1)
|
|
|
|
|
|
return f"{provider_type}/{model_parts[-1]}"
|
|
|
|
|
|
return model
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def generate_embedding(
|
|
|
|
|
|
text: str,
|
|
|
|
|
|
model: str | None = None,
|
|
|
|
|
|
db: "AsyncSession | None" = None,
|
|
|
|
|
|
tenant_id: "uuid.UUID | None" = None,
|
|
|
|
|
|
) -> list[float]:
|
2026-07-18 11:21:51 +02:00
|
|
|
|
"""Generate a single embedding via LiteLLM.
|
|
|
|
|
|
|
2026-07-19 19:21:49 +02:00
|
|
|
|
Uses OpenRouter with text-embedding-3-small (768 dimensions).
|
|
|
|
|
|
|
2026-07-18 11:21:51 +02:00
|
|
|
|
Args:
|
|
|
|
|
|
text: Input text (truncated to 8000 chars).
|
2026-07-19 19:21:49 +02:00
|
|
|
|
model: Embedding model name (default: openai/text-embedding-3-small).
|
2026-07-19 02:22:25 +02:00
|
|
|
|
db: Optional DB session for API key lookup.
|
|
|
|
|
|
tenant_id: Optional tenant ID for API key lookup.
|
2026-07-18 11:21:51 +02:00
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
Embedding vector as list of floats.
|
|
|
|
|
|
"""
|
2026-07-19 19:21:49 +02:00
|
|
|
|
model = model or OPENROUTER_EMBEDDING_MODEL
|
2026-07-18 11:21:51 +02:00
|
|
|
|
truncated = text[:MAX_INPUT_CHARS]
|
|
|
|
|
|
try:
|
2026-07-19 02:22:25 +02:00
|
|
|
|
api_key, api_base, provider_type = await _get_api_credentials(db, tenant_id)
|
|
|
|
|
|
litellm_model = _build_model(model, provider_type)
|
|
|
|
|
|
|
|
|
|
|
|
litellm_kwargs: dict[str, Any] = dict(
|
|
|
|
|
|
model=litellm_model,
|
2026-07-18 11:21:51 +02:00
|
|
|
|
input=truncated,
|
|
|
|
|
|
)
|
2026-07-19 02:22:25 +02:00
|
|
|
|
if api_key:
|
|
|
|
|
|
litellm_kwargs["api_key"] = api_key
|
|
|
|
|
|
if api_base:
|
|
|
|
|
|
litellm_kwargs["api_base"] = api_base
|
|
|
|
|
|
|
2026-07-19 19:21:49 +02:00
|
|
|
|
# Request 768 dimensions to match DB vector(768) column
|
|
|
|
|
|
if 'text-embedding-3' in litellm_model:
|
|
|
|
|
|
litellm_kwargs['dimensions'] = EMBEDDING_DIMENSIONS
|
|
|
|
|
|
|
2026-07-19 02:22:25 +02:00
|
|
|
|
response = await litellm.aembedding(**litellm_kwargs)
|
2026-07-18 11:21:51 +02:00
|
|
|
|
return response.data[0]["embedding"]
|
|
|
|
|
|
except Exception:
|
2026-07-19 19:21:49 +02:00
|
|
|
|
logger.warning("Failed to generate embedding", exc_info=True)
|
2026-07-18 11:21:51 +02:00
|
|
|
|
return []
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def generate_embeddings_batch(
|
2026-07-19 02:22:25 +02:00
|
|
|
|
texts: list[str],
|
|
|
|
|
|
model: str | None = None,
|
|
|
|
|
|
db: "AsyncSession | None" = None,
|
|
|
|
|
|
tenant_id: "uuid.UUID | None" = None,
|
2026-07-18 11:21:51 +02:00
|
|
|
|
) -> list[list[float]]:
|
|
|
|
|
|
"""Generate embeddings for multiple texts in a single API call.
|
|
|
|
|
|
|
|
|
|
|
|
Args:
|
|
|
|
|
|
texts: List of input texts.
|
2026-07-19 19:21:49 +02:00
|
|
|
|
model: Embedding model name (default: openai/text-embedding-3-small).
|
2026-07-19 02:22:25 +02:00
|
|
|
|
db: Optional DB session for API key lookup.
|
|
|
|
|
|
tenant_id: Optional tenant ID for API key lookup.
|
2026-07-18 11:21:51 +02:00
|
|
|
|
|
|
|
|
|
|
Returns:
|
|
|
|
|
|
List of embedding vectors.
|
|
|
|
|
|
"""
|
2026-07-19 19:21:49 +02:00
|
|
|
|
model = model or OPENROUTER_EMBEDDING_MODEL
|
2026-07-18 11:21:51 +02:00
|
|
|
|
truncated = [t[:MAX_INPUT_CHARS] for t in texts]
|
|
|
|
|
|
try:
|
2026-07-19 02:22:25 +02:00
|
|
|
|
api_key, api_base, provider_type = await _get_api_credentials(db, tenant_id)
|
|
|
|
|
|
litellm_model = _build_model(model, provider_type)
|
|
|
|
|
|
|
|
|
|
|
|
litellm_kwargs: dict[str, Any] = dict(
|
|
|
|
|
|
model=litellm_model,
|
2026-07-18 11:21:51 +02:00
|
|
|
|
input=truncated,
|
|
|
|
|
|
)
|
2026-07-19 02:22:25 +02:00
|
|
|
|
if api_key:
|
|
|
|
|
|
litellm_kwargs["api_key"] = api_key
|
|
|
|
|
|
if api_base:
|
|
|
|
|
|
litellm_kwargs["api_base"] = api_base
|
|
|
|
|
|
|
2026-07-19 19:21:49 +02:00
|
|
|
|
if 'text-embedding-3' in litellm_model:
|
|
|
|
|
|
litellm_kwargs['dimensions'] = EMBEDDING_DIMENSIONS
|
|
|
|
|
|
|
2026-07-19 02:22:25 +02:00
|
|
|
|
response = await litellm.aembedding(**litellm_kwargs)
|
2026-07-18 11:21:51 +02:00
|
|
|
|
return [d["embedding"] for d in response.data]
|
|
|
|
|
|
except Exception:
|
2026-07-19 02:22:25 +02:00
|
|
|
|
logger.warning("Failed to generate batch embeddings", exc_info=True)
|
2026-07-18 11:21:51 +02:00
|
|
|
|
return [[] for _ in texts]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def index_entity(
|
|
|
|
|
|
entity_type: str,
|
|
|
|
|
|
entity_id: uuid.UUID,
|
|
|
|
|
|
tenant_id: uuid.UUID,
|
2026-07-19 02:22:25 +02:00
|
|
|
|
db: "AsyncSession",
|
2026-07-18 11:21:51 +02:00
|
|
|
|
) -> bool:
|
|
|
|
|
|
"""Generate and store embedding for a single entity.
|
|
|
|
|
|
|
|
|
|
|
|
Uses the provider registry to get embedding text, generates embedding,
|
|
|
|
|
|
and updates the entity's embedding column.
|
|
|
|
|
|
|
|
|
|
|
|
Returns True on success, False on failure.
|
|
|
|
|
|
"""
|
|
|
|
|
|
from app.plugins.builtins.unified_search.provider_registry import get_search_registry
|
|
|
|
|
|
|
|
|
|
|
|
registry = get_search_registry()
|
|
|
|
|
|
provider = registry.get(entity_type)
|
|
|
|
|
|
if provider is None:
|
|
|
|
|
|
logger.warning("No provider for entity_type=%s", entity_type)
|
|
|
|
|
|
return False
|
|
|
|
|
|
|
|
|
|
|
|
try:
|
|
|
|
|
|
text = await provider.get_embedding_text(db, entity_id, tenant_id)
|
|
|
|
|
|
if not text.strip():
|
|
|
|
|
|
logger.debug("Empty embedding text for %s/%s", entity_type, entity_id)
|
|
|
|
|
|
return False
|
|
|
|
|
|
|
2026-07-19 02:22:25 +02:00
|
|
|
|
embedding = await generate_embedding(text, db=db, tenant_id=tenant_id)
|
2026-07-18 11:21:51 +02:00
|
|
|
|
if not embedding:
|
|
|
|
|
|
return False
|
|
|
|
|
|
|
|
|
|
|
|
# Update the entity's embedding column
|
|
|
|
|
|
from sqlalchemy import text as sql_text
|
|
|
|
|
|
|
|
|
|
|
|
table_map = {
|
|
|
|
|
|
"contact": "contacts",
|
|
|
|
|
|
"company": "companies",
|
|
|
|
|
|
"mail": "mails",
|
|
|
|
|
|
"file": "files",
|
|
|
|
|
|
"event": "calendar_entries",
|
|
|
|
|
|
}
|
|
|
|
|
|
table = table_map.get(entity_type)
|
|
|
|
|
|
if not table:
|
|
|
|
|
|
logger.warning("Unknown entity_type=%s for embedding storage", entity_type)
|
|
|
|
|
|
return False
|
|
|
|
|
|
|
|
|
|
|
|
sql = sql_text(
|
|
|
|
|
|
f"UPDATE {table} SET embedding = cast(:emb AS vector) "
|
|
|
|
|
|
f"WHERE id = :eid AND tenant_id = :tid"
|
|
|
|
|
|
)
|
|
|
|
|
|
await db.execute(
|
|
|
|
|
|
sql,
|
|
|
|
|
|
{"emb": str(embedding), "eid": entity_id, "tid": tenant_id},
|
|
|
|
|
|
)
|
|
|
|
|
|
await db.commit()
|
|
|
|
|
|
return True
|
|
|
|
|
|
except Exception:
|
|
|
|
|
|
logger.exception("Failed to index entity %s/%s", entity_type, entity_id)
|
|
|
|
|
|
return False
|