284 lines
11 KiB
Python
284 lines
11 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""
|
||
AI 分析引擎:DeepSeek 深度研报 + RAG 增强
|
||
- 检索:股票相关新闻(向量语义)+ 公司概况 + 机构动向/基金持仓(结构化)
|
||
- 生成:结构化工研报(公司概况/基本面/技术面/消息面/机构动向/风险提示/操作建议)
|
||
"""
|
||
import logging
|
||
import threading
|
||
import time
|
||
import requests
|
||
|
||
from config import (LLM_API_KEY, LLM_BASE_URL, LLM_MAX_TOKENS, LLM_MODEL,
|
||
LLM_TEMPERATURE, LLM_TIMEOUT, CHROMA_NEWS_COLLECTION,
|
||
CHROMA_PROFILE_COLLECTION)
|
||
from database import query, query_one, execute
|
||
from rag.vector_store import query_vectors
|
||
|
||
log = logging.getLogger("analyst")
|
||
|
||
_jobs = {} # code -> {status, report, error, ts}
|
||
_jobs_lock = threading.Lock()
|
||
|
||
|
||
# ------------------------------------------------------------------ LLM
|
||
def llm_chat(messages, max_tokens=None, temperature=None, timeout=None):
|
||
"""调用 DeepSeek(OpenAI 兼容)。返回最终 content(忽略推理过程)
|
||
模型/Key/地址从设置区读取(settings 覆盖 config 默认)"""
|
||
from settings import llm_config
|
||
cfg = llm_config()
|
||
resp = requests.post(
|
||
f"{cfg['base_url']}/chat/completions",
|
||
headers={"Authorization": f"Bearer {cfg['api_key']}"},
|
||
json={
|
||
"model": cfg["model"],
|
||
"messages": messages,
|
||
"max_tokens": max_tokens or cfg["max_tokens"],
|
||
"temperature": cfg["temperature"] if temperature is None else temperature,
|
||
"stream": False,
|
||
},
|
||
timeout=timeout or cfg["timeout"],
|
||
)
|
||
resp.raise_for_status()
|
||
data = resp.json()
|
||
try:
|
||
return data["choices"][0]["message"].get("content") or ""
|
||
except (KeyError, IndexError):
|
||
return ""
|
||
|
||
|
||
# ------------------------------------------------------------------ RAG 检索
|
||
def _rag_news(code, stock_name, query_text, top_k=6):
|
||
"""检索个股相关新闻(向量语义,按 code 过滤)"""
|
||
try:
|
||
where = {"code": code}
|
||
hits = query_vectors(query_text, n_results=top_k, where=where,
|
||
name=CHROMA_NEWS_COLLECTION)
|
||
out = []
|
||
for h in hits:
|
||
m = h.get("metadata", {})
|
||
out.append({
|
||
"title": m.get("title", ""),
|
||
"date": m.get("date", ""),
|
||
"sentiment": m.get("sentiment", 0),
|
||
"text": h.get("document", "")[:400],
|
||
"distance": round(h.get("distance", 0), 3),
|
||
})
|
||
return out
|
||
except Exception as e:
|
||
log.warning("RAG news fail: %s", e)
|
||
return []
|
||
|
||
|
||
def _rag_profile(code):
|
||
try:
|
||
hits = query_vectors("公司主营业务与基本面", n_results=1,
|
||
where={"code": code}, name=CHROMA_PROFILE_COLLECTION)
|
||
if hits:
|
||
return hits[0].get("document", "")
|
||
except Exception:
|
||
pass
|
||
return ""
|
||
|
||
|
||
def _inst_summary(code):
|
||
"""机构评级 + 基金持仓摘要(结构化)"""
|
||
ratings = query(
|
||
"SELECT inst_name, rating, target_price, rating_date, prev_rating "
|
||
"FROM inst_ratings WHERE stock_code=? ORDER BY rating_date DESC LIMIT 6", (code,))
|
||
holdings = query(
|
||
"SELECT inst_name, quarter, hold_value, change_pct FROM fund_holdings "
|
||
"WHERE stock_code=? ORDER BY quarter DESC, hold_value DESC LIMIT 6", (code,))
|
||
return ratings, holdings
|
||
|
||
|
||
# ------------------------------------------------------------------ 报告生成
|
||
def _fmt_indicators(ind):
|
||
if not ind:
|
||
return "(暂无技术数据)"
|
||
lines = [
|
||
f"- 最新价 {ind.get('close')},当日 {ind.get('change_pct', 0):+.2f}%",
|
||
f"- MA5={ind.get('ma5')} / MA10={ind.get('ma10')} / MA20={ind.get('ma20')} / MA60={ind.get('ma60')}",
|
||
f"- RSI(14)={ind.get('rsi')},KDJ K/D/J={ind.get('kdj_k')}/{ind.get('kdj_d')}/{ind.get('kdj_j')}",
|
||
f"- MACD DIF={ind.get('dif')} / DEA={ind.get('dea')} / 柱={ind.get('macd')}",
|
||
f"- 量比 {ind.get('vol_ratio')},5日涨幅 {ind.get('chg_5d', 0):+.2f}%,20日涨幅 {ind.get('chg_20d', 0):+.2f}%",
|
||
f"- 近120日区间 {ind.get('low_52w')} ~ {ind.get('high_52w')},20日波动率 {ind.get('volatility')}%",
|
||
]
|
||
return "\n".join(lines)
|
||
|
||
|
||
def _build_prompt(stock, ind, news_hits, profile, ratings, holdings, score, focus):
|
||
rated = "、".join(f"{r['inst_name']}({r['rating']},目标{r['target_price']})" for r in ratings) or "暂无"
|
||
held = ";".join(f"{h['inst_name']} {h['quarter']}持仓{h['hold_value']:.0f}万 环比{h['change_pct']:+.1f}%" for h in holdings) or "暂无"
|
||
news_text = "\n\n".join(
|
||
f"【{n['date']}|{n['title']}】(情感{n['sentiment']:+.2f})\n{n['text']}" for n in news_hits
|
||
) or "(检索到相关资讯较少)"
|
||
|
||
return f"""你是资深A股投顾,请基于下方【资料】对股票 {stock['name']}({stock['code']}) 输出一份结构化工研报。
|
||
|
||
【资料】
|
||
公司概况:
|
||
{profile or stock.get('description', '暂无')}
|
||
|
||
技术面:
|
||
{_fmt_indicators(ind)}
|
||
|
||
综合评分:{score.get('total')} 分(评级:{score.get('rating')}),分项:趋势{score.get('trend')}/动量{score.get('momentum')}/技术{score.get('technical')}/量能{score.get('volume')}/消息{score.get('news')}/机构{score.get('institutional')}
|
||
|
||
机构评级:{rated}
|
||
基金持仓:{held}
|
||
|
||
相关资讯(RAG 语义检索):
|
||
{news_text}
|
||
|
||
用户关注点:{focus or '整体投资价值'}
|
||
|
||
【输出要求】用 Markdown 输出,结构如下:
|
||
## 一、公司概况与基本面
|
||
## 二、技术面解读
|
||
## 三、消息面与市场情绪
|
||
## 四、机构动向
|
||
## 五、风险提示
|
||
## 六、操作建议(给出 目标区间 / 支撑位 / 压力位,说明短线与中线思路)
|
||
注意:内容需严格基于上述资料,数据为模拟数据,结尾加一句「以上内容基于模拟数据生成,仅供系统演示,不构成投资建议」。"""
|
||
|
||
|
||
def generate_report_sync(code, focus=""):
|
||
"""同步生成报告(后台线程调用),并记录历史 + 数据源"""
|
||
import json
|
||
stock = query_one("SELECT * FROM stocks WHERE code=?", (code,))
|
||
if not stock:
|
||
return {"error": "股票不存在"}
|
||
ind = _indicators_for(code)
|
||
score = _score_for(code, ind)
|
||
hits = _rag_news(code, stock["name"], f"{stock['name']} {focus or '投资价值 业绩 利好利空'} {ind.get('close','')}")
|
||
profile = _rag_profile(code)
|
||
ratings, holdings = _inst_summary(code)
|
||
prompt = _build_prompt(stock, ind, hits, profile, ratings, holdings, score, focus)
|
||
# 记录大模型参考的数据源(供详情页展示)
|
||
sources = {
|
||
"focus": focus,
|
||
"score": score,
|
||
"indicators": _fmt_indicators(ind),
|
||
"profile": profile or stock.get("description", ""),
|
||
"news": hits,
|
||
"ratings": ratings,
|
||
"holdings": holdings,
|
||
"prompt": prompt,
|
||
}
|
||
try:
|
||
report = llm_chat([
|
||
{"role": "system", "content": "你是一名严谨专业的A股投资顾问,输出结构化、简洁、可执行的研报。"},
|
||
{"role": "user", "content": prompt},
|
||
])
|
||
report = report.strip()
|
||
if not report:
|
||
raise RuntimeError("LLM 返回为空")
|
||
execute("INSERT OR REPLACE INTO analysis_cache(code, report, created_at) VALUES(?,?,datetime('now','localtime'))",
|
||
(code, report))
|
||
execute(
|
||
"INSERT INTO analysis_history(code, stock_name, focus, report, sources, created_at) "
|
||
"VALUES(?,?,?,?,?,datetime('now','localtime'))",
|
||
(code, stock["name"], focus, report, json.dumps(sources, ensure_ascii=False)))
|
||
return {"report": report, "ts": time.time()}
|
||
except Exception as e:
|
||
log.exception("gen report fail")
|
||
return {"error": str(e)}
|
||
|
||
|
||
def _indicators_for(code):
|
||
"""从 DB 读取日线并算指标(避免循环依赖 app)"""
|
||
from engine.indicators import compute_indicators
|
||
rows = query("SELECT date,open,high,low,close,volume FROM stock_daily WHERE code=? ORDER BY date ASC", (code,))
|
||
return compute_indicators(rows)
|
||
|
||
|
||
def _score_for(code, ind):
|
||
from engine.scoring import score_stock
|
||
# 新闻情感(与 app 端口径一致:精确/前缀/后缀三种匹配)
|
||
n = query_one(
|
||
"SELECT AVG(sentiment) AS s FROM news WHERE (related_stocks=? OR related_stocks LIKE ? OR related_stocks LIKE ?) "
|
||
"AND publish_date >= date('now','-7 day')",
|
||
(code, f"%,{code}", f"{code},%"))
|
||
news_score = n["s"] if n and n["s"] is not None else 0.0
|
||
# 机构热度
|
||
st = query_one(
|
||
"SELECT COUNT(*) AS c FROM inst_ratings WHERE stock_code=? AND rating IN ('买入','增持') "
|
||
"AND rating_date >= date('now','-30 day')", (code,))
|
||
inst_count = st["c"] if st else 0
|
||
inst_score = min(1.0, inst_count / 4.0)
|
||
return score_stock(ind, news_score, inst_score)
|
||
|
||
|
||
# ------------------------------------------------------------------ 异步任务
|
||
def submit_report(code, focus=""):
|
||
"""提交后台生成任务,立即返回"""
|
||
with _jobs_lock:
|
||
if _jobs.get(code, {}).get("status") == "running":
|
||
return {"status": "running"}
|
||
_jobs[code] = {"status": "running", "report": None, "error": None, "ts": time.time()}
|
||
threading.Thread(target=_run_job, args=(code, focus), daemon=True).start()
|
||
return {"status": "running"}
|
||
|
||
|
||
def _run_job(code, focus):
|
||
try:
|
||
res = generate_report_sync(code, focus)
|
||
with _jobs_lock:
|
||
if res.get("error"):
|
||
_jobs[code] = {"status": "error", "error": res["error"], "ts": time.time()}
|
||
else:
|
||
_jobs[code] = {"status": "done", "report": res["report"], "ts": time.time()}
|
||
except Exception as e:
|
||
with _jobs_lock:
|
||
_jobs[code] = {"status": "error", "error": str(e), "ts": time.time()}
|
||
|
||
|
||
def report_status(code):
|
||
with _jobs_lock:
|
||
return dict(_jobs.get(code, {}))
|
||
|
||
|
||
def get_cached_report(code):
|
||
return query_one("SELECT report, created_at FROM analysis_cache WHERE code=?", (code,))
|
||
|
||
|
||
# ------------------------------------------------------------------ 历史记录
|
||
def list_history(code, limit=20):
|
||
"""某股票的历史 AI 分析记录(不含正文,只返回摘要)"""
|
||
rows = query(
|
||
"SELECT id, code, stock_name, focus, created_at, sources, LENGTH(report) AS len, "
|
||
"SUBSTR(report, 1, 60) AS excerpt FROM analysis_history "
|
||
"WHERE code=? ORDER BY id DESC LIMIT ?", (code, limit))
|
||
out = []
|
||
for r in rows:
|
||
try:
|
||
import json
|
||
src = json.loads(r.get("sources") or "{}")
|
||
except Exception:
|
||
src = {}
|
||
out.append({
|
||
"id": r["id"], "code": r["code"], "stock_name": r["stock_name"],
|
||
"focus": r["focus"], "created_at": r["created_at"],
|
||
"chars": r["len"], "excerpt": (r["excerpt"] or "").strip(),
|
||
"news_count": len(src.get("news") or []),
|
||
})
|
||
return out
|
||
|
||
|
||
def get_history(aid):
|
||
"""单条分析详情:正文 + 数据源 JSON"""
|
||
import json
|
||
row = query_one("SELECT * FROM analysis_history WHERE id=?", (aid,))
|
||
if not row:
|
||
return None
|
||
try:
|
||
src = json.loads(row.get("sources") or "{}")
|
||
except Exception:
|
||
src = {}
|
||
return {
|
||
"id": row["id"], "code": row["code"], "stock_name": row["stock_name"],
|
||
"focus": row["focus"], "report": row["report"], "created_at": row["created_at"],
|
||
"sources": src,
|
||
}
|