aps-agent/server/knowledge/embedding.py

194 lines
6.7 KiB
Python
Raw Permalink Normal View History

# ============================================================
# 混合嵌入(moduleId: knowledge-embedding, 可重生 ✅)
# 本地 sentence-transformers → API embeddings → None(检索降级 bigram)
# ============================================================
from __future__ import annotations
import json
import math
import os
import tempfile
import threading
from typing import Any
import httpx
_PATH_ENV = "APS_EMBEDDINGS_PATH"
def _cosine(a: list[float], b: list[float]) -> float:
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))
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 EmbeddingProvider:
"""混合嵌入:local → api → disabled(检索走 bigram)。"""
def __init__(self) -> None:
self.mode = (os.environ.get("EMBEDDING_PROVIDER") or "auto").strip().lower()
self.api_key = os.environ.get("EMBEDDING_API_KEY") or os.environ.get("LLM_API_KEY") or ""
self.base_url = (os.environ.get("EMBEDDING_BASE_URL")
or os.environ.get("LLM_BASE_URL") or "").strip()
self.model = (os.environ.get("EMBEDDING_MODEL") or "text-embedding-3-small").strip()
self.local_model = os.environ.get("EMBEDDING_LOCAL_MODEL") or "BAAI/bge-small-zh-v1.5"
self._local = None
self._backend = self._resolve_backend()
def _resolve_backend(self) -> str:
want = self.mode
if want in ("off", "none", "bigram"):
return "none"
if want in ("local", "auto"):
try:
from sentence_transformers import SentenceTransformer # noqa: F401
return "local"
except Exception:
if want == "local":
return "none"
if want in ("api", "auto") and self.api_key and self.base_url:
return "api"
return "none"
@property
def enabled(self) -> bool:
return self._backend in ("local", "api")
@property
def backend(self) -> str:
return self._backend
def _ensure_local(self):
if self._local is None:
from sentence_transformers import SentenceTransformer
self._local = SentenceTransformer(self.local_model)
return self._local
def embed(self, texts: list[str]) -> list[list[float]] | None:
"""批量嵌入;不可用返回 None。"""
if not texts:
return []
if self._backend == "local":
model = self._ensure_local()
vecs = model.encode(texts, normalize_embeddings=True)
return [v.tolist() for v in vecs]
if self._backend == "api":
try:
with httpx.Client(timeout=30.0) as client:
resp = client.post(
f"{self.base_url.rstrip('/')}/embeddings",
headers={"Authorization": f"Bearer {self.api_key}"},
json={"model": self.model, "input": texts},
)
resp.raise_for_status()
data = resp.json().get("data") or []
data = sorted(data, key=lambda x: x.get("index", 0))
return [d["embedding"] for d in data]
except Exception:
return None
return None
def embed_one(self, text: str) -> list[float] | None:
out = self.embed([text])
return out[0] if out else None
class EmbeddingStore:
"""向量仓:chunk/asset id → vector,JSON 落盘。"""
def __init__(self, path: str | None = None) -> None:
self.path = path or os.environ.get(_PATH_ENV, "server/data/embeddings.json")
self._lock = threading.Lock()
self.vectors: dict[str, list[float]] = self._load()
def _load(self) -> dict[str, list[float]]:
try:
with open(self.path, "r", encoding="utf-8") as f:
return json.load(f).get("vectors", {})
except (FileNotFoundError, json.JSONDecodeError):
return {}
def _write(self) -> None:
os.makedirs(os.path.dirname(self.path) or ".", exist_ok=True)
fd, tmp = tempfile.mkstemp(dir=os.path.dirname(self.path) or ".", suffix=".tmp")
try:
with os.fdopen(fd, "w", encoding="utf-8") as f:
json.dump({"vectors": self.vectors}, f, ensure_ascii=False)
os.replace(tmp, self.path)
except BaseException:
if os.path.exists(tmp):
os.unlink(tmp)
raise
def upsert_many(self, items: dict[str, list[float]]) -> None:
with self._lock:
self.vectors.update(items)
self._write()
def search(self, query_vec: list[float], top_k: int = 5) -> list[tuple[str, float]]:
scored = [(_id, _cosine(query_vec, vec)) for _id, vec in self.vectors.items()]
scored.sort(key=lambda x: -x[1])
return scored[:top_k]
_provider: EmbeddingProvider | None = None
_stores: dict[str, EmbeddingStore] = {}
_stores_lock = threading.Lock()
def get_embedding_provider() -> EmbeddingProvider:
global _provider
if _provider is None:
_provider = EmbeddingProvider()
return _provider
def get_embedding_store() -> EmbeddingStore:
from server.auth.context import get_identity
from server.state.store import world_path_for
tenant_uuid = get_identity().tenant_uuid
with _stores_lock:
store = _stores.get(tenant_uuid)
if store is None:
if tenant_uuid == "platform":
path = os.environ.get(_PATH_ENV, "server/data/embeddings.json")
else:
world_path = world_path_for("knowledge", tenant_uuid)
path = os.path.join(os.path.dirname(os.path.dirname(world_path)), "embeddings.json")
store = EmbeddingStore(path)
_stores[tenant_uuid] = store
return store
def index_units(units: list[dict[str, Any]]) -> int:
"""为检索单元生成/更新向量;返回成功条数。"""
prov = get_embedding_provider()
if not prov.enabled:
return 0
store = get_embedding_store()
texts, keys = [], []
for u in units:
key = u.get("chunkId") or u["assetId"]
if key in store.vectors:
continue
texts.append((u.get("title") or "") + "\n" + (u.get("content") or ""))
keys.append(f"{u['assetId']}:{key}")
if not texts:
return 0
# 分批
added = 0
batch = 16
for i in range(0, len(texts), batch):
vecs = prov.embed(texts[i:i + batch])
if not vecs:
break
store.upsert_many({keys[i + j]: vecs[j] for j in range(len(vecs))})
added += len(vecs)
return added