""" GCAS Search Engine – Provider-agnostic embedding layer Supported providers ------------------- "local" – sentence-transformers (no API key required, runs on CPU) "openai" – OpenAI Embeddings API (requires OPENAI_API_KEY) """ from __future__ import annotations import logging from typing import List, Optional import numpy as np from config import settings logger = logging.getLogger(__name__) # Lazy-loaded singleton for the local model so the heavy import # only happens once per worker process. _local_model = None def _get_local_model(): global _local_model if _local_model is None: try: from sentence_transformers import SentenceTransformer except ImportError as exc: raise ImportError( "sentence-transformers is required for local embeddings. " "Run: pip install sentence-transformers" ) from exc logger.info("Loading local embedding model: %s", settings.local_embedding_model) _local_model = SentenceTransformer(settings.local_embedding_model) logger.info("Local model ready.") return _local_model def embed_texts( texts: List[str], *, provider: Optional[str] = None, api_key: Optional[str] = None, batch_size: int = 512, ) -> np.ndarray: """ Embed a list of strings into L2-normalised float32 vectors. Parameters ---------- texts : list of strings to embed provider : "local" | "openai" – overrides config when specified api_key : override API key (only for "openai") batch_size : chunk size for local model (ignored for OpenAI) Returns ------- np.ndarray of shape (len(texts), embedding_dim), dtype float32 """ provider = provider or settings.embedding_provider if provider == "local": model = _get_local_model() embeddings = model.encode( texts, batch_size=batch_size, show_progress_bar=len(texts) > 200, normalize_embeddings=True, convert_to_numpy=True, ) return embeddings.astype(np.float32) elif provider == "openai": try: from openai import OpenAI except ImportError as exc: raise ImportError("openai package required: pip install openai") from exc key = api_key or settings.openai_api_key if not key: raise ValueError( "OPENAI_API_KEY is not set. " "Set it in .env or pass api_key in the request." ) client = OpenAI(api_key=key) # OpenAI allows up to 2 048 texts per call; batch for safety. all_embeddings: List[np.ndarray] = [] chunk = 512 for i in range(0, len(texts), chunk): batch = texts[i : i + chunk] response = client.embeddings.create( model=settings.openai_embedding_model, input=batch, ) vecs = np.array([e.embedding for e in response.data], dtype=np.float32) # L2 normalise norms = np.linalg.norm(vecs, axis=1, keepdims=True) norms = np.where(norms == 0, 1.0, norms) all_embeddings.append(vecs / norms) return np.vstack(all_embeddings) else: raise ValueError( f"Unknown embedding provider: '{provider}'. Choose 'local' or 'openai'." ) def embed_query( query: str, *, provider: Optional[str] = None, api_key: Optional[str] = None, ) -> np.ndarray: """Convenience wrapper – embed a single query string.""" return embed_texts([query], provider=provider, api_key=api_key)[0]