"""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 time from datetime import date, datetime, timedelta from decimal import Decimal from typing import Any from app.domain.entities.market import ( AdjustFactor, DailyBar, FinancialIndicator, Stock, 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: if value is None or value == "": return None return datetime.strptime(str(value)[:10], _TS_DATE).date() 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]]) -> list[Stock]: 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 ""), industry=rec.get("industry"), area=rec.get("area"), market=rec.get("market"), exchange=rec.get("exchange"), list_date=_to_date(rec.get("list_date")) or date.min, delist_date=_to_date(rec.get("delist_date")), status=str(rec.get("status") or "L"), ) ) 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_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_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[Stock]: records = self._call( "stock_basic", fields="ts_code,symbol,name,area,industry,market,exchange,list_date,delist_date,status", ) return self.normalize_stock(records) 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 _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)