deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/knowledge/embeddings.py
2026-09-07 18:24:55 +08:00

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 []