| """ |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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) |
|
|
|
|
| |
| |
| |
|
|
| 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() |
| |
| vec: list[float] = [] |
| for i in range(_EMBEDDING_DIM): |
| |
| h = hashlib.md5(digest + struct.pack("<I", i)).digest() |
| |
| raw = struct.unpack("<I", h[:4])[0] |
| vec.append((raw / 0xFFFFFFFF) * 2 - 1) |
| return vec |
|
|
|
|
| |
| |
| |
|
|
| async def _call_embedding_api( |
| texts: list[str], |
| base_url: str, |
| api_key: str, |
| extra_headers: dict[str, str] | None = None, |
| ) -> 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() |
|
|
| |
| embeddings_data = sorted(data["data"], key=lambda d: d["index"]) |
| return [item["embedding"] for item in embeddings_data] |
|
|
|
|
| |
| |
| |
|
|
| 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() |
|
|
| |
| 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] |
|
|
| |
| 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 |
|
|