139 lines
5.1 KiB
Python
139 lines
5.1 KiB
Python
# -*- 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)
|