智能荐股系统 v1.0.0:股票池/行情/新闻RAG/机构持仓/多因子评分/AI深度研报

This commit is contained in:
2026-08-19 19:36:32 +08:00
commit 90a9f9b212
29 changed files with 3332 additions and 0 deletions
+143
View File
@@ -0,0 +1,143 @@
# -*- 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