Files
qlib/backend/app/infrastructure/data_sources/tushare.py
T
Simon 442999f701 feat(data): 财务/日线同步增量 + 新浪「两边一致」校验兜底 + 逐只进度
- financial 默认增量:按 A 股披露节奏判断已最新并跳过;--full 强制全量重拉
- Tushare fina_indicator 增加报告期窗口与 100 条/请求自动分页(修复老报告期静默截断)
- 新浪兜底收紧为校验兜底:两源重叠历史一致才导入缺失键,行标记 source=sina;
  财务可比字段取 eps/销售毛利率(ROE 两端口径不同不作依据),日线只比较最近重叠交易日
- CLI 输出逐只进度与导入内容描述(来源/行数/报告期与公告区间),失败股票留待重跑
- financial_indicator 增 source 列(迁移 d3f6c9a21b04);新增一致性/分页/服务测试
2026-09-08 21:48:09 +08:00

267 lines
9.3 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 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)