Spaces:
Runtime error
Runtime error
| """Embedding abstraction for the v2 backend. | |
| Default backend is a free, local sentence-transformers model (MiniLM). Setting | |
| ``embedding_provider='openai'`` switches to OpenAI embeddings. The active | |
| embedder exposes ``embed_dim`` so the FAISS store can size its index, and a | |
| deterministic fake backend is used when neither dependency is available (tests). | |
| """ | |
| from __future__ import annotations | |
| import hashlib | |
| import logging | |
| import struct | |
| from backend.config import settings | |
| logger = logging.getLogger(__name__) | |
| class Embedder: | |
| """Minimal embedding interface used by the RAG store.""" | |
| embed_dim: int | |
| def embed_documents(self, texts: list[str]) -> list[list[float]]: # pragma: no cover - interface | |
| raise NotImplementedError | |
| def embed_query(self, text: str) -> list[float]: # pragma: no cover - interface | |
| raise NotImplementedError | |
| class LocalEmbedder(Embedder): | |
| """sentence-transformers backend (default).""" | |
| def __init__(self, model_name: str) -> None: | |
| from sentence_transformers import SentenceTransformer | |
| self._model = SentenceTransformer(model_name) | |
| self.embed_dim = int(self._model.get_sentence_embedding_dimension()) | |
| logger.info("LocalEmbedder loaded %s (dim=%d)", model_name, self.embed_dim) | |
| def embed_documents(self, texts: list[str]) -> list[list[float]]: | |
| vecs = self._model.encode( | |
| list(texts), normalize_embeddings=True, convert_to_numpy=True | |
| ) | |
| return [v.tolist() for v in vecs] | |
| def embed_query(self, text: str) -> list[float]: | |
| return self.embed_documents([text])[0] | |
| class OpenAIEmbedder(Embedder): | |
| """OpenAI embeddings backend.""" | |
| _DIMS = { | |
| "text-embedding-3-small": 1536, | |
| "text-embedding-3-large": 3072, | |
| "text-embedding-ada-002": 1536, | |
| } | |
| def __init__(self, model_name: str, api_key: str) -> None: | |
| from openai import OpenAI | |
| self._client = OpenAI( | |
| api_key=api_key, | |
| timeout=float(settings.openai_request_timeout_seconds), | |
| ) | |
| self._model = model_name | |
| self.embed_dim = self._DIMS.get(model_name, 1536) | |
| logger.info("OpenAIEmbedder using %s (dim=%d)", model_name, self.embed_dim) | |
| def embed_documents(self, texts: list[str]) -> list[list[float]]: | |
| from backend.llm.openai_client import pipeline_timeout | |
| resp = self._client.embeddings.create( | |
| model=self._model, | |
| input=list(texts), | |
| timeout=pipeline_timeout(), | |
| ) | |
| return [d.embedding for d in resp.data] | |
| def embed_query(self, text: str) -> list[float]: | |
| return self.embed_documents([text])[0] | |
| class FakeEmbedder(Embedder): | |
| """Deterministic hash-based embedder for offline tests.""" | |
| def __init__(self, dim: int = 384) -> None: | |
| self.embed_dim = dim | |
| def _vec(self, text: str) -> list[float]: | |
| out: list[float] = [] | |
| seed = (text or "").encode("utf-8") | |
| counter = 0 | |
| while len(out) < self.embed_dim: | |
| h = hashlib.sha256(seed + struct.pack(">I", counter)).digest() | |
| for i in range(0, len(h), 4): | |
| if len(out) >= self.embed_dim: | |
| break | |
| (val,) = struct.unpack(">I", h[i : i + 4]) | |
| out.append((val / 0xFFFFFFFF) * 2.0 - 1.0) | |
| counter += 1 | |
| norm = sum(x * x for x in out) ** 0.5 or 1.0 | |
| return [x / norm for x in out] | |
| def embed_documents(self, texts: list[str]) -> list[list[float]]: | |
| return [self._vec(t) for t in texts] | |
| def embed_query(self, text: str) -> list[float]: | |
| return self._vec(text) | |
| _instance: Embedder | None = None | |
| def get_embedder() -> Embedder: | |
| """Return the cached embedder singleton selected by configuration.""" | |
| global _instance | |
| if _instance is not None: | |
| return _instance | |
| provider = (settings.embedding_provider or "local").lower() | |
| try: | |
| if provider == "openai" and settings.openai_api_key: | |
| _instance = OpenAIEmbedder(settings.openai_embedding_model, settings.openai_api_key) | |
| else: | |
| _instance = LocalEmbedder(settings.local_embedding_model) | |
| except Exception as exc: # noqa: BLE001 - dependency/model load failures → fake | |
| logger.warning("Embedder init failed (%s); using FakeEmbedder (TEST ONLY).", exc) | |
| _instance = FakeEmbedder() | |
| return _instance | |
| def reset_embedder() -> None: | |
| """Reset the cached embedder (tests / config reloads).""" | |
| global _instance | |
| _instance = None | |