Files
stock-advisor/engine/agent.py
T

394 lines
17 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 -*-
"""
持仓跟踪智能体(AI Agent + 工作流)
对持仓/自选股票进行定期深度跟踪,工作流:
1. 上下文构建 读取产业链知识库(上游/下游/同业 + 关键词)
2. 数据采集 个股动态 + 上游产业链 + 下游产业链 + 同业动态(DB + RAG 向量检索)
3. 智能体分析 大模型扮演产业链跟踪分析师,输出深度专业分析 + 影响度判定
4. 沉淀与通知 结果入库 tracking_reports;重大变化(impact_score 达阈值)邮件通知
「智能体」体现在:大模型自主综合个股与产业链多环节信息,输出结构化的
个股动态 / 上游供给成本 / 下游需求景气 / 同业竞争 / 传导影响 / 风险与关注要点,
并给出 0-100 影响度评分与变化性质判定。
"""
import json
import logging
import re
import threading
import time
from database import query, query_one, execute
from settings import tracking_config, tracking_state, set_tracking_state, mail_config
from engine.chain_data import get_chain
from engine.indicators import compute_indicators
from rag.vector_store import query_vectors
from config import CHROMA_NEWS_COLLECTION
from engine.analyst import llm_chat
log = logging.getLogger("agent")
_track_lock = threading.Lock()
_cycle_running = False
_jobs = {} # code -> {status, error, ts}
# ===================================================================== 数据采集
def _db_news(codes, days=40, limit=8):
"""按关联股票代码取新闻(近 days 天)"""
if not codes:
return []
conds, args = [], []
for c in codes:
conds.append("(related_stocks=? OR related_stocks LIKE ? OR related_stocks LIKE ?)")
args += [c, f"%,{c}", f"{c},%"]
args.append(limit)
return query(
f"SELECT id,title,content,source,category,publish_date,sentiment FROM news "
f"WHERE ({' OR '.join(conds)}) AND publish_date >= date('now','-{days} day') "
f"ORDER BY publish_date DESC LIMIT ?", args)
def _rag_news_text(question, where=None, top_k=4):
"""向量语义检索,返回紧凑文本"""
try:
hits = query_vectors(question, n_results=top_k, where=where, name=CHROMA_NEWS_COLLECTION)
out = []
for h in hits:
m = h.get("metadata", {})
out.append(f" [{m.get('date','')}] {m.get('title','')} (情感{m.get('sentiment',0):+.2f}) "
f"{h.get('document','')[:90]}")
return out
except Exception as e:
log.warning("rag fail: %s", e)
return []
def _fmt_news(items):
return "\n".join(
f" [{n['publish_date']}] {n['title']} (情感{n['sentiment']:+.2f}) {n['content'][:90]}"
for n in items) or " (暂无)"
def collect_chain(code):
"""采集个股 + 产业链各环节资讯,返回结构化 dict"""
stock = query_one("SELECT * FROM stocks WHERE code=?", (code,))
if not stock:
return None
chain = get_chain(stock["industry"])
seg = {}
# 1. 个股直接动态
direct = _db_news([code], days=40, limit=8)
direct_rag = _rag_news_text(f"{stock['name']} 最新动态 业绩 公告 重大事项", where={"code": code}, top_k=4)
seg["direct"] = {"db": _fmt_news(direct), "rag": "\n".join(direct_rag) or " (暂无)"}
# 2. 上游
up_db = _db_news(chain["upstream"]["codes"], days=40, limit=6)
up_rag = _rag_news_text(f"{stock['industry']} 上游 {chain['upstream']['keywords']}", top_k=4)
seg["upstream"] = {"codes": chain["upstream"]["codes"],
"keywords": chain["upstream"]["keywords"],
"db": _fmt_news(up_db), "rag": "\n".join(up_rag) or " (暂无)"}
# 3. 下游
dn_db = _db_news(chain["downstream"]["codes"], days=40, limit=6)
dn_rag = _rag_news_text(f"{stock['industry']} 下游需求 景气 {chain['downstream']['keywords']}", top_k=4)
seg["downstream"] = {"codes": chain["downstream"]["codes"],
"keywords": chain["downstream"]["keywords"],
"db": _fmt_news(dn_db), "rag": "\n".join(dn_rag) or " (暂无)"}
# 4. 同业
peer_db = _db_news(chain["peers"], days=40, limit=5)
peer_rag = _rag_news_text(f"{stock['industry']} 竞争格局 同业 {chain['peers']}", top_k=3)
seg["peers"] = {"codes": chain["peers"],
"db": _fmt_news(peer_db), "rag": "\n".join(peer_rag) or " (暂无)"}
# 5. 技术面 + 机构
ind = compute_indicators(query("SELECT date,open,high,low,close,volume FROM stock_daily "
"WHERE code=? ORDER BY date ASC", (code,)))
ratings = query("SELECT inst_name, rating, target_price, rating_date FROM inst_ratings "
"WHERE stock_code=? ORDER BY rating_date DESC LIMIT 5", (code,))
holdings = query("SELECT inst_name, quarter, hold_value, change_pct FROM fund_holdings "
"WHERE stock_code=? ORDER BY quarter DESC LIMIT 5", (code,))
seg["stock"] = stock
seg["chain"] = chain
seg["indicators"] = ind
seg["ratings"] = ratings
seg["holdings"] = holdings
return seg
# ===================================================================== 分析提示词
def _fmt_ind(ind):
return (f"最新价 {ind.get('close')}{ind.get('change_pct',0):+.2f}%),5日{ind.get('chg_5d',0):+.2f}% / "
f"20日{ind.get('chg_20d',0):+.2f}%RSI={ind.get('rsi')},量比{ind.get('vol_ratio')}"
f"MA20={ind.get('ma20')}")
def build_prompt(seg):
s = seg["stock"]
chain = seg["chain"]
rated = "、".join(f"{r['inst_name']}({r['rating']},目标{r['target_price']})" for r in seg["ratings"]) or "暂无"
held = "".join(f"{h['inst_name']} {h['quarter']}持仓{h['hold_value']:.0f}万 环比{h['change_pct']:+.1f}%"
for h in seg["holdings"]) or "暂无"
return f"""你是资深产业链跟踪分析师,正在对【持仓标的】{s['name']}({s['code']}) 进行深度跟踪。请综合【个股】与【产业链上下游/同业】的全部动态,输出一份专业、有洞察的产业链跟踪分析。
【个股基本面】
{s.get('description','')}
【技术面】{_fmt_ind(seg['indicators'])}
【机构动向】评级:{rated} 持仓:{held}
【一、个股直接动态】(公告/新闻/机构观点)
{seg['direct']['db']}
{seg['direct']['rag']}
【二、上游产业链】(供给/原材料/成本端;关联股票 {seg['upstream']['codes'] or '无'},关键词:{seg['upstream']['keywords']}
{seg['upstream']['db']}
{seg['upstream']['rag']}
【三、下游产业链】(需求/客户/景气端;关联股票 {seg['downstream']['codes'] or '无'},关键词:{seg['downstream']['keywords']}
{seg['downstream']['db']}
{seg['downstream']['rag']}
【四、同业动态】(竞争格局;{seg['peers']['codes'] or '无'}
{seg['peers']['db']}
{seg['peers']['rag']}
【输出要求】
第一步,先输出一个 json 代码块(必须最先输出,内容为本次判定,不要包含其他内容):
```json
{{"significance":"high|medium|low","impact_score":0到100的整数,"change_kind":"利好/利空/中性/震荡","summary":"一句话总结","chain_trend":"产业链趋势判断"}}
```
第二步,再输出 Markdown 分析报告,结构如下:
## 一、个股最新动态
## 二、上游产业链分析(供给、原材料、成本端变化及其传导)
## 三、下游产业链分析(需求、客户、景气度变化及其传导)
## 四、同行业竞争格局
## 五、产业链传导与投资启示(上游→中游→下游,对{ s['name']}的影响路径)
## 六、风险提示
## 关注要点(3-5条)
分析须严格基于提供的资讯,避免编造。impact_score 反映本次跟踪发现的动态对股价的潜在影响程度:>=65 视为重大变化。"""
def parse_judge(text):
"""从容错地从 LLM 输出中提取 JSON 判定(支持 json 代码块/截断/缺失字段)"""
t = text.strip()
# 0) 优先取 json 代码块
m = re.search(r"```json\s*(.*?)\s*```", t, re.S)
if m:
try:
j = json.loads(m.group(1))
if "significance" in j or "impact_score" in j:
return j
except Exception:
pass
# 1) 整体 JSON 对象解析
for m in re.finditer(r"\{[^{}]*\}", t, re.S):
try:
j = json.loads(m.group(0))
if "significance" in j or "impact_score" in j:
return j
except Exception:
continue
# 2) 逐字段容错提取(末尾被截断时也能拿到已输出字段)
out = {}
m = re.search(r'"significance"\s*:\s*"(high|medium|low)"', t)
if m:
out["significance"] = m.group(1)
m = re.search(r'"impact_score"\s*:\s*(\d+)', t)
if m:
out["impact_score"] = int(m.group(1))
m = re.search(r'"change_kind"\s*:\s*"([^"]{1,20})"', t)
if m:
out["change_kind"] = m.group(1)
m = re.search(r'"summary"\s*:\s*"((?:[^"\\]|\\.){1,200})"', t)
if m:
out["summary"] = m.group(1)
m = re.search(r'"chain_trend"\s*:\s*"((?:[^"\\]|\\.){1,120})"', t)
if m:
out["chain_trend"] = m.group(1)
return out or None
def strip_json_block(text):
"""从报告文本中剥离最前面的 json 代码块,保留纯 Markdown"""
m = re.search(r"```json\s*.*?```\s*", text, re.S)
if m:
return text[m.end():].strip()
return text
# ===================================================================== 执行
def track_stock(code, focus=""):
"""执行一次跟踪,返回 {ok, report_id, meta, ...}"""
with _track_lock:
seg = collect_chain(code)
if not seg:
return {"error": "股票不存在"}
s = seg["stock"]
prompt = build_prompt(seg)
try:
reply = llm_chat([
{"role": "system", "content": "你是一名严谨专业的产业链跟踪分析师,输出结构化、有数据支撑的分析。"},
{"role": "user", "content": prompt},
]).strip()
if not reply:
raise RuntimeError("LLM 返回为空")
judge = parse_judge(reply)
# 摘要兜底:解析不到则取报告首个标题;并剥离 json 块保留纯净 Markdown
clean_report = strip_json_block(reply)
summary = (judge or {}).get("summary") or _first_heading(clean_report)
meta = {
"significance": (judge or {}).get("significance", "medium"),
"impact_score": int((judge or {}).get("impact_score", 50)),
"change_kind": (judge or {}).get("change_kind", "中性"),
"summary": summary,
"chain_trend": (judge or {}).get("chain_trend", ""),
"news_counts": {
"direct": _count(seg["direct"]),
"upstream": _count(seg["upstream"]),
"downstream": _count(seg["downstream"]),
"peers": _count(seg["peers"]),
},
"focus": focus,
}
sources = {
"direct": seg["direct"], "upstream": seg["upstream"],
"downstream": seg["downstream"], "peers": seg["peers"],
"indicators": _fmt_ind(seg["indicators"]),
"ratings": seg["ratings"], "holdings": seg["holdings"],
}
execute(
"INSERT INTO tracking_reports(code, stock_name, industry, report, meta, sources, status, created_at) "
"VALUES(?,?,?,?,?,?,'done',datetime('now','localtime'))",
(code, s["name"], s["industry"], clean_report, json.dumps(meta, ensure_ascii=False),
json.dumps(sources, ensure_ascii=False)))
rid = query_one("SELECT MAX(id) id FROM tracking_reports")["id"]
_notify_if_significant(rid, s, meta)
return {"ok": True, "report_id": rid, "meta": meta}
except Exception as e:
log.exception("track %s fail", code)
return {"error": str(e)}
def _count(seg):
return (seg["db"].count("[") + seg["rag"].count("[")) // 1
def _first_heading(text):
"""取报告第一行非空文本作为摘要兜底"""
for line in (text or "").splitlines():
line = line.strip().lstrip("#* ").strip()
if line:
return line[:60]
return ""
def _notify_if_significant(rid, stock, meta):
"""影响度达阈值且开启通知 → 邮件"""
cfg = tracking_config()
try:
if cfg["notify"] and int(meta["impact_score"]) >= int(cfg["impact_threshold"]):
from engine.notifier import send_email
mc = mail_config()
send_email(
f"[持仓跟踪] {stock['name']} 出现{meta.get('change_kind','')}动态(影响度{meta['impact_score']}",
f"""<html><body style="font-family:Microsoft YaHei;padding:20px;background:#f5f6f8;">
<div style="max-width:640px;margin:auto;background:#fff;border-radius:8px;border:1px solid #e5e7eb;overflow:hidden;">
<div style="background:#1e293b;color:#fff;padding:14px 20px;font-size:17px;font-weight:bold;">🧭 持仓跟踪 · {stock['name']}{stock['code']}</div>
<div style="padding:16px 20px;">
<p><b>影响度:</b>{meta['impact_score']}/100{'🔴 重大' if meta['impact_score']>=65 else '🟡 关注'}<br>
<b>性质:</b>{meta.get('change_kind','')} <b>显著性:</b>{meta.get('significance','')}</p>
<p style="font-size:15px;"><b>摘要:</b>{meta.get('summary','')}</p>
<p style="color:#555;"><b>产业链趋势:</b>{meta.get('chain_trend','')}</p>
<p style="color:#888;font-size:12px;">个股资讯 {meta['news_counts']['direct']} 条 / 上游 {meta['news_counts']['upstream']} 条 / 下游 {meta['news_counts']['downstream']} 条 / 同业 {meta['news_counts']['peers']} 条</p>
</div></div></body></html>""",
cfg=mc)
set_tracking_state(last_alert=int(meta["impact_score"]))
except Exception as e:
log.warning("track notify fail: %s", e)
# ===================================================================== 批量与调度
def track_watchlist(progress=None):
"""串行跟踪自选股(持仓)。同一时刻只允许一个跟踪任务(防重复)"""
global _cycle_running
if _cycle_running:
return {"tracked": 0, "msg": "已有跟踪任务进行中,请稍后再试"}
_cycle_running = True
try:
stocks = query("SELECT w.code, s.name FROM watchlist w JOIN stocks s ON s.code=w.code ORDER BY w.added_at")
if not stocks:
return {"tracked": 0, "msg": "自选股为空,请先在股票池添加"}
results = []
for i, st in enumerate(stocks):
r = track_stock(st["code"])
results.append({"code": st["code"], "name": st["name"], **r})
set_tracking_state(last_run=time.strftime("%Y-%m-%d %H:%M:%S"), last_stock=st["name"])
if progress:
progress(i + 1, len(stocks))
return {"tracked": len(results), "results": results}
finally:
_cycle_running = False
def latest_reports(code, limit=5):
return query("SELECT id, code, stock_name, industry, meta, status, created_at "
"FROM tracking_reports WHERE code=? ORDER BY id DESC LIMIT ?", (code, limit))
def list_reports(limit=30):
return query("SELECT id, code, stock_name, industry, meta, status, created_at "
"FROM tracking_reports ORDER BY id DESC LIMIT ?", (limit,))
def get_report(rid):
return query_one("SELECT * FROM tracking_reports WHERE id=?", (rid,))
class TrackingThread(threading.Thread):
"""后台调度:定期跟踪全部持仓股票"""
def __init__(self):
super().__init__(daemon=True, name="tracking")
self._stop = threading.Event()
def stop(self):
self._stop.set()
def run(self):
log.info("持仓跟踪调度器启动")
first = True
while not self._stop.is_set():
try:
cfg = tracking_config()
if first:
# 启动后等待一个完整间隔再首跑,避免重启即烧一轮 LLM、与手动操作冲突
self._stop.wait(cfg.get("interval_min", 60) * 60)
first = False
continue
if cfg["enabled"]:
try:
r = track_watchlist()
log.info("tracking cycle: %s", r)
except Exception as e:
log.warning("tracking cycle error: %s", e)
except Exception as e:
log.warning("tracking loop error: %s", e)
self._stop.wait(cfg.get("interval_min", 60) * 60)
log.info("持仓跟踪调度器停止")
_tracking = None
def start_tracking():
global _tracking
if _tracking and _tracking.is_alive():
return _tracking
_tracking = TrackingThread()
_tracking.start()
return _tracking