"""Embedding + reranking client. - Embeddings: BAAI/bge-m3 via vLLM OpenAI-compat endpoint, 1024-dim, 8K context. - Reranker: BAAI/bge-reranker-v2-m3 cross-encoder via vllm-rerank-api endpoint. Both run on the same host (10.11.10.15) on different ports (8200 / 8100). """ from __future__ import annotations import math from collections.abc import AsyncIterator, Sequence from contextlib import asynccontextmanager from dataclasses import dataclass from typing import Any import httpx from tenacity import ( AsyncRetrying, retry_if_exception_type, stop_after_attempt, wait_exponential, ) from shared.config import settings from shared.logging import get_logger log = get_logger(__name__) class EmbeddingError(Exception): pass @dataclass(slots=True, frozen=True) class RerankResult: """One reranked document, sorted by score descending.""" index: int # original position in the input documents list score: float # cross-encoder score (higher = more relevant) document: str # the text of that document class EmbeddingClient: """Async client for BGE-M3 embeddings + BGE-reranker-v2-m3 reranking.""" def __init__( self, *, embed_url: str | None = None, rerank_url: str | None = None, timeout: float = 60.0, max_attempts: int = 3, ): self._embed_base = (embed_url or settings.embedding_url).rstrip("/") self._rerank_base = (rerank_url or settings.reranker_url).rstrip("/") self._max_attempts = max_attempts embed_headers = {"Content-Type": "application/json"} if settings.embedding_api_key: embed_headers["Authorization"] = f"Bearer {settings.embedding_api_key}" rerank_headers = {"Content-Type": "application/json"} if settings.reranker_api_key: rerank_headers["Authorization"] = f"Bearer {settings.reranker_api_key}" self._embed_http = httpx.AsyncClient( base_url=self._embed_base, headers=embed_headers, timeout=httpx.Timeout(timeout, connect=10.0), ) self._rerank_http = httpx.AsyncClient( base_url=self._rerank_base, headers=rerank_headers, timeout=httpx.Timeout(timeout, connect=10.0), ) async def aclose(self) -> None: await self._embed_http.aclose() await self._rerank_http.aclose() async def __aenter__(self) -> EmbeddingClient: return self async def __aexit__(self, *args: Any) -> None: await self.aclose() # ------------------------------------------------------------------ embed async def embed(self, texts: Sequence[str]) -> list[list[float]]: """Embed a batch of texts. Returns 1024-dim vectors in input order. BGE-M3 max input is 8192 tokens; we don't truncate here — caller is responsible for chunking long inputs (Atomic does this server-side). Empty input returns empty list. """ if not texts: return [] payload = {"model": settings.embedding_model, "input": list(texts)} async for attempt in AsyncRetrying( stop=stop_after_attempt(self._max_attempts), wait=wait_exponential(multiplier=1, min=1, max=10), retry=retry_if_exception_type((httpx.HTTPError, httpx.TimeoutException)), reraise=True, ): with attempt: resp = await self._embed_http.post("/v1/embeddings", json=payload) if resp.status_code >= 400: raise EmbeddingError( f"HTTP {resp.status_code} from embedding endpoint: " f"{resp.text[:500]}" ) data = resp.json() items = data.get("data") if not items or len(items) != len(texts): raise EmbeddingError( f"Embedding count mismatch: got {len(items or [])} for " f"{len(texts)} inputs" ) vecs = [item["embedding"] for item in items] if vecs and len(vecs[0]) != settings.embedding_dim: raise EmbeddingError( f"Embedding dim mismatch: got {len(vecs[0])}, " f"expected {settings.embedding_dim}" ) return vecs raise EmbeddingError("retry loop exited without result") # unreachable async def embed_one(self, text: str) -> list[float]: """Convenience for single-text embedding.""" result = await self.embed([text]) return result[0] # ----------------------------------------------------------------- rerank async def rerank( self, query: str, documents: Sequence[str], *, top_n: int | None = None, ) -> list[RerankResult]: """Cross-encoder rerank query × documents. Returns sorted by score desc. Use this AFTER first-stage embedding retrieval to dramatically boost precision on the top-N candidates (typical pattern: embed retrieves 50, rerank picks 10). """ if not documents: return [] payload: dict[str, Any] = { "model": settings.reranker_model, "query": query, "documents": list(documents), } if top_n is not None: payload["top_n"] = top_n async for attempt in AsyncRetrying( stop=stop_after_attempt(self._max_attempts), wait=wait_exponential(multiplier=1, min=1, max=10), retry=retry_if_exception_type((httpx.HTTPError, httpx.TimeoutException)), reraise=True, ): with attempt: resp = await self._rerank_http.post("/v1/rerank", json=payload) if resp.status_code >= 400: raise EmbeddingError( f"HTTP {resp.status_code} from reranker: {resp.text[:500]}" ) data = resp.json() results = data.get("results", []) return [ RerankResult( index=int(r["index"]), score=float(r["score"]), document=r.get("document") or documents[int(r["index"])], ) for r in results ] raise EmbeddingError("retry loop exited without result") # unreachable # --------------------------------------------------------------------- math util def cosine(a: Sequence[float], b: Sequence[float]) -> float: """Cosine similarity. BGE-M3 vectors are NOT pre-normalized; we compute it.""" if len(a) != len(b): raise ValueError(f"vector dim mismatch: {len(a)} vs {len(b)}") dot = sum(x * y for x, y in zip(a, b, strict=True)) na = math.sqrt(sum(x * x for x in a)) nb = math.sqrt(sum(x * x for x in b)) if na == 0 or nb == 0: return 0.0 return dot / (na * nb) @asynccontextmanager async def embedding_client() -> AsyncIterator[EmbeddingClient]: client = EmbeddingClient() try: yield client finally: await client.aclose()