194 lines
6.7 KiB
Python
194 lines
6.7 KiB
Python
# ============================================================
|
||
# 混合嵌入(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
|