Initial commit: cc-cursor 全链路量化研究平台
7 Sprints 全部完成: Sprint 0: 基础设施 (DataManager + MariaDB) Sprint 1: 因子引擎 (34因子/12分类) Sprint 2: VectorBT 回测 (5策略+截面) Sprint 3: Optuna 优化 (+Walk-Forward) Sprint 4: ML 模型 (LightGBM+CatBoost) Sprint 5: Qwen 情绪因子 (三源新闻+日期对齐) Sprint 6: Agent 系统 (4Agent+日报.md/.html) 生产加固 (15项): Tushare双源fallback, SSH自动恢复, pool_pre_ping, save_daily先删后插, load_dotenv绝对路径, 日报5d/20d修复, RiskAgent改上证指数, 昨日对比+数据截止, mac_report utf8mb4, CLAUDE-*.md 9条已知Bug, demo全参数化, djapi数据源归一化, indexDatas API修正 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,373 @@
|
||||
"""
|
||||
新闻公告数据源。
|
||||
|
||||
支持三种数据来源:
|
||||
1. AkShare stock_news_em(东方财富个股新闻)
|
||||
2. MariaDB xwlb_daily_ext(新闻联播分割数据)
|
||||
3. MCP trendradar-news(外部新闻聚合服务)
|
||||
|
||||
统一返回格式:DataFrame (date, title, content, source, url)
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
import akshare as ak
|
||||
import pandas as pd
|
||||
import requests
|
||||
|
||||
# 确保 .env 已加载
|
||||
import config.settings # noqa: F401
|
||||
|
||||
|
||||
class NewsSource:
|
||||
"""
|
||||
新闻数据源聚合器。
|
||||
|
||||
参数:
|
||||
use_akshare: 启用 AkShare 新闻接口
|
||||
use_xwlb: 启用新闻联播数据库
|
||||
use_mcp: 启用 MCP 新闻服务
|
||||
mcp_url: MCP 服务地址
|
||||
request_delay: 请求间隔(秒),避免被限频
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
use_akshare: bool = True,
|
||||
use_xwlb: bool = True,
|
||||
use_mcp: bool = True,
|
||||
mcp_url: str | None = None,
|
||||
request_delay: float = 1.0,
|
||||
):
|
||||
self.use_akshare = use_akshare
|
||||
self.use_xwlb = use_xwlb
|
||||
self.use_mcp = use_mcp
|
||||
self.mcp_url = mcp_url or os.getenv("NEWS_MCP_URL", "http://192.168.1.160:3333/mcp")
|
||||
self.request_delay = request_delay
|
||||
self._mcp_session_id: str | None = None
|
||||
|
||||
# ── 统一入口 ──────────────────────────────────────────
|
||||
|
||||
def fetch(
|
||||
self,
|
||||
ts_code: str,
|
||||
start: str | None = None,
|
||||
end: str | None = None,
|
||||
max_news: int | None = None,
|
||||
) -> pd.DataFrame:
|
||||
"""
|
||||
获取个股相关新闻(聚合多源)。
|
||||
|
||||
参数:
|
||||
ts_code: 如 '000001.SZ'
|
||||
start: 起始日期 'YYYYMMDD',默认 30 天前
|
||||
end: 结束日期 'YYYYMMDD',默认今天
|
||||
max_news: 最多返回条数,默认从 .env SENTIMENT_MAX_NEWS_PER_STOCK 读取,fallback 30
|
||||
|
||||
返回:
|
||||
DataFrame (date, title, content, source, url)
|
||||
"""
|
||||
if max_news is None:
|
||||
max_news = int(os.getenv("SENTIMENT_MAX_NEWS_PER_STOCK", "30"))
|
||||
|
||||
if end is None:
|
||||
end = time.strftime("%Y%m%d")
|
||||
if start is None:
|
||||
start = (datetime.strptime(end, "%Y%m%d") - timedelta(days=30)).strftime("%Y%m%d")
|
||||
|
||||
frames = []
|
||||
|
||||
if self.use_akshare:
|
||||
try:
|
||||
df = self._fetch_akshare(ts_code, start, end)
|
||||
if not df.empty:
|
||||
frames.append(df)
|
||||
time.sleep(self.request_delay)
|
||||
except Exception as e:
|
||||
print(" [WARN] AkShare 新闻获取失败 ({}): {}".format(ts_code, e))
|
||||
|
||||
if self.use_xwlb:
|
||||
try:
|
||||
# xwlb: news_date 是播出日,+1day = 交易日
|
||||
# 所以查询时 start/end 各减 1 天
|
||||
xwlb_start = (datetime.strptime(start, "%Y%m%d") - timedelta(days=1)).strftime("%Y%m%d")
|
||||
xwlb_end = (datetime.strptime(end, "%Y%m%d") - timedelta(days=1)).strftime("%Y%m%d")
|
||||
df = self._fetch_xwlb(xwlb_start, xwlb_end)
|
||||
if not df.empty:
|
||||
frames.append(df)
|
||||
except Exception as e:
|
||||
print(" [WARN] xwlb 新闻获取失败: {}".format(e))
|
||||
|
||||
if self.use_mcp:
|
||||
try:
|
||||
df = self._fetch_mcp(ts_code, start, end)
|
||||
if not df.empty:
|
||||
frames.append(df)
|
||||
except Exception as e:
|
||||
print(" [WARN] MCP 新闻获取失败: {}".format(e))
|
||||
|
||||
if not frames:
|
||||
return pd.DataFrame(columns=["date", "title", "content", "source", "url"])
|
||||
|
||||
result = pd.concat(frames, ignore_index=True)
|
||||
result = result.drop_duplicates(subset=["title", "date"])
|
||||
result = result.sort_values("date", ascending=False)
|
||||
|
||||
if len(result) > max_news:
|
||||
result = result.head(max_news)
|
||||
|
||||
return result.reset_index(drop=True)
|
||||
|
||||
# ── AkShare 数据源 ────────────────────────────────────
|
||||
|
||||
def _fetch_akshare(self, ts_code: str, start: str = "", end: str = "") -> pd.DataFrame:
|
||||
"""
|
||||
东方财富个股新闻。
|
||||
|
||||
stock_news_em 返回最新约 10 条(无日期筛选),
|
||||
如需更多应使用 InfoMoney 或自建爬虫。
|
||||
"""
|
||||
symbol = ts_code.replace(".SZ", "").replace(".SH", "").replace(".BJ", "")
|
||||
try:
|
||||
df = ak.stock_news_em(symbol=symbol.zfill(6))
|
||||
except Exception:
|
||||
return pd.DataFrame()
|
||||
|
||||
if df is None or df.empty:
|
||||
return pd.DataFrame()
|
||||
|
||||
df = df.rename(columns={
|
||||
"新闻标题": "title",
|
||||
"新闻内容": "content",
|
||||
"发布时间": "date",
|
||||
"文章来源": "source",
|
||||
"新闻链接": "url",
|
||||
})
|
||||
df["date"] = pd.to_datetime(df["date"]).dt.strftime("%Y%m%d")
|
||||
df["source"] = "akshare_" + df["source"].fillna("unknown")
|
||||
|
||||
# 日期筛选
|
||||
if start:
|
||||
df = df[df["date"] >= start]
|
||||
if end:
|
||||
df = df[df["date"] <= end]
|
||||
|
||||
return df[["date", "title", "content", "source", "url"]]
|
||||
|
||||
# ── 新闻联播数据源 ────────────────────────────────────
|
||||
|
||||
def _fetch_xwlb(self, start: str, end: str) -> pd.DataFrame:
|
||||
"""
|
||||
从 MariaDB xwlb_daily_ext 表获取新闻联播分割数据。
|
||||
|
||||
news_date 是播出日期,+1day 后成为对市场产生影响的交易日。
|
||||
调用方应传入 (target_start-1, target_end-1)。
|
||||
|
||||
连接失败时 database/connection.py 自动重建 SSH 隧道。
|
||||
"""
|
||||
from database.connection import get_engine
|
||||
|
||||
start_fmt = "{}-{}-{}".format(start[:4], start[4:6], start[6:8])
|
||||
end_fmt = "{}-{}-{}".format(end[:4], end[4:6], end[6:8])
|
||||
|
||||
day_span = (datetime.strptime(end, "%Y%m%d") - datetime.strptime(start, "%Y%m%d")).days
|
||||
row_limit = max(100, (day_span + 1) * 40)
|
||||
|
||||
sql = """
|
||||
SELECT news_date, news_title, news_content
|
||||
FROM xwlb_daily_ext
|
||||
WHERE news_date BETWEEN %(start)s AND %(end)s
|
||||
ORDER BY news_date DESC
|
||||
LIMIT {}
|
||||
""".format(row_limit)
|
||||
|
||||
try:
|
||||
df = pd.read_sql(
|
||||
sql, get_engine(),
|
||||
params={"start": start_fmt, "end": end_fmt},
|
||||
)
|
||||
if df.empty:
|
||||
return pd.DataFrame()
|
||||
|
||||
df = df.rename(columns={
|
||||
"news_date": "date",
|
||||
"news_title": "title",
|
||||
"news_content": "content",
|
||||
})
|
||||
df["date"] = (
|
||||
pd.to_datetime(df["date"]) + pd.Timedelta(days=1)
|
||||
).dt.strftime("%Y%m%d")
|
||||
df["source"] = "xwlb"
|
||||
df["url"] = ""
|
||||
return df[["date", "title", "content", "source", "url"]]
|
||||
|
||||
except Exception:
|
||||
return pd.DataFrame()
|
||||
|
||||
# ── MCP 数据源 ────────────────────────────────────────
|
||||
|
||||
def _fetch_mcp(self, ts_code: str, start: str, end: str) -> pd.DataFrame:
|
||||
"""通过 MCP streamable HTTP 调用 trendradar-news 服务。"""
|
||||
session_id = self._mcp_initialize()
|
||||
if not session_id:
|
||||
return pd.DataFrame()
|
||||
|
||||
symbol = ts_code.replace(".SZ", "").replace(".SH", "").replace(".BJ", "")
|
||||
try:
|
||||
result = self._mcp_call_tool(
|
||||
session_id,
|
||||
"search_news",
|
||||
{"keyword": symbol, "start_date": start, "end_date": end, "limit": 30},
|
||||
)
|
||||
if not result:
|
||||
return pd.DataFrame()
|
||||
|
||||
records = self._parse_mcp_result(result)
|
||||
if not records:
|
||||
return pd.DataFrame()
|
||||
|
||||
df = pd.DataFrame(records)
|
||||
df["source"] = "mcp_trendradar"
|
||||
return df[["date", "title", "content", "source", "url"]]
|
||||
except Exception:
|
||||
return pd.DataFrame()
|
||||
|
||||
def _mcp_initialize(self) -> str | None:
|
||||
"""MCP streamable HTTP 初始化,获取 session ID。"""
|
||||
if self._mcp_session_id:
|
||||
return self._mcp_session_id
|
||||
try:
|
||||
resp = requests.post(
|
||||
self.mcp_url,
|
||||
json={
|
||||
"jsonrpc": "2.0",
|
||||
"method": "initialize",
|
||||
"params": {
|
||||
"protocolVersion": "2024-11-05",
|
||||
"capabilities": {},
|
||||
"clientInfo": {"name": "cc-cursor", "version": "1.0"},
|
||||
},
|
||||
"id": 1,
|
||||
},
|
||||
headers={
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json, text/event-stream",
|
||||
},
|
||||
timeout=10,
|
||||
)
|
||||
# headers 是大小写不敏感的,尝试多种格式
|
||||
session_id = (
|
||||
resp.headers.get("mcp-session-id")
|
||||
or resp.headers.get("Mcp-Session-Id")
|
||||
or resp.headers.get("MCP-Session-Id")
|
||||
)
|
||||
if session_id:
|
||||
self._mcp_session_id = session_id
|
||||
else:
|
||||
print(" [WARN] MCP initialize 未返回 session-id")
|
||||
return session_id
|
||||
except Exception as e:
|
||||
print(" [WARN] MCP 连接失败: {}".format(e))
|
||||
return None
|
||||
|
||||
def _mcp_call_tool(
|
||||
self, session_id: str, tool_name: str, arguments: dict
|
||||
) -> dict | None:
|
||||
"""调用 MCP 工具。"""
|
||||
try:
|
||||
resp = requests.post(
|
||||
self.mcp_url,
|
||||
json={
|
||||
"jsonrpc": "2.0",
|
||||
"method": "tools/call",
|
||||
"params": {"name": tool_name, "arguments": arguments},
|
||||
"id": 2,
|
||||
},
|
||||
headers={
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json, text/event-stream",
|
||||
"mcp-session-id": session_id,
|
||||
},
|
||||
timeout=15,
|
||||
)
|
||||
# 解析 SSE 响应
|
||||
for line in resp.text.split("\n"):
|
||||
if line.startswith("data:"):
|
||||
data = json.loads(line[5:].strip())
|
||||
return data.get("result", {})
|
||||
return None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _parse_mcp_result(result: dict) -> list[dict]:
|
||||
"""解析 MCP 返回内容为统一格式。"""
|
||||
records = []
|
||||
content = result.get("content", [])
|
||||
for item in content:
|
||||
text = item.get("text", "") if isinstance(item, dict) else str(item)
|
||||
if text:
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
if isinstance(parsed, list):
|
||||
for p in parsed:
|
||||
records.append({
|
||||
"date": str(p.get("date", p.get("publish_date", ""))),
|
||||
"title": str(p.get("title", "")),
|
||||
"content": str(p.get("content", p.get("summary", ""))),
|
||||
"url": str(p.get("url", "")),
|
||||
})
|
||||
elif isinstance(parsed, dict):
|
||||
records.append({
|
||||
"date": str(parsed.get("date", parsed.get("publish_date", ""))),
|
||||
"title": str(parsed.get("title", "")),
|
||||
"content": str(parsed.get("content", parsed.get("summary", ""))),
|
||||
"url": str(parsed.get("url", "")),
|
||||
})
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
records.append({
|
||||
"date": time.strftime("%Y%m%d"),
|
||||
"title": text[:100],
|
||||
"content": text,
|
||||
"url": "",
|
||||
})
|
||||
return records
|
||||
|
||||
|
||||
# ── 日期对齐工具 ──────────────────────────────────────────
|
||||
|
||||
def align_news_to_trading_days(
|
||||
news_df: pd.DataFrame,
|
||||
trading_calendar: pd.DatetimeIndex,
|
||||
) -> pd.DataFrame:
|
||||
"""
|
||||
将新闻日期对齐到最近交易日。
|
||||
|
||||
非交易日的新闻归入最近的下一个交易日。
|
||||
例:周六新闻 → 下周一交易日的因子值中体现。
|
||||
|
||||
参数:
|
||||
news_df: 新闻 DataFrame,含 'date' 列 (YYYYMMDD)
|
||||
trading_calendar: 交易日 DatetimeIndex
|
||||
|
||||
返回:
|
||||
news_df,'date' 列已替换为对齐后的交易日
|
||||
"""
|
||||
calendar = trading_calendar.sort_values()
|
||||
|
||||
news_dates = pd.to_datetime(news_df["date"], format="%Y%m%d")
|
||||
aligned = []
|
||||
|
||||
for nd in news_dates:
|
||||
future_dates = calendar[calendar >= nd]
|
||||
if len(future_dates) > 0:
|
||||
aligned.append(future_dates[0])
|
||||
else:
|
||||
aligned.append(calendar[-1])
|
||||
|
||||
news_df = news_df.copy()
|
||||
news_df["date"] = pd.DatetimeIndex(aligned).strftime("%Y%m%d")
|
||||
return news_df
|
||||
@@ -0,0 +1,229 @@
|
||||
"""
|
||||
Qwen API 封装(DashScope / 本地 Ollama)。
|
||||
|
||||
文本 → 情绪分析 → 标准化 JSON 输出。
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
import requests
|
||||
|
||||
# 确保 .env 已加载(无论从哪个路径导入)
|
||||
import config.settings # noqa: F401
|
||||
|
||||
|
||||
class QwenClient:
|
||||
"""
|
||||
Qwen 情绪分析客户端。
|
||||
|
||||
支持两种后端:
|
||||
1. DashScope API(云)- 需要 QWEN_API_KEY
|
||||
2. 本地 Ollama(本地)- 需要 QWEN_LOCAL_BASE_URL
|
||||
|
||||
优先级:本地 > API
|
||||
|
||||
参数:
|
||||
api_key: DashScope API Key,默认从环境变量 QWEN_API_KEY 读取
|
||||
model: 模型名称,默认从 QWEN_MODEL 读取
|
||||
local_base_url: Ollama 地址,默认从 QWEN_LOCAL_BASE_URL 读取
|
||||
local_model: Ollama 模型名,默认从 QWEN_LOCAL_MODEL 读取
|
||||
max_retries: 失败重试次数
|
||||
cache_enabled: 是否启用结果缓存(按文本 hash)
|
||||
"""
|
||||
|
||||
# 情绪分析 Prompt
|
||||
SYSTEM_PROMPT = """你是一个专业的金融情绪分析专家。分析以下 A 股相关的新闻/公告文本,返回 JSON 格式的分析结果。
|
||||
|
||||
分析维度:
|
||||
1. sentiment_score: -1.0(极度利空) ~ 0(中性) ~ +1.0(极度利好),保留1位小数
|
||||
2. impact_duration: "short"(1-3个交易日) | "medium"(1-2周) | "long"(1个月以上)
|
||||
3. confidence: 0.0~1.0 置信度
|
||||
4. key_topics: 涉及的关键主题列表(最多5个)
|
||||
5. affected_factors: 可能影响的基本面/技术面因子类型列表
|
||||
|
||||
重要规则:
|
||||
- 只考虑对股价的直接影响
|
||||
- 中性新闻(常规报道、例行公告)给 0 分
|
||||
- 重大利好(业绩超预期、政策支持、大订单)给 >0.5 分
|
||||
- 重大利空(亏损、处罚、减持、诉讼)给 <-0.5 分
|
||||
- 如果不确定,confidences 要低 (<0.5)
|
||||
|
||||
返回格式(仅 JSON):
|
||||
{"sentiment_score": 0.0, "impact_duration": "short", "confidence": 0.5, "key_topics": [], "affected_factors": []}"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
model: str | None = None,
|
||||
local_base_url: str | None = None,
|
||||
local_model: str | None = None,
|
||||
max_retries: int = 3,
|
||||
cache_enabled: bool = True,
|
||||
):
|
||||
self.api_key = api_key or os.getenv("QWEN_API_KEY", "")
|
||||
self.model = model or os.getenv("QWEN_MODEL", "qwen-turbo")
|
||||
self.local_base_url = local_base_url or os.getenv("QWEN_LOCAL_BASE_URL", "")
|
||||
self.local_model = local_model or os.getenv("QWEN_LOCAL_MODEL", "qwen2.5:7b")
|
||||
self.max_retries = max_retries
|
||||
self.cache_enabled = cache_enabled
|
||||
self._cache: dict[str, dict] = {}
|
||||
self._is_local = bool(self.local_base_url)
|
||||
|
||||
# ── 单条分析 ──────────────────────────────────────────
|
||||
|
||||
def analyze_sentiment(self, text: str) -> dict:
|
||||
"""
|
||||
分析单条文本的情绪。
|
||||
|
||||
返回:
|
||||
{"sentiment_score": float, "impact_duration": str, "confidence": float,
|
||||
"key_topics": list, "affected_factors": list}
|
||||
"""
|
||||
if not text or not text.strip():
|
||||
return self._empty_result()
|
||||
|
||||
cache_key = str(hash(text))
|
||||
if self.cache_enabled and cache_key in self._cache:
|
||||
return self._cache[cache_key]
|
||||
|
||||
for attempt in range(self.max_retries):
|
||||
try:
|
||||
if self._is_local:
|
||||
result = self._call_ollama(text)
|
||||
else:
|
||||
result = self._call_dashscope(text)
|
||||
|
||||
if result and "sentiment_score" in result:
|
||||
if self.cache_enabled:
|
||||
self._cache[cache_key] = result
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
if attempt < self.max_retries - 1:
|
||||
time.sleep(1 * (attempt + 1))
|
||||
else:
|
||||
print(f" [ERROR] Qwen 分析失败(重试{self.max_retries}次): {e}")
|
||||
|
||||
return self._empty_result()
|
||||
|
||||
# ── 批量分析 ──────────────────────────────────────────
|
||||
|
||||
def analyze_batch(
|
||||
self,
|
||||
texts: list[str],
|
||||
batch_size: int = 10,
|
||||
progress_callback: callable = None, # type: ignore
|
||||
) -> list[dict]:
|
||||
"""
|
||||
批量分析(逐条调用,节省并发成本)。
|
||||
|
||||
参数:
|
||||
texts: 文本列表
|
||||
batch_size: 每批次数量
|
||||
progress_callback: 进度回调 fn(done, total)
|
||||
|
||||
返回:
|
||||
[{"sentiment_score": ..., ...}, ...]
|
||||
"""
|
||||
results = []
|
||||
total = len(texts)
|
||||
for i, text in enumerate(texts):
|
||||
result = self.analyze_sentiment(text)
|
||||
results.append(result)
|
||||
if progress_callback:
|
||||
progress_callback(i + 1, total)
|
||||
# 批次间短暂休息,避免触发限频
|
||||
if (i + 1) % batch_size == 0 and i < total - 1:
|
||||
time.sleep(0.5)
|
||||
return results
|
||||
|
||||
# ── DashScope API ──────────────────────────────────────
|
||||
|
||||
def _call_dashscope(self, text: str) -> dict:
|
||||
"""调用 DashScope API。"""
|
||||
resp = requests.post(
|
||||
"https://dashscope.aliyuncs.com/compatible-mode/v1/chat/completions",
|
||||
headers={
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
json={
|
||||
"model": self.model,
|
||||
"messages": [
|
||||
{"role": "system", "content": self.SYSTEM_PROMPT},
|
||||
{"role": "user", "content": text[:4000]}, # 截断长文本
|
||||
],
|
||||
"temperature": 0.1,
|
||||
"max_tokens": 500,
|
||||
},
|
||||
timeout=30,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
body = resp.json()
|
||||
content = body["choices"][0]["message"]["content"]
|
||||
return self._parse_json(content)
|
||||
|
||||
# ── Ollama 本地调用 ────────────────────────────────────
|
||||
|
||||
def _call_ollama(self, text: str) -> dict:
|
||||
"""调用本地 Ollama。"""
|
||||
resp = requests.post(
|
||||
f"{self.local_base_url.rstrip('/')}/chat/completions",
|
||||
json={
|
||||
"model": self.local_model,
|
||||
"messages": [
|
||||
{"role": "system", "content": self.SYSTEM_PROMPT},
|
||||
{"role": "user", "content": text[:4000]},
|
||||
],
|
||||
"temperature": 0.1,
|
||||
"max_tokens": 500,
|
||||
},
|
||||
timeout=60,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
body = resp.json()
|
||||
content = body["choices"][0]["message"]["content"]
|
||||
return self._parse_json(content)
|
||||
|
||||
# ── JSON 解析 ──────────────────────────────────────────
|
||||
|
||||
@staticmethod
|
||||
def _parse_json(content: str) -> dict:
|
||||
"""从模型输出中提取 JSON。"""
|
||||
# 尝试直接解析
|
||||
try:
|
||||
return json.loads(content)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
# 尝试提取 ```json ... ``` 代码块
|
||||
if "```json" in content:
|
||||
start = content.find("```json") + 7
|
||||
end = content.find("```", start)
|
||||
if end > start:
|
||||
try:
|
||||
return json.loads(content[start:end].strip())
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
# 尝试找 { ... }
|
||||
brace_start = content.find("{")
|
||||
brace_end = content.rfind("}")
|
||||
if brace_start >= 0 and brace_end > brace_start:
|
||||
try:
|
||||
return json.loads(content[brace_start:brace_end + 1])
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
return QwenClient._empty_result()
|
||||
|
||||
@staticmethod
|
||||
def _empty_result() -> dict:
|
||||
return {
|
||||
"sentiment_score": 0.0,
|
||||
"impact_duration": "short",
|
||||
"confidence": 0.0,
|
||||
"key_topics": [],
|
||||
"affected_factors": [],
|
||||
}
|
||||
@@ -0,0 +1,311 @@
|
||||
"""
|
||||
情绪因子计算引擎。
|
||||
|
||||
全链路:数据获取 → Qwen 分析 → 因子计算 → 缓存管理。
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from data.data_manager import DataManager
|
||||
from factors.sentiment.news_source import NewsSource, align_news_to_trading_days
|
||||
from factors.sentiment.qwen_client import QwenClient
|
||||
from factors.sentiment.sentiment_factor import (
|
||||
NewsSentimentFactor,
|
||||
SentimentMomentumFactor,
|
||||
SentimentConfidenceFactor,
|
||||
)
|
||||
|
||||
|
||||
def _get_or_fetch_price(dm, ts_code, fallback_code="000001.SZ"):
|
||||
"""
|
||||
获取交易日历(价格 DataFrame)。
|
||||
|
||||
1. 优先从 DB 缓存读取目标股票
|
||||
2. 无缓存则 sync_daily 补齐
|
||||
3. 补齐失败则 fallback 到备用股票(仅用作交易日历,不影响新闻分析)
|
||||
"""
|
||||
import pandas as pd
|
||||
from database.dao import get_latest_trade_date
|
||||
|
||||
# 检查 DB 缓存
|
||||
if get_latest_trade_date(ts_code):
|
||||
daily = dm.get_daily(ts_code)
|
||||
if daily is not None and not daily.empty:
|
||||
return daily.set_index("trade_date").sort_index()
|
||||
|
||||
# 无缓存 → 尝试补齐
|
||||
print(" [SentimentEngine] {} 无 DB 缓存,尝试 sync_daily...".format(ts_code))
|
||||
try:
|
||||
n = dm.sync_daily(ts_code)
|
||||
if n > 0:
|
||||
daily = dm.get_daily(ts_code)
|
||||
if daily is not None and not daily.empty:
|
||||
return daily.set_index("trade_date").sort_index()
|
||||
except Exception as e:
|
||||
print(" [SentimentEngine] sync_daily 失败: {}".format(e))
|
||||
|
||||
# Fallback
|
||||
print(" [SentimentEngine] 使用 {} 交易日历作为 fallback".format(fallback_code))
|
||||
try:
|
||||
fallback = dm.get_daily(fallback_code)
|
||||
if fallback is not None and not fallback.empty:
|
||||
return fallback.set_index("trade_date").sort_index()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return None
|
||||
|
||||
|
||||
class SentimentEngine:
|
||||
"""
|
||||
情绪因子计算引擎。
|
||||
|
||||
参数:
|
||||
data_manager: DataManager 实例
|
||||
qwen_client: QwenClient 实例(可选,默认从环境变量创建)
|
||||
news_source: NewsSource 实例(可选)
|
||||
cache_dir: 缓存目录
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
data_manager: DataManager,
|
||||
qwen_client: QwenClient | None = None,
|
||||
news_source: NewsSource | None = None,
|
||||
cache_dir: str | None = None,
|
||||
):
|
||||
self._dm = data_manager
|
||||
self._qwen = qwen_client or QwenClient()
|
||||
self._news_source = news_source or NewsSource()
|
||||
self._cache_dir = Path(cache_dir or os.path.join(
|
||||
os.path.dirname(__file__), "..", "..", ".cache", "sentiment"
|
||||
))
|
||||
self._cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# ── 核心方法 ──────────────────────────────────────────
|
||||
|
||||
def compute(
|
||||
self,
|
||||
ts_code: str,
|
||||
start: str | None = None,
|
||||
end: str | None = None,
|
||||
max_news: int = 20,
|
||||
) -> pd.DataFrame:
|
||||
"""
|
||||
计算单只股票的情绪因子(含缓存)。
|
||||
|
||||
返回:
|
||||
DataFrame (trade_date, news_sent_5, news_conf_5, sent_delta_5)
|
||||
"""
|
||||
from database.dao import get_latest_trade_date
|
||||
|
||||
# 1. 获取交易日历:优先目标股票 DB 缓存,无则尝试补齐,失败则 fallback
|
||||
price = _get_or_fetch_price(self._dm, ts_code)
|
||||
if price is None:
|
||||
return pd.DataFrame()
|
||||
daily_idx = pd.to_datetime(price.index, format="%Y%m%d", errors="coerce")
|
||||
|
||||
# 2. 检查缓存
|
||||
cache_key = "sent_{}_{}_{}".format(ts_code, start or "all", end or "now")
|
||||
cached = self._load_cache(cache_key)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
# 3. 获取新闻
|
||||
max_n = max_news or int(os.getenv("SENTIMENT_MAX_NEWS_PER_STOCK", "30"))
|
||||
news_df = self._news_source.fetch(ts_code, start=start, end=end, max_news=max_n)
|
||||
|
||||
# 4. 对齐到交易日
|
||||
if not news_df.empty:
|
||||
news_df = align_news_to_trading_days(news_df, daily_idx)
|
||||
|
||||
# 5. Qwen 情绪分析(有 API key 时才执行)
|
||||
sentiment_df = self._analyze_news(news_df)
|
||||
|
||||
# 6. 计算因子
|
||||
factor_dfs = {}
|
||||
if not sentiment_df.empty:
|
||||
for factor_cls, kwargs in [
|
||||
(NewsSentimentFactor, {"window": 5, "sentiment_df": sentiment_df}),
|
||||
(SentimentConfidenceFactor, {"window": 5, "sentiment_df": sentiment_df}),
|
||||
(SentimentMomentumFactor, {"period": 5, "sentiment_df": sentiment_df}),
|
||||
]:
|
||||
f = factor_cls(**kwargs)
|
||||
series = f.calculate(price)
|
||||
factor_dfs[f.name] = series
|
||||
|
||||
if factor_dfs:
|
||||
result = pd.DataFrame(factor_dfs)
|
||||
else:
|
||||
result = pd.DataFrame(index=price.index)
|
||||
|
||||
# 7. 缓存
|
||||
self._save_cache(cache_key, result)
|
||||
|
||||
return result
|
||||
|
||||
def compute_batch(
|
||||
self,
|
||||
ts_codes: list[str],
|
||||
date: str | None = None,
|
||||
max_news: int = 20,
|
||||
) -> dict[str, pd.DataFrame]:
|
||||
"""
|
||||
批量计算多只股票的情绪因子。
|
||||
|
||||
返回:
|
||||
{ts_code: factor_df}
|
||||
"""
|
||||
results = {}
|
||||
total = len(ts_codes)
|
||||
for i, ts_code in enumerate(ts_codes):
|
||||
try:
|
||||
results[ts_code] = self.compute(
|
||||
ts_code, start=None if date is None else date, end=date,
|
||||
max_news=max_news,
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"[WARN] {ts_code} 情绪因子失败: {e}")
|
||||
if (i + 1) % 10 == 0:
|
||||
print(f"[SentimentEngine] {i + 1}/{total}")
|
||||
return results
|
||||
|
||||
# ── 分析范围解析 ──────────────────────────────────────
|
||||
|
||||
def get_scope_stocks(self) -> list[str]:
|
||||
"""
|
||||
根据 .env 配置解析分析范围。
|
||||
|
||||
优先级: custom > index > sector > all
|
||||
"""
|
||||
scope_type = os.getenv("SENTIMENT_SCOPE_TYPE", "index")
|
||||
all_stocks = self._dm.get_stock_list()
|
||||
|
||||
if scope_type == "custom":
|
||||
codes = os.getenv("SENTIMENT_SCOPE_CUSTOM", "")
|
||||
return [c.strip() for c in codes.split(",") if c.strip()]
|
||||
|
||||
if scope_type == "sector":
|
||||
return self._get_sector_stocks(all_stocks)
|
||||
|
||||
if scope_type == "index":
|
||||
return self._get_index_stocks()
|
||||
|
||||
# all
|
||||
return list(all_stocks.index)
|
||||
|
||||
def _get_index_stocks(self) -> list[str]:
|
||||
"""根据指数成分股获取股票列表。"""
|
||||
indexes_str = os.getenv("SENTIMENT_SCOPE_INDEXES", "000300")
|
||||
indexes = [i.strip() for i in indexes_str.split(",") if i.strip()]
|
||||
stocks = set()
|
||||
import akshare as ak
|
||||
for idx_code in indexes:
|
||||
try:
|
||||
df = ak.index_stock_cons(symbol=idx_code)
|
||||
if df is not None and not df.empty:
|
||||
# 列名可能是 '品种代码' 或 'stock_code'
|
||||
col = next((c for c in df.columns if "代码" in c or "code" in c.lower()), df.columns[0])
|
||||
for code in df[col]:
|
||||
stocks.add(f"{code}.SZ" if code.startswith("0") or code.startswith("3") else f"{code}.SH")
|
||||
except Exception as e:
|
||||
print(f" [WARN] 指数 {idx_code} 成分股获取失败: {e}")
|
||||
return list(stocks)
|
||||
|
||||
def _get_sector_stocks(self, all_stocks: pd.DataFrame) -> list[str]:
|
||||
"""根据板块名称筛选股票。"""
|
||||
sectors_str = os.getenv("SENTIMENT_SCOPE_SECTORS", "")
|
||||
sectors = [s.strip() for s in sectors_str.split(",") if s.strip()]
|
||||
if not sectors:
|
||||
return list(all_stocks.index)
|
||||
# 利用 AkShare 行业分类筛选
|
||||
import akshare as ak
|
||||
all_industry = ak.stock_board_industry_name_em()
|
||||
matched = all_industry[all_industry["板块名称"].isin(sectors)]
|
||||
stocks = set()
|
||||
for _, row in matched.iterrows():
|
||||
try:
|
||||
df = ak.stock_board_industry_cons_em(symbol=row["板块名称"])
|
||||
if df is not None and not df.empty:
|
||||
code_col = next((c for c in df.columns if "代码" in c), df.columns[0])
|
||||
for code in df[code_col]:
|
||||
stocks.add(code)
|
||||
except Exception:
|
||||
continue
|
||||
return list(stocks)
|
||||
|
||||
# ── 新闻情绪分析 ──────────────────────────────────────
|
||||
|
||||
def _analyze_news(self, news_df: pd.DataFrame) -> pd.DataFrame:
|
||||
"""
|
||||
对新闻列表执行情绪分析。
|
||||
|
||||
返回:
|
||||
DataFrame (date, sentiment_score, confidence, impact_duration, key_topics, title, content)
|
||||
"""
|
||||
if news_df.empty:
|
||||
return pd.DataFrame()
|
||||
|
||||
# 检查是否有 Qwen API 配置
|
||||
if not self._qwen.api_key and not self._qwen.local_base_url:
|
||||
return pd.DataFrame()
|
||||
|
||||
results = []
|
||||
for _, row in news_df.iterrows():
|
||||
title = row.get("title", "")
|
||||
content = row.get("content", "")
|
||||
# 优先用标题+内容,文本过短时只用标题
|
||||
text = f"{title}\n{content}" if len(str(content)) > 20 else title
|
||||
if len(text.strip()) < 10:
|
||||
continue
|
||||
|
||||
analysis = self._qwen.analyze_sentiment(text)
|
||||
results.append({
|
||||
"date": row["date"],
|
||||
"title": title,
|
||||
"sentiment_score": analysis.get("sentiment_score", 0.0),
|
||||
"confidence": analysis.get("confidence", 0.0),
|
||||
"impact_duration": analysis.get("impact_duration", "short"),
|
||||
"key_topics": json.dumps(analysis.get("key_topics", [])),
|
||||
})
|
||||
|
||||
if not results:
|
||||
return pd.DataFrame()
|
||||
|
||||
return pd.DataFrame(results)
|
||||
|
||||
# ── 缓存管理 ──────────────────────────────────────────
|
||||
|
||||
def _load_cache(self, key: str) -> pd.DataFrame | None:
|
||||
"""加载缓存。缓存有效期 6 小时。"""
|
||||
cache_file = self._cache_dir / f"{key.replace('/', '_')}.parquet"
|
||||
if not cache_file.exists():
|
||||
return None
|
||||
mtime = cache_file.stat().st_mtime
|
||||
if time.time() - mtime > 6 * 3600:
|
||||
return None
|
||||
try:
|
||||
return pd.read_parquet(cache_file)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def _save_cache(self, key: str, df: pd.DataFrame) -> None:
|
||||
cache_file = self._cache_dir / f"{key.replace('/', '_')}.parquet"
|
||||
try:
|
||||
df.to_parquet(cache_file)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def clear_cache(self) -> int:
|
||||
"""清空缓存,返回删除文件数。"""
|
||||
count = 0
|
||||
for f in self._cache_dir.glob("*.parquet"):
|
||||
f.unlink()
|
||||
count += 1
|
||||
return count
|
||||
@@ -0,0 +1,168 @@
|
||||
"""
|
||||
情绪因子。
|
||||
|
||||
将 Qwen 输出的情绪分数转换为量化因子值。
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
|
||||
from factors.base import BaseFactor
|
||||
|
||||
|
||||
class NewsSentimentFactor(BaseFactor):
|
||||
"""
|
||||
新闻情绪因子。
|
||||
|
||||
将多条新闻的情绪分数按时间加权聚合到每个交易日。
|
||||
|
||||
参数:
|
||||
window: 滚动窗口(交易日)
|
||||
decay: 指数衰减系数(越大衰减越快),0 表示等权
|
||||
sentiment_df: 情绪分析结果 DataFrame
|
||||
(date, sentiment_score, confidence, title, content)
|
||||
"""
|
||||
|
||||
category = "sentiment"
|
||||
|
||||
def __init__(self, window: int = 5, decay: float = 0.3, sentiment_df: pd.DataFrame | None = None):
|
||||
self.window = window
|
||||
self.decay = decay
|
||||
self.sentiment_df = sentiment_df
|
||||
self.name = f"news_sent_{window}"
|
||||
|
||||
def calculate(self, df: pd.DataFrame) -> pd.Series:
|
||||
if self.sentiment_df is None or self.sentiment_df.empty:
|
||||
return pd.Series(float("nan"), index=df.index, name=self.name)
|
||||
|
||||
return _aggregate_sentiment(
|
||||
df, self.sentiment_df, self.window, self.decay
|
||||
)
|
||||
|
||||
def get_required_columns(self) -> list[str]:
|
||||
return []
|
||||
|
||||
|
||||
class SentimentMomentumFactor(BaseFactor):
|
||||
"""
|
||||
情绪动量因子。
|
||||
|
||||
当前情绪 - N 日前情绪,衡量情绪变化方向。
|
||||
"""
|
||||
|
||||
category = "sentiment"
|
||||
|
||||
def __init__(self, period: int = 5, sentiment_df: pd.DataFrame | None = None):
|
||||
self.period = period
|
||||
self.sentiment_df = sentiment_df
|
||||
self.name = f"sent_delta_{period}"
|
||||
|
||||
def calculate(self, df: pd.DataFrame) -> pd.Series:
|
||||
base = NewsSentimentFactor(window=1, sentiment_df=self.sentiment_df).calculate(df)
|
||||
return base.diff(self.period)
|
||||
|
||||
def get_required_columns(self) -> list[str]:
|
||||
return []
|
||||
|
||||
|
||||
class SentimentConfidenceFactor(BaseFactor):
|
||||
"""
|
||||
情绪置信度因子。
|
||||
|
||||
新闻情绪分析的置信度越高,因子绝对值越大(方向同 sentiment_score)。
|
||||
sentiment_score × confidence → 高置信利好=正值大,高置信利空=负值大。
|
||||
"""
|
||||
|
||||
category = "sentiment"
|
||||
|
||||
def __init__(self, window: int = 5, sentiment_df: pd.DataFrame | None = None):
|
||||
self.window = window
|
||||
self.sentiment_df = sentiment_df
|
||||
self.name = f"news_conf_{window}"
|
||||
|
||||
def calculate(self, df: pd.DataFrame) -> pd.Series:
|
||||
if self.sentiment_df is None or self.sentiment_df.empty:
|
||||
return pd.Series(float("nan"), index=df.index, name=self.name)
|
||||
|
||||
sdf = self.sentiment_df.copy()
|
||||
# 设置加权分数
|
||||
if "confidence" in sdf.columns and "sentiment_score" in sdf.columns:
|
||||
sdf["weighted_score"] = sdf["sentiment_score"] * sdf["confidence"]
|
||||
else:
|
||||
return pd.Series(float("nan"), index=df.index, name=self.name)
|
||||
|
||||
return _aggregate_sentiment(df, sdf, self.window, decay=0.3, score_col="weighted_score")
|
||||
|
||||
def get_required_columns(self) -> list[str]:
|
||||
return []
|
||||
|
||||
|
||||
# ── 情绪聚合工具函数 ──────────────────────────────────────
|
||||
|
||||
def _aggregate_sentiment(
|
||||
daily_df: pd.DataFrame,
|
||||
sentiment_df: pd.DataFrame,
|
||||
window: int,
|
||||
decay: float,
|
||||
score_col: str = "sentiment_score",
|
||||
) -> pd.Series:
|
||||
"""
|
||||
将情绪分数按时间加权聚合到交易日。
|
||||
|
||||
逻辑:
|
||||
1. 对每个交易日 t,找到 [t - window + 1, t] 范围内的所有新闻
|
||||
2. 按 time_decay = exp(-decay * days_from_t) 加权
|
||||
3. 按 confidence(如有)加权
|
||||
4. 返回加权平均情绪分数
|
||||
|
||||
参数:
|
||||
daily_df: 日线 DataFrame(提供 index 和日期对齐)
|
||||
sentiment_df: 情绪 DataFrame(date 列 + score_col)
|
||||
window: 窗口大小
|
||||
decay: 衰减系数
|
||||
score_col: 情绪分数列名
|
||||
"""
|
||||
if sentiment_df.empty:
|
||||
return pd.Series(float("nan"), index=daily_df.index, name=score_col)
|
||||
|
||||
# 统一日期格式
|
||||
sdf = sentiment_df.copy()
|
||||
sdf["date"] = pd.to_datetime(sdf["date"], format="%Y%m%d", errors="coerce")
|
||||
sdf = sdf.dropna(subset=["date"])
|
||||
sdf = sdf.sort_values("date")
|
||||
|
||||
daily_idx = pd.to_datetime(daily_df.index, format="%Y%m%d", errors="coerce")
|
||||
if daily_idx.isna().all():
|
||||
daily_idx = pd.to_datetime(daily_df.index)
|
||||
|
||||
result = pd.Series(float("nan"), index=daily_df.index)
|
||||
|
||||
# 对 news 日期建立搜索索引
|
||||
news_dates = sdf["date"].values
|
||||
|
||||
for i, dt in enumerate(daily_idx):
|
||||
if pd.isna(dt):
|
||||
continue
|
||||
# 窗口起始
|
||||
window_start = dt - pd.Timedelta(days=window * 2) # 宽窗覆盖非交易日
|
||||
|
||||
mask = (news_dates >= window_start) & (news_dates <= dt)
|
||||
candidates = sdf[mask]
|
||||
|
||||
if candidates.empty:
|
||||
continue
|
||||
|
||||
# 时间衰减权重
|
||||
days_diff = (dt - candidates["date"]).dt.days
|
||||
time_weights = np.exp(-decay * days_diff)
|
||||
|
||||
# 置信度权重(如有)
|
||||
conf_weights = candidates.get("confidence", pd.Series(1.0, index=candidates.index)).fillna(0.5)
|
||||
scores = candidates[score_col].fillna(0.0)
|
||||
|
||||
total_weight = (time_weights * conf_weights).sum()
|
||||
if total_weight > 0:
|
||||
result.iloc[i] = (scores * time_weights * conf_weights).sum() / total_weight
|
||||
|
||||
result.name = score_col
|
||||
return result.astype("float64")
|
||||
Reference in New Issue
Block a user