""" Embedding generation service for DocDoe AI. Provider cascade: 1. OpenRouter API (openai/text-embedding-3-small, 1536 dims) 2. Sarvam AI (same OpenAI-compatible /embeddings endpoint) 3. Local fallback (deterministic bag-of-words hasher — never crashes) """ from __future__ import annotations import hashlib import json import logging import struct from typing import Any import httpx from app.core.config import get_settings logger = logging.getLogger(__name__) _EMBEDDING_DIM = 1536 _BATCH_SIZE = 20 _EMBEDDING_MODEL = "openai/text-embedding-3-small" _TIMEOUT_SECONDS = 60 # --------------------------------------------------------------------------- # Public helpers # --------------------------------------------------------------------------- def embedding_to_str(vec: list[float]) -> str: """Serialise an embedding vector to a JSON string for DB storage.""" return json.dumps(vec) def str_to_embedding(s: str) -> list[float]: """Deserialise a JSON string back to an embedding vector.""" return json.loads(s) # --------------------------------------------------------------------------- # Local deterministic fallback (offline dev mode) # --------------------------------------------------------------------------- def _local_hash_embedding(text: str) -> list[float]: """ Produce a deterministic pseudo-vector from text using SHA-256 seeded bag-of-words hashing. Not semantically meaningful — only guarantees that the same input always yields the same 1536-float vector so that upstream code never crashes when no API key is available. """ digest = hashlib.sha256(text.encode("utf-8")).digest() # Expand the 32-byte digest to fill 1536 floats deterministically vec: list[float] = [] for i in range(_EMBEDDING_DIM): # Re-hash with index to get unique bytes per dimension h = hashlib.md5(digest + struct.pack(" list[list[float]]: """Call an OpenAI-compatible /embeddings endpoint.""" headers: dict[str, str] = { "Authorization": f"Bearer {api_key}", "Content-Type": "application/json", } if extra_headers: headers.update(extra_headers) payload: dict[str, Any] = { "model": _EMBEDDING_MODEL, "input": texts, } async with httpx.AsyncClient(timeout=_TIMEOUT_SECONDS) as client: response = await client.post( f"{base_url}/embeddings", headers=headers, json=payload, ) response.raise_for_status() data = response.json() # OpenAI format: data.data[i].embedding embeddings_data = sorted(data["data"], key=lambda d: d["index"]) return [item["embedding"] for item in embeddings_data] # --------------------------------------------------------------------------- # Main public function # --------------------------------------------------------------------------- async def generate_embeddings(texts: list[str]) -> list[list[float]]: """ Generate embedding vectors for a list of texts. Tries OpenRouter first, then Sarvam, then falls back to a local deterministic hasher so offline development never crashes. """ if not texts: return [] settings = get_settings() # Determine provider provider: str | None = None base_url: str = "" api_key: str = "" extra_headers: dict[str, str] = {} if settings.openrouter_api_key: provider = "openrouter" base_url = settings.openrouter_base_url.rstrip("/") api_key = settings.openrouter_api_key extra_headers = { "HTTP-Referer": settings.openrouter_site_url, "X-Title": settings.openrouter_app_name, } elif settings.sarvam_api_key: provider = "sarvam" base_url = settings.sarvam_base_url.rstrip("/") api_key = settings.sarvam_api_key if provider is None: if settings.environment == "production" or not settings.ai_fallback_to_mock: raise RuntimeError("No embedding API key configured") logger.info( "No embedding API key configured — using local hash fallback for %d texts", len(texts), ) return [_local_hash_embedding(t) for t in texts] # Batch and call API all_embeddings: list[list[float]] = [] for batch_start in range(0, len(texts), _BATCH_SIZE): batch = texts[batch_start : batch_start + _BATCH_SIZE] try: batch_embeddings = await _call_embedding_api( batch, base_url, api_key, extra_headers ) all_embeddings.extend(batch_embeddings) except Exception as exc: if settings.environment == "production" or not settings.ai_fallback_to_mock: raise RuntimeError(f"Embedding generation failed: {exc}") from exc logger.warning( "Embedding API (%s) failed for batch starting at index %d, " "falling back to local hash", provider, batch_start, exc_info=True, ) all_embeddings.extend([_local_hash_embedding(t) for t in batch]) return all_embeddings