Files
qlib/backend/app/infrastructure/data_sources/tushare.py
T
Simon 56254172b3 feat(data): Tushare 限速退避 + 新浪兜底(财务 getFinanceReport2022 / 日K 前复权),source+adjust 口径标记
- TushareProvider:频率超限按指数退避重试(不再一次 200/min 即中断),最长等待 30s
- SinaProvider 重构(参考 cc-cursor 公开接口实现):
  · 新增财务通道 CompanyFinanceService.getFinanceReport2022(source=gjzb) → FinancialIndicator
    (report_date / announce_date=publish_date),与 Tushare fina_indicator schema 一致
  · 日 K 保留 jsonp(前复权),统一 UA + 重试
  · 不支持方法仍抛 DataSourceNotSupported(复权因子/交易日历/基础信息)
- FailoverProvider 现在可对 daily 与 financial 兜底(CLI _failover_provider 接 SinaProvider)
- DailyBar + stock_daily 表新增 source/adjust 列:新浪兜底行标记 sina/qfq,
  Tushare 恢复后 --resume 按同键覆盖回不复权 → 两源格式一致且可追溯
- 迁移 91c4e27a03fb 已生成;执行需在全市场同步结束后:uv run alembic upgrade head
- 测试 34+ 项(新浪财务解析/格式一致/限速退避等)通过
2026-09-06 20:52:13 +08:00

234 lines
8.2 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
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) -> list[FinancialIndicator]:
records = self._call("fina_indicator", ts_code=symbol)
return self.normalize_financial(records)
# ---- 内部 ----
_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)