DocDoeAI / app /services /embedding_service.py
asnannp's picture
Deploy backend cd4237ff: support routes + rate limit + exam_date nullable + upload 413 fix
7c6ffa6
Raw
History Blame Contribute Delete
5.7 kB
"""
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("<I", i)).digest()
# Convert first 4 bytes to a float in [-1, 1]
raw = struct.unpack("<I", h[:4])[0]
vec.append((raw / 0xFFFFFFFF) * 2 - 1)
return vec
# ---------------------------------------------------------------------------
# API embedding call
# ---------------------------------------------------------------------------
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()
# 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