Files

284 lines
11 KiB
Python
Raw Permalink 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 -*-
"""
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):
"""调用 DeepSeekOpenAI 兼容)。返回最终 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,
}