# -*- coding: utf-8 -*- """ 向量检索层:Embedding(16011) + Chroma(16010) 纯 REST 实现,无第三方客户端依赖。 - embedding:OpenAI 兼容 /v1/embeddings(bge-large-zh-v1.5,1024维) - rerank :Cohere 兼容 /v1/rerank(bge-reranker-v2-m3,可选增强) - chroma :REST /api/v2(add/query 用 UUID,get/delete 用名字) 扩展其他球类/联赛时:向量索引按 collection 隔离(nba_fan_knowledge_v1 / cba_... ),互不影响。 """ import json import logging import threading import urllib.request import urllib.error import requests from config import (CHROMA_HOST, CHROMA_PORT, CHROMA_COLLECTION, EMBEDDING_API_URL, EMBEDDING_MODEL, USE_RERANK, RERANK_API_URL, RERANK_MODEL) log = logging.getLogger("vector") _BASE = f"http://{CHROMA_HOST}:{CHROMA_PORT}/api/v2/tenants/default_tenant/databases/default_database/collections" _lock = threading.Lock() # ------------------------------------------------------------------ Embedding def embed_texts(texts): """返回 [[float...], ...] 向量列表""" if isinstance(texts, str): texts = [texts] resp = requests.post(EMBEDDING_API_URL, json={"model": EMBEDDING_MODEL, "input": list(texts)}, timeout=120) resp.raise_for_status() data = resp.json().get("data", []) return [d["embedding"] for d in sorted(data, key=lambda x: x["index"])] def rerank(query, docs, top_k=5): """docs: [{"id":..,"text":..}] → 按分数降序返回 top_k(带 score)""" if not USE_RERANK or not docs: return docs try: payload = {"query": query, "documents": docs, "model": RERANK_MODEL, "top_k": top_k} resp = requests.post(RERANK_API_URL, json=payload, timeout=60) resp.raise_for_status() data = resp.json().get("data", []) return [{"id": d["document"]["id"], "text": d["document"]["text"], "score": d["score"]} for d in data] except Exception as e: # rerank 失败不阻断主流程 log.warning("rerank failed: %s", e) return docs # ------------------------------------------------------------------ Chroma def _http(method, url, payload=None, timeout=30): req = urllib.request.Request(url, method=method) if payload is not None: req.add_header("Content-Type", "application/json") req.data = json.dumps(payload).encode("utf-8") try: with urllib.request.urlopen(req, timeout=timeout) as r: body = r.read().decode("utf-8") return r.status, (json.loads(body) if body else {}) except urllib.error.HTTPError as e: body = e.read().decode("utf-8", "ignore") raise RuntimeError(f"Chroma {method} {url} -> {e.code}: {body[:300]}") def _get_collection_id(name): """按名字查集合,返回 (id, 是否存在)""" _, data = _http("GET", f"{_BASE}/{name}") return data.get("id"), True def ensure_collection(name=CHROMA_COLLECTION, space="cosine"): """获取或创建集合,返回 collection_id""" with _lock: try: cid, _ = _get_collection_id(name) return cid except RuntimeError: pass _, data = _http("POST", _BASE, { "name": name, "configuration": {"hnsw": {"space": space}}, "get_or_create": True, }) return data["id"] def collection_count(name=CHROMA_COLLECTION): try: cid, _ = _get_collection_id(name) _, data = _http("GET", f"{_BASE}/{cid}/count") return int(data) if isinstance(data, int) else int(data.get("count", 0)) except Exception: return 0 def add_documents(ids, documents, metadatas, name=CHROMA_COLLECTION): """按文档批量写入(内部自动 embedding)""" if not ids: return cid = ensure_collection(name) vectors = embed_texts(documents) # 批量计算向量 payload = {"ids": list(ids), "embeddings": vectors, "documents": list(documents), "metadatas": list(metadatas)} _http("POST", f"{_BASE}/{cid}/add", payload, timeout=120) def query_vectors(query_text, n_results=5, where=None, name=CHROMA_COLLECTION): """语义检索:返回 [{id, document, distance, metadata}, ...](升序按相似度)""" try: cid, _ = _get_collection_id(name) except RuntimeError: return [] vec = embed_texts(query_text)[0] payload = {"query_embeddings": [vec], "n_results": n_results, "include": ["documents", "metadatas", "distances"]} if where: payload["where"] = where _, data = _http("POST", f"{_BASE}/{cid}/query", payload) out = [] for i, doc in enumerate(data.get("documents", [[]])[0]): out.append({ "id": data["ids"][0][i], "document": doc, "distance": data["distances"][0][i], "metadata": data["metadatas"][0][i] if data.get("metadatas") else {}, }) return out def reset_collection(name=CHROMA_COLLECTION): """重建集合(清空全部数据),用于重新灌库""" try: _http("DELETE", f"{_BASE}/{name}") except RuntimeError: pass return ensure_collection(name)