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