Files
stock-advisor/rag/vector_store.py
T

144 lines
4.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""
向量检索层:Embedding(16011) + Chroma(16010) 纯 REST 实现,无第三方客户端依赖。
- embeddingOpenAI 兼容 /v1/embeddingsbge-large-zh-v1.51024维)
- chroma REST /api/v2add/query 用 UUIDget/delete 用名字)
两个集合:
- stock_news_v1 财经新闻正文语义索引(metadata: code/title/date/category/sentiment/news_id
- stock_profiles_v1 公司概况语义索引(metadata: code/name/industry
"""
import json
import logging
import threading
import urllib.request
import urllib.error
import requests
from config import (CHROMA_HOST, CHROMA_PORT, EMBEDDING_API_URL,
EMBEDDING_MODEL, RERANK_API_URL, RERANK_MODEL, USE_RERANK)
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):
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:
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, 是否存在)"""
try:
_, data = _http("GET", f"{_BASE}/{name}")
return data.get("id"), True
except RuntimeError:
return None, False
def ensure_collection(name, space="cosine"):
"""获取或创建集合,返回 collection_id"""
with _lock:
cid, exists = _get_collection_id(name)
if exists:
return cid
_, data = _http("POST", _BASE, {
"name": name, "configuration": {"hnsw": {"space": space}},
"get_or_create": True,
})
return data["id"]
def collection_count(name):
try:
cid, exists = _get_collection_id(name)
if not exists:
return 0
_, 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):
"""按文档批量写入(内部自动 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=180)
def query_vectors(query_text, n_results=5, where=None, name=None):
"""语义检索:返回 [{id, document, distance, metadata}, ...](升序按相似度)"""
try:
cid, exists = _get_collection_id(name)
if not exists:
return []
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 delete_collection(name):
"""删除集合(用于重灌)"""
try:
_http("DELETE", f"{_BASE}/{name}")
except RuntimeError:
pass