aps-agent/server/knowledge/embedding.py

194 lines
6.7 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# ============================================================
# 混合嵌入(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