80 lines
3.1 KiB
Python
80 lines
3.1 KiB
Python
"""Embedding generation for vector RAG (phase 2).
|
|
|
|
Talks to an OpenAI-compatible ``/embeddings`` endpoint configured under
|
|
``knowledge.embedding``. Everything here is **best-effort**: if embeddings are
|
|
disabled, unconfigured, or the endpoint errors, callers get an empty result and
|
|
the search layer falls back to keyword scoring — so vector RAG never breaks a
|
|
deployment that has no embedding model.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
import math
|
|
from typing import TYPE_CHECKING
|
|
|
|
if TYPE_CHECKING:
|
|
from deerflow.config.knowledge_config import KnowledgeEmbeddingConfig
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def cosine_similarity(a: list[float], b: list[float]) -> float:
|
|
"""Cosine similarity of two equal-length vectors (0 when degenerate)."""
|
|
if not a or not b or len(a) != len(b):
|
|
return 0.0
|
|
dot = sum(x * y for x, y in zip(a, b, strict=False))
|
|
na = math.sqrt(sum(x * x for x in a))
|
|
nb = math.sqrt(sum(y * y for y in b))
|
|
if na == 0 or nb == 0:
|
|
return 0.0
|
|
return dot / (na * nb)
|
|
|
|
|
|
class EmbeddingClient:
|
|
"""Thin async client over an OpenAI-compatible embeddings endpoint."""
|
|
|
|
def __init__(self, config: KnowledgeEmbeddingConfig) -> None:
|
|
self.config = config
|
|
|
|
def is_enabled(self) -> bool:
|
|
return bool(self.config.enabled and self.config.base_url and self.config.model)
|
|
|
|
async def embed_texts(self, texts: list[str]) -> list[list[float]]:
|
|
"""Return one vector per input text; ``[]`` on any failure/disabled."""
|
|
if not self.is_enabled() or not texts:
|
|
return []
|
|
try:
|
|
import httpx
|
|
except Exception: # pragma: no cover - httpx is a core dep
|
|
logger.warning("httpx unavailable; knowledge embeddings disabled")
|
|
return []
|
|
|
|
url = self.config.base_url.rstrip("/") + "/embeddings"
|
|
headers = {"Content-Type": "application/json"}
|
|
if self.config.api_key:
|
|
headers["Authorization"] = f"Bearer {self.config.api_key}"
|
|
|
|
vectors: list[list[float]] = []
|
|
batch = max(1, self.config.batch_size)
|
|
try:
|
|
async with httpx.AsyncClient(timeout=self.config.timeout_seconds) as client:
|
|
for start in range(0, len(texts), batch):
|
|
chunk = texts[start : start + batch]
|
|
payload: dict = {"model": self.config.model, "input": chunk}
|
|
if self.config.dimensions:
|
|
payload["dimensions"] = self.config.dimensions
|
|
resp = await client.post(url, headers=headers, json=payload)
|
|
resp.raise_for_status()
|
|
data = resp.json()
|
|
items = sorted(data.get("data", []), key=lambda d: d.get("index", 0))
|
|
vectors.extend([item.get("embedding", []) for item in items])
|
|
except Exception:
|
|
logger.exception("Knowledge embedding request failed; falling back to keyword search")
|
|
return []
|
|
return vectors
|
|
|
|
async def embed_query(self, text: str) -> list[float]:
|
|
vecs = await self.embed_texts([text])
|
|
return vecs[0] if vecs else []
|