Files
qlib/backend/app/infrastructure/data_sources/tushare.py
T
Simon 23972e7063 feat: 股息率案例口径 + 策略库与图表统一 + 回测存档完整化
汇总三轮未提交的开发(每轮均在本机 MariaDB + 真实浏览器上验证):

1) 股息率案例(全市场股息率最高 n 只,默认 20,每 m 月择股)
   - 新增日频估值表 daily_basic + 迁移;股息率因子(dv_ratio / dividend_yield / TTM)
   - 名称历史表 stock_name_history:剔除 ST 按**择股日当时名称**判定,消除
     「曾高股息后 ST」的股息陷阱(实测 3.70pp 偏差)
   - 区间择股/调仓双周期(m 择股 / y 调仓)、指数成分与白名单、停牌近似剔除
   - 复权因子口径核对(4,164,742 行、缺失 0.0%)、收盘价成交与涨跌停拦单
   - 案例实测:2020-01-01~2026-09-04 总收益 +24.86%(年化 3.52%、回撤 -28.58%)

2) 策略库与前端统一
   - strategy 表 + CRUD/PUT 原地更新 + `describe_strategy` 按 spec 真实推导
     「一句话说明 + 计算公式 + 执行步骤 + 注意事项」(与引擎实执行规则同源)
   - 任何出现股票代码处都成对显示名称且可点击进个股页
   - 全站图表基座统一 TradingView Lightweight Charts(ECharts 依赖、
     锁文件、组件与文档标注一并清除),买卖点标记只落在真实交易日上

3) 回测存档完整化(可往复查看)
   - 同步端点(POST /api/backtests、/api/factor-tests)此前完全不落库 → 现在同样归档,
     归档 id 经响应头 X-Experiment-Id 返回(不破坏 response_model)
   - data_version 首次真实写入(数据快照指纹:最新交易日 + 各表规模)
   - 个股收益曲线默认**全量保存**(此前硬截断 60 只);超出体积预算才裁剪,
     并写 archive_meta(机器可读)+ unimplemented(人可读)如实标注
   - 列表 kind/q 过滤 + X-Total-Count(此前 limit=50 静默截断)、DELETE 归档
   - 只读归档页 /experiments/{id}(Server Component,SSR 直出**选股条件**与
     **交易执行依据**);结果视图按 kind 分发(backtest/factor_test/selection),
     非回测归档不套用回测口径
   - 新增 CLI:prune_experiments(保留策略,默认 dry-run)、
     restore_experiment_from_job(从 Job 副本按原 id 重建被删的历史归档,默认 dry-run)

门禁:pytest 388 passed、ruff All checks passed、tsc 0 错误、图表单测 7 passed、
next build 成功、契约脚本 verify_strategy_workspace 59/59(含按 kind 逐类验证归档页)。
2026-09-20 07:31:04 +08:00

475 lines
19 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.
"""Tushare Provider —— 首选数据源实现。
依赖注入:pro 客户端(tushare.pro.client 或测试 Fake)。真实运行时惰性加载
tushare 库(pyproject optional:uv sync --extra datasource-tushare)。
归一化函数只依赖 list[dict],便于无 pandas 环境下单测。
"""
from __future__ import annotations
import importlib
import logging
import re
import time
from datetime import date, datetime, timedelta
from decimal import Decimal
from typing import Any
from app.domain.entities.index import IndexWeight
from app.domain.entities.market import (
AdjustFactor,
DailyBar,
DailyBasic,
FinancialIndicator,
Stock,
StockNameHistory,
TradingCalendar,
)
from app.infrastructure.data_sources.errors import (
DataSourceAuthenticationError,
DataSourceError,
)
logger = logging.getLogger(__name__)
_TS_DATE = "%Y%m%d"
def _to_date(value: str | None) -> date | None:
"""日期归一:None / NaN / 空串 → None。
pandas 读到的缺失日期是 float NaN(namechange 的 end_date、部分财务字段),
若不拦住会抛 `time data 'nan' does not match format '%Y%m%d'` 并**中断整批拉取**
——实测 namechange 按年分片时 32/37 片因此失败。
"""
if value is None or value == "":
return None
if isinstance(value, float) and value != value: # NaN
return None
text = str(value).strip()
if text == "" or text.lower() in {"nan", "none", "null", "nat"}:
return None
return datetime.strptime(text[:10], _TS_DATE).date()
# 本地 symbol 规范:6 位数字 + 交易所后缀(与 Stock 实体的 pattern 校验一致)
_SYMBOL_RE = re.compile(r"^\d{6}\.(SH|SZ|BJ)$")
def _to_opt_str(value) -> str | None:
"""可选字符串字段归一:None/NaN/空串 → None。
pandas 读到的缺失值是 float NaN(如退市股的 industry/area),直接塞进
`str | None` 字段会被 pydantic 拒绝(string_type)——实测退市股拉取时命中。
"""
if value is None:
return None
if isinstance(value, float) and value != value: # NaN
return None
text = str(value).strip()
if text == "" or text.lower() in {"nan", "none", "null"}:
return None
return text
def _to_decimal(value) -> Decimal | None:
if value is None:
return None
try:
num = float(value)
except (ValueError, TypeError):
return None
if num != num: # NaN
return None
return Decimal(str(num))
class TushareProvider:
"""封装 Tushare Pro(ts.pro_api)。所有输出已归一化为领域实体。"""
name = "tushare"
def __init__(
self,
token: str = "",
*,
pro: object | None = None,
max_retries: int = 3,
rate_limit_wait: float = 30.0,
) -> None:
self._pro = pro if pro is not None else _build_pro(token)
self._max_retries = max_retries
self._rate_limit_wait = rate_limit_wait
# ---- 归一化(纯函数,输入 list[dict],可单测) ----
@staticmethod
def normalize_stock(
records: list[dict[str, Any]], default_status: str = "L"
) -> list[Stock]:
"""归一化为 Stock。
`default_status`:tushare `stock_basic(list_status='D')` 返回的 status 字段
为空(实测 None),若一律兜底成 "L" 会把退市股标成在市 → 调用方按查询的
list_status 传入,保证 status 与 delist_date 语义一致。
"""
stocks: list[Stock] = []
for rec in records:
stocks.append(
Stock(
symbol=str(rec.get("ts_code") or rec.get("symbol") or ""),
name=str(rec.get("name") or "").strip(),
industry=_to_opt_str(rec.get("industry")),
area=_to_opt_str(rec.get("area")),
market=_to_opt_str(rec.get("market")),
exchange=_to_opt_str(rec.get("exchange")),
list_date=_to_date(rec.get("list_date")) or date.min,
delist_date=_to_date(rec.get("delist_date")),
status=_to_opt_str(rec.get("status")) or default_status,
)
)
return stocks
@staticmethod
def normalize_calendar(records: list[dict[str, Any]]) -> list[TradingCalendar]:
return [
TradingCalendar(
calendar_date=_to_date(rec.get("cal_date")) or date.min,
is_open=bool(rec.get("is_open")),
)
for rec in records
]
@staticmethod
def normalize_index_weight(
records: list[dict[str, Any]], index_code_fallback: str = ""
) -> list[IndexWeight]:
"""index_weight 接口行 → IndexWeight(index_code/con_code/trade_date/weight)。"""
out: list[IndexWeight] = []
for rec in records:
code = str(rec.get("index_code") or index_code_fallback or "")
symbol = str(rec.get("con_code") or "")
if not code or not symbol:
continue
out.append(
IndexWeight(
index_code=code,
index_name=rec.get("index_name"),
trade_date=_to_date(rec.get("trade_date")) or date.min,
symbol=symbol,
weight=_to_decimal(rec.get("weight")),
)
)
return out
@staticmethod
def normalize_daily(records: list[dict[str, Any]]) -> list[DailyBar]:
bars: list[DailyBar] = []
for rec in records:
vol = _to_decimal(rec.get("vol"))
amount = _to_decimal(rec.get("amount"))
bars.append(
DailyBar(
symbol=str(rec.get("ts_code") or ""),
trade_date=_to_date(rec.get("trade_date")) or date.min,
source="tushare",
adjust="none",
open=_to_decimal(rec.get("open")),
high=_to_decimal(rec.get("high")),
low=_to_decimal(rec.get("low")),
close=_to_decimal(rec.get("close")),
volume=vol * 100 if vol is not None else None,
amount=amount * 1000 if amount is not None else None,
)
)
return bars
@staticmethod
def normalize_adj_factor(records: list[dict[str, Any]]) -> list[AdjustFactor]:
return [
AdjustFactor(
symbol=str(rec.get("ts_code") or ""),
trade_date=_to_date(rec.get("trade_date")) or date.min,
factor=_to_decimal(rec.get("adj_factor")) or Decimal(1),
)
for rec in records
]
@staticmethod
def normalize_daily_basic(records: list[dict[str, Any]]) -> list[DailyBasic]:
"""daily_basic → DailyBasic。
单位保持 Tushare 原样(不做隐式换算,避免口径漂移):
- dv_ratio / dv_ttm / turnover_rate / volume_ratio / pe / pb / ps … 为百分数或倍数
- total_share / float_share / free_share 单位万股;total_mv / circ_mv 单位万元
- close 为**不复权**收盘价,与 stock_daily(adjust=none) 同口径
"""
rows: list[DailyBasic] = []
for rec in records:
rows.append(
DailyBasic(
symbol=str(rec.get("ts_code") or ""),
trade_date=_to_date(rec.get("trade_date")) or date.min,
source="tushare",
close=_to_decimal(rec.get("close")),
turnover_rate=_to_decimal(rec.get("turnover_rate")),
volume_ratio=_to_decimal(rec.get("volume_ratio")),
pe=_to_decimal(rec.get("pe")),
pe_ttm=_to_decimal(rec.get("pe_ttm")),
pb=_to_decimal(rec.get("pb")),
ps=_to_decimal(rec.get("ps")),
ps_ttm=_to_decimal(rec.get("ps_ttm")),
dv_ratio=_to_decimal(rec.get("dv_ratio")),
dv_ttm=_to_decimal(rec.get("dv_ttm")),
total_share=_to_decimal(rec.get("total_share")),
float_share=_to_decimal(rec.get("float_share")),
free_share=_to_decimal(rec.get("free_share")),
total_mv=_to_decimal(rec.get("total_mv")),
circ_mv=_to_decimal(rec.get("circ_mv")),
)
)
return rows
@staticmethod
def normalize_name_history(records: list[dict[str, Any]]) -> list[StockNameHistory]:
"""namechange → StockNameHistory(名称生效区间)。
注意:`namechange` 的区间是**完整历史**(一行一个名称生效段),
`end_date` 为 NaN 表示「至今有效」;`change_reason` 为 ST/*ST/撤销ST 等。
"""
rows: list[StockNameHistory] = []
for rec in records:
symbol = _to_opt_str(rec.get("ts_code"))
start = _to_date(rec.get("start_date"))
name = _to_opt_str(rec.get("name"))
if not symbol or not start or not name or not _SYMBOL_RE.match(symbol):
continue
rows.append(
StockNameHistory(
symbol=symbol,
name=name,
start_date=start,
end_date=_to_date(rec.get("end_date")),
ann_date=_to_date(rec.get("ann_date")),
change_reason=_to_opt_str(rec.get("change_reason")),
source="tushare",
)
)
return rows
@staticmethod
def normalize_financial(records: list[dict[str, Any]]) -> list[FinancialIndicator]:
rows: list[FinancialIndicator] = []
for rec in records:
rows.append(
FinancialIndicator(
symbol=str(rec.get("ts_code") or ""),
report_date=_to_date(rec.get("end_date")) or date.min,
announce_date=_to_date(rec.get("ann_date")) or date.min,
eps=_to_decimal(rec.get("eps")),
roe=_to_decimal(rec.get("roe")),
net_profit=_to_decimal(rec.get("n_income_attr_p")),
gross_margin=_to_decimal(rec.get("grossprofit_margin")),
)
)
return rows
# ---- 接口调用 ----
def get_stock_basic(self, list_status: str = "L") -> list[Stock]:
"""股票基础信息(list_status: L=上市 / D=退市 / P=暂停上市)。
tushare `stock_basic` 不带 list_status 时**只返回在市股票**,因此
`delist_date` 恒为空、退市股整体缺失 → 回测存在幸存者偏差。
需要退市股时必须显式传 "D"(实测 2019-12 之后退市 230 只)。
"""
records = self._call(
"stock_basic",
list_status=list_status,
fields="ts_code,symbol,name,area,industry,market,exchange,list_date,delist_date,status",
)
# 代码规范过滤:tushare 退市表含极少数非本地代码规范的记录
# (实测 'T600018.SH' = 上港集箱(退),2006 年退市,T 前缀表示转入三板),
# 直接归一化会因 symbol 正则校验失败而**中断整个列表** —— 跳过并如实告警,
# 不做静默丢弃(AGENT.md §24)。
kept, skipped = [], []
for rec in records:
code = str(rec.get("ts_code") or rec.get("symbol") or "")
(kept if _SYMBOL_RE.match(code) else skipped).append(rec)
if skipped:
logger.warning(
"tushare.stock_basic(list_status=%s) 跳过 %d 条不符合本地代码规范的记录:%s",
list_status,
len(skipped),
[r.get("ts_code") for r in skipped[:5]],
)
return self.normalize_stock(kept, default_status=list_status)
def get_trade_cal(self, start: date, end: date) -> list[TradingCalendar]:
records = self._call(
"trade_cal",
exchange="SSE",
start_date=start.strftime(_TS_DATE),
end_date=end.strftime(_TS_DATE),
is_open="",
)
return self.normalize_calendar(records)
def get_daily(self, symbol: str, start: date, end: date) -> list[DailyBar]:
records = self._call(
"daily",
ts_code=symbol,
start_date=start.strftime(_TS_DATE),
end_date=end.strftime(_TS_DATE),
)
return self.normalize_daily(records)
def get_adjust_factor(self, symbol: str, start: date, end: date) -> list[AdjustFactor]:
records = self._call(
"adj_factor",
ts_code=symbol,
start_date=start.strftime(_TS_DATE),
end_date=end.strftime(_TS_DATE),
)
return self.normalize_adj_factor(records)
def get_financial(
self,
symbol: str,
start: date | None = None,
end: date | None = None,
) -> list[FinancialIndicator]:
"""fina_indicator:报告期窗口 + 100 条/请求上限自动分页。
Tushare 单次请求最多返回 100 条(超出按最新 100 条截断),因此
全量历史必须按报告期窗口回卷分页,否则老报告期会被静默丢弃。
"""
lo = start or date(1990, 1, 1)
hi = end or date.today()
raw: list[dict[str, Any]] = []
while lo <= hi:
batch = self._call(
"fina_indicator",
ts_code=symbol,
start_date=lo.strftime(_TS_DATE),
end_date=hi.strftime(_TS_DATE),
)
raw += batch
if len(batch) < 100:
break
ends = [
datetime.strptime(str(r["end_date"])[:8], _TS_DATE).date()
for r in batch
if r.get("end_date")
]
if not ends:
break
next_hi = min(ends) - timedelta(days=1)
if next_hi < lo: # 无进展保护(边界簇被截断等极端情况)
break
hi = next_hi
return self.normalize_financial(raw)
# ---- 内部 ----
_RATE_LIMIT_MARKERS = ("频率超限", "每分钟", "frequenc", "too many")
def get_index_weight(self, index_code: str) -> list[IndexWeight]:
"""指数成分(Tushare index_weight 全历史,ts_code 过滤)。"""
records = self._call("index_weight", ts_code=index_code)
return self.normalize_index_weight(records, index_code_fallback=index_code)
# Tushare 单次接口返回上限(实测 daily_basic 全市场单日 3700~5600 行、
# namechange 2020+ 区间 4031 行):取满即告警,避免静默截断。
MAX_ROWS_PER_CALL = 6000
# namechange 单次请求上限同样约 6000 行;实测 2020+ 区间仅 4031 行,
# 但全历史(1990 起)会超限 —— 由 Syncer 按年分片调用,避免静默截断。
_NAMECHANGE_FIELDS = "ts_code,name,start_date,end_date,ann_date,change_reason"
def get_name_changes(self, start: date, end: date) -> list[StockNameHistory]:
"""区间内全市场名称变更(Tushare namechange,按公告/生效区间批量取)。"""
records = self._call(
"namechange",
start_date=start.strftime(_TS_DATE),
end_date=end.strftime(_TS_DATE),
fields=self._NAMECHANGE_FIELDS,
)
if len(records) >= self.MAX_ROWS_PER_CALL:
logger.warning(
"tushare.namechange(%s~%s) 返回 %d 行,可能触及单次上限被截断,"
"请缩小区间后重跑",
start,
end,
len(records),
)
return self.normalize_name_history(records)
# daily_basic 单次请求上限 6000 行(全市场一日约 3700~5600 行),按交易日调用即可
_DAILY_BASIC_FIELDS = (
"ts_code,trade_date,close,turnover_rate,volume_ratio,pe,pe_ttm,pb,ps,ps_ttm,"
"dv_ratio,dv_ttm,total_share,float_share,free_share,total_mv,circ_mv"
)
def get_daily_basic(self, trade_date: date) -> list[DailyBasic]:
"""单交易日全市场每日指标(daily_basic)。
注意:Tushare 单次 6000 行上限 —— 全市场单日实测 3700~5600 行
(2020 年约 3700,2026 年约 5560),当前安全;但若未来上市公司数
逼近 6000,需要按 ts_code 分片。此处对「恰好取满 6000 行」做告警,
避免静默截断(AGENTS §7 数据可追溯)。
"""
records = self._call(
"daily_basic",
trade_date=trade_date.strftime(_TS_DATE),
fields=self._DAILY_BASIC_FIELDS,
)
if len(records) >= self.MAX_ROWS_PER_CALL:
logger.warning(
"daily_basic %s 返回 %d 行(达到 %d 行上限),可能被截断,需按 ts_code 分片",
trade_date,
len(records),
self.MAX_ROWS_PER_CALL,
)
return self.normalize_daily_basic(records)
def _call(self, api: str, **kwargs) -> list[dict[str, Any]]:
"""带限速退避的调用:频率超限按指数退避(最长 _rate_limit_wait)等待后重试。"""
last_error: Exception | None = None
for attempt in range(self._max_retries):
try:
fn = getattr(self._pro, api)
result = fn(**kwargs)
if result is None:
return []
if hasattr(result, "to_dict"):
return result.to_dict("records")
if isinstance(result, list):
return result
return []
except Exception as exc: # noqa: BLE001 —— tushare 异常无统一类型,逐一归类
last_error = exc
msg = str(exc)
if "权限" in msg or "积分" in msg or "token" in msg.lower():
raise DataSourceAuthenticationError(msg) from exc
if any(marker in msg for marker in self._RATE_LIMIT_MARKERS):
wait = min(self._rate_limit_wait, 2 ** (attempt + 1))
logger.warning("tushare.%s 频率超限,退避 %.1fs 后重试", api, wait)
time.sleep(wait)
raise DataSourceError(
f"tushare.{api} 重试 {self._max_retries} 次仍失败: {last_error}"
) from last_error
def _build_pro(token: str):
if not token:
raise DataSourceAuthenticationError(
"缺少 TUSHARE_TOKEN:请 cp .env.example .env 并填入 Tushare Pro token"
)
try:
ts = importlib.import_module("tushare")
except ImportError as exc: # pragma: no cover —— 环境相关
raise DataSourceError(
"未安装 tushare 客户端:cd backend && uv sync --extra datasource-tushare"
) from exc
return ts.pro_api(token)