"""Embedding generation — local SentenceTransformers or remote TEI/HF HTTP API. Default model: Alibaba-NLP/gte-multilingual-base (768 dims). When EMBEDDING_URL is set, vectors are fetched over HTTPS (no local torch), which is the preferred path for a no-laptop cloud pilot. """ from __future__ import annotations import logging from functools import lru_cache import httpx from .config import settings from .weave_trace import redact_inputs, summarize_embedding, traced logger = logging.getLogger(__name__) def _embedding_auth_headers() -> dict[str, str]: token = ( settings.embedding_api_key or settings.hf_token or settings.llm_api_key or settings.wandb_api_key or "" ).strip() if not token: return {} return {"Authorization": f"Bearer {token}"} def _normalize(vec: list[float]) -> list[float]: import math norm = math.sqrt(sum(x * x for x in vec)) or 1.0 return [x / norm for x in vec] def _embed_remote(texts: list[str]) -> list[list[float]]: """Call TEI / OpenAI-compatible embeddings, or HF feature-extraction style.""" base = settings.embedding_url.rstrip("/") headers = {"Content-Type": "application/json", **_embedding_auth_headers()} # Prefer OpenAI-compatible /v1/embeddings (TEI and many HF endpoints). url = base if base.endswith("/embeddings") else f"{base}/v1/embeddings" payload = {"model": settings.embedding_model, "input": texts} with httpx.Client(timeout=120.0) as client: resp = client.post(url, json=payload, headers=headers) if resp.status_code == 404 and not base.endswith("/embeddings"): # Fallback: TEI native /embed resp = client.post( f"{base}/embed", json={"inputs": texts if len(texts) > 1 else texts[0]}, headers=headers, ) resp.raise_for_status() data = resp.json() if isinstance(data, list) and data and isinstance(data[0], (int, float)): return [_normalize([float(x) for x in data])] if isinstance(data, list) and data and isinstance(data[0], list): return [_normalize([float(x) for x in row]) for row in data] if isinstance(data, dict) and "data" in data: rows = sorted(data["data"], key=lambda r: r.get("index", 0)) return [_normalize([float(x) for x in r["embedding"]]) for r in rows] raise RuntimeError(f"Unrecognized embedding response shape from {url}") @lru_cache(maxsize=1) def _model(): from sentence_transformers import SentenceTransformer logger.info("Loading embedding model %s", settings.embedding_model) # gte-multilingual-base requires trust_remote_code for its custom pooling. return SentenceTransformer(settings.embedding_model, trust_remote_code=True) @traced( "embed_text", kind="tool", postprocess_inputs=redact_inputs, postprocess_output=summarize_embedding, ) def embed_text(text: str) -> list[float]: """Return a single embedding for the given text.""" if settings.embedding_url: return _embed_remote([text])[0] vec = _model().encode( [text], normalize_embeddings=True, convert_to_numpy=True )[0] return vec.tolist() @traced( "embed_batch", kind="tool", postprocess_inputs=redact_inputs, postprocess_output=summarize_embedding, ) def embed_batch(texts: list[str]) -> list[list[float]]: """Batch embed multiple texts (used during ingestion).""" if not texts: return [] if settings.embedding_url: out: list[list[float]] = [] # Keep batches modest for remote rate limits. chunk = 32 for i in range(0, len(texts), chunk): out.extend(_embed_remote(texts[i : i + chunk])) return out vecs = _model().encode( texts, normalize_embeddings=True, convert_to_numpy=True, batch_size=32 ) return [v.tolist() for v in vecs]