"""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)