Recruitment_Copilot / backend /services /embedding_service.py
Ashgen12's picture
Embedding fix + ingestion hardening
4fb3cce verified
Raw
History Blame Contribute Delete
2.3 kB
"""Gemini embeddings.
Switched to the modern ``google-genai`` SDK (already in requirements.txt) and
the current GA embedding model ``gemini-embedding-001``. The legacy
``google-generativeai`` SDK + ``text-embedding-004`` returns
``404 models/text-embedding-004 is not found for API version v1beta`` against
recent API regions, so we use the new SDK + new model.
"""
from __future__ import annotations
import logging
from google import genai
from google.genai import types
from core.config import get_settings
logger = logging.getLogger(__name__)
def _normalise_model(name: str) -> str:
name = (name or "").strip()
# Older config used ``models/text-embedding-004``; map it to the new GA model.
if not name or "text-embedding-004" in name:
return "gemini-embedding-001"
if name.startswith("models/"):
name = name.split("/", 1)[1]
return name
class EmbeddingService:
def __init__(self) -> None:
self.settings = get_settings()
self._client: genai.Client | None = None
self._model_name = _normalise_model(self.settings.embedding_model)
if self.settings.gemini_api_key:
self._client = genai.Client(api_key=self.settings.gemini_api_key)
@property
def is_ready(self) -> bool:
return self._client is not None
def embed_document(self, text: str) -> list[float]:
return self._embed(text=text, task_type="RETRIEVAL_DOCUMENT")
def embed_query(self, text: str) -> list[float]:
return self._embed(text=text, task_type="RETRIEVAL_QUERY")
def _embed(self, text: str, task_type: str) -> list[float]:
if self._client is None:
raise RuntimeError("GEMINI_API_KEY is required to generate embeddings.")
response = self._client.models.embed_content(
model=self._model_name,
contents=text,
config=types.EmbedContentConfig(task_type=task_type),
)
embeddings = getattr(response, "embeddings", None) or []
if not embeddings:
raise RuntimeError("Embedding response did not include any vectors.")
first = embeddings[0]
values = getattr(first, "values", None)
if not values:
raise RuntimeError("Embedding entry had no values.")
return list(values)