- financial 默认增量:按 A 股披露节奏判断已最新并跳过;--full 强制全量重拉 - Tushare fina_indicator 增加报告期窗口与 100 条/请求自动分页(修复老报告期静默截断) - 新浪兜底收紧为校验兜底:两源重叠历史一致才导入缺失键,行标记 source=sina; 财务可比字段取 eps/销售毛利率(ROE 两端口径不同不作依据),日线只比较最近重叠交易日 - CLI 输出逐只进度与导入内容描述(来源/行数/报告期与公告区间),失败股票留待重跑 - financial_indicator 增 source 列(迁移 d3f6c9a21b04);新增一致性/分页/服务测试
267 lines
9.3 KiB
Python
267 lines
9.3 KiB
Python
"""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)
|