138 lines
5.1 KiB
Python
138 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)
|