203 lines
7.1 KiB
Python
203 lines
7.1 KiB
Python
"""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()
|