Files
qlib/backend/tests/test_tushare_provider.py
Simon 23972e7063 feat: 股息率案例口径 + 策略库与图表统一 + 回测存档完整化
汇总三轮未提交的开发(每轮均在本机 MariaDB + 真实浏览器上验证):

1) 股息率案例(全市场股息率最高 n 只,默认 20,每 m 月择股)
   - 新增日频估值表 daily_basic + 迁移;股息率因子(dv_ratio / dividend_yield / TTM)
   - 名称历史表 stock_name_history:剔除 ST 按**择股日当时名称**判定,消除
     「曾高股息后 ST」的股息陷阱(实测 3.70pp 偏差)
   - 区间择股/调仓双周期(m 择股 / y 调仓)、指数成分与白名单、停牌近似剔除
   - 复权因子口径核对(4,164,742 行、缺失 0.0%)、收盘价成交与涨跌停拦单
   - 案例实测:2020-01-01~2026-09-04 总收益 +24.86%(年化 3.52%、回撤 -28.58%)

2) 策略库与前端统一
   - strategy 表 + CRUD/PUT 原地更新 + `describe_strategy` 按 spec 真实推导
     「一句话说明 + 计算公式 + 执行步骤 + 注意事项」(与引擎实执行规则同源)
   - 任何出现股票代码处都成对显示名称且可点击进个股页
   - 全站图表基座统一 TradingView Lightweight Charts(ECharts 依赖、
     锁文件、组件与文档标注一并清除),买卖点标记只落在真实交易日上

3) 回测存档完整化(可往复查看)
   - 同步端点(POST /api/backtests、/api/factor-tests)此前完全不落库 → 现在同样归档,
     归档 id 经响应头 X-Experiment-Id 返回(不破坏 response_model)
   - data_version 首次真实写入(数据快照指纹:最新交易日 + 各表规模)
   - 个股收益曲线默认**全量保存**(此前硬截断 60 只);超出体积预算才裁剪,
     并写 archive_meta(机器可读)+ unimplemented(人可读)如实标注
   - 列表 kind/q 过滤 + X-Total-Count(此前 limit=50 静默截断)、DELETE 归档
   - 只读归档页 /experiments/{id}(Server Component,SSR 直出**选股条件**与
     **交易执行依据**);结果视图按 kind 分发(backtest/factor_test/selection),
     非回测归档不套用回测口径
   - 新增 CLI:prune_experiments(保留策略,默认 dry-run)、
     restore_experiment_from_job(从 Job 副本按原 id 重建被删的历史归档,默认 dry-run)

门禁:pytest 388 passed、ruff All checks passed、tsc 0 错误、图表单测 7 passed、
next build 成功、契约脚本 verify_strategy_workspace 59/59(含按 kind 逐类验证归档页)。
2026-09-20 07:31:04 +08:00

329 lines
12 KiB
Python
Raw Permalink 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 测试:归一化、重试与鉴权错误归类(用 Fake pro,不触网)。"""
from __future__ import annotations
from datetime import date
from decimal import Decimal
import pytest
from app.infrastructure.data_sources.errors import (
DataSourceAuthenticationError,
DataSourceError,
)
from app.infrastructure.data_sources.tushare import TushareProvider
class FakePro:
"""模拟 tushare.pro 客户端:方法返回 records(list[dict]) 或抛错。"""
def __init__(self, *, payload=None, error: Exception | None = None) -> None:
self.payload = payload or []
self.error = error
self.calls: list[str] = []
# 关键:必须记录关键字参数,否则「list_status 是否真的透传给 tushare」无法被断言
# (退市股同步完全依赖该参数,见 TestDelistedStocks)
self.call_kwargs: list[dict] = []
def __getattr__(self, api: str):
def _run(**kwargs):
self.calls.append(api)
self.call_kwargs.append(kwargs)
if self.error is not None:
raise self.error
return self.payload
return _run
def _pro(payload=None, error=None, retries: int = 2) -> TushareProvider:
return TushareProvider(
token="fake-token", pro=FakePro(payload=payload, error=error), max_retries=retries
)
class TestNormalize:
def test_stock_records(self) -> None:
stocks = TushareProvider.normalize_stock(
[
{
"ts_code": "600519.SH",
"symbol": "600519",
"name": "贵州茅台",
"area": "贵州",
"industry": "白酒",
"list_date": "20010827",
}
]
)
assert stocks[0].symbol == "600519.SH"
assert stocks[0].list_date == date(2001, 8, 27)
assert stocks[0].delist_date is None
def test_daily_volume_amount_scaled(self) -> None:
bars = TushareProvider.normalize_daily(
[
{
"ts_code": "600519.SH",
"trade_date": "20240102",
"open": "100.0",
"high": "101.5",
"low": "99.0",
"close": "100.5",
"vol": "10000.0",
"amount": "1010000.0",
}
]
)
bar = bars[0]
assert bar.trade_date == date(2024, 1, 2)
assert bar.volume == Decimal("1000000") # 手 → 股(×100)
assert bar.amount == Decimal("1010000000") # 千元 → 元(×1000)
def test_daily_nan_dropped(self) -> None:
bars = TushareProvider.normalize_daily(
[
{
"ts_code": "600519.SH",
"trade_date": "20240102",
"open": None,
"high": float("nan"),
"close": "10.0",
"vol": None,
"amount": None,
}
]
)
bar = bars[0]
assert bar.open is None
assert bar.high is None
assert bar.close == Decimal("10.0")
def test_financial_maps_announce_date(self) -> None:
rows = TushareProvider.normalize_financial(
[
{
"ts_code": "600519.SH",
"end_date": "20240630",
"ann_date": "20240831",
"eps": "1.23",
"roe": "15.5",
}
]
)
fin = rows[0]
assert fin.report_date == date(2024, 6, 30)
assert fin.announce_date == date(2024, 8, 31)
assert fin.eps == Decimal("1.23")
class TestCall:
def test_empty_result_returns_empty_list(self) -> None:
provider = _pro(payload=[])
assert provider.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 31)) == []
def test_success_records_returned(self) -> None:
payload = [
{
"ts_code": "600519.SH",
"trade_date": "20240102",
"open": "100",
"close": "101",
"vol": "1",
"amount": "1",
}
]
provider = _pro(payload=payload)
bars = provider.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 31))
assert len(bars) == 1
assert bars[0].symbol == "600519.SH"
def test_transient_error_retries_then_raises(self) -> None:
provider = _pro(error=RuntimeError("network down"), retries=2)
with pytest.raises(DataSourceError, match="重试 2 次仍失败"):
provider.get_stock_basic()
assert len(provider._pro.calls) == 2 # noqa: SLF001 —— 测试探针
def test_permission_error_raises_immediately(self) -> None:
provider = _pro(error=RuntimeError("抱歉,您没有访问该接口的权限,请升级积分"), retries=3)
with pytest.raises(DataSourceAuthenticationError):
provider.get_stock_basic()
assert len(provider._pro.calls) == 1 # noqa: SLF001
def test_missing_token_rejected(self) -> None:
with pytest.raises(DataSourceAuthenticationError, match="TUSHARE_TOKEN"):
TushareProvider(token="")
class TestRateLimitBackoff:
def test_rate_limit_retries_with_sleep(self) -> None:
"""频率超限:按退避等待后重试,最终抛错带原始信息(不当作鉴权错误)。"""
import app.infrastructure.data_sources.tushare as ts_mod
orig_sleep = ts_mod.time.sleep
sleeps: list[float] = []
ts_mod.time.sleep = lambda w: sleeps.append(w) # noqa: SLF001 —— 测试桩
try:
provider = _pro(
error=RuntimeError("抱歉,您访问接口(adj_factor)频率超限(200次/分钟)"), retries=3
)
provider._rate_limit_wait = 0.1 # noqa: SLF001
with pytest.raises(DataSourceError, match="频率超限"):
provider.get_stock_basic()
finally:
ts_mod.time.sleep = orig_sleep
assert len(provider._pro.calls) == 3 # noqa: SLF001 —— 完整重试 3 次
assert len(sleeps) >= 2
def _quarter_ends(count: int) -> list[str]:
"""最近 count 个季度末(YYYYMMDD,降序)。"""
ends: list[str] = []
y, m = 2026, 6
while len(ends) < count:
ends.append(f"{y}{m:02d}30" if m in (6, 9) else f"{y}{m:02d}31")
m -= 3
if m <= 0:
m += 12
y -= 1
return ends
class _FakeProQueue:
"""按调用顺序弹出 payload 的 Fake pro(模拟分页)。"""
def __init__(self, payloads: list[list[dict]]) -> None:
self.payloads = list(payloads)
self.kwargs: list[dict] = []
def fina_indicator(self, **kwargs):
self.kwargs.append(kwargs)
return self.payloads.pop(0)
class TestGetFinancialWindow:
def _record(self, end: str) -> dict:
return {"ts_code": "600519.SH", "end_date": end, "ann_date": end, "eps": "1.0"}
def test_window_args_passed(self) -> None:
ends = _quarter_ends(10)
fake = _FakeProQueue([[self._record(e) for e in ends]])
provider = TushareProvider(token="t", pro=fake)
rows = provider.get_financial("600519.SH", date(2024, 1, 1), date(2026, 6, 30))
assert len(rows) == 10
assert fake.kwargs[0]["start_date"] == "20240101"
assert fake.kwargs[0]["end_date"] == "20260630"
def test_paging_over_100_row_cap(self) -> None:
"""单请求最多 100 条 → 超过必须回卷报告期窗口继续取,老数据不丢。"""
newest = _quarter_ends(100)
older = _quarter_ends(140)[100:] # 100 条之外更早的 40 个季度
# 二次请求 end_date 必须早于首请求(分页回卷)
fake = _FakeProQueue(
[[self._record(e) for e in newest], [self._record(e) for e in older]]
)
provider = TushareProvider(token="t", pro=fake)
rows = provider.get_financial("600519.SH", date(2000, 1, 1), date(2026, 6, 30))
assert len(rows) == 100 + 40
assert len(fake.kwargs) == 2
assert fake.kwargs[1]["end_date"] < fake.kwargs[0]["end_date"]
assert {r.report_date for r in rows} == {date.fromisoformat(e[:4] + "-" + e[4:6] + "-" + e[6:]) for e in newest + older}
class TestIndexWeight:
def test_normalize_mapping(self) -> None:
from app.infrastructure.data_sources.tushare import TushareProvider
rows = TushareProvider.normalize_index_weight(
[
{"index_code": "000300.SH", "con_code": "600519.SH",
"trade_date": "20240628", "weight": 1.53},
{"con_code": "000001.SZ", "trade_date": "20240628", "weight": 0.9},
],
index_code_fallback="000300.SH",
)
assert len(rows) == 2
assert rows[0].index_code == "000300.SH"
assert rows[0].symbol == "600519.SH"
assert rows[0].trade_date.isoformat() == "2024-06-28"
assert float(rows[0].weight) == 1.53
# 无 index_code 时用 fallback;con_code 缺失跳过
assert rows[1].index_code == "000300.SH"
def test_provider_calls_index_weight(self) -> None:
from app.infrastructure.data_sources.tushare import TushareProvider
pro = FakePro(
payload=[{"index_code": "000300.SH", "con_code": "600519.SH",
"trade_date": "20240628", "weight": 1.0}]
)
p = TushareProvider(token="t", pro=pro)
rows = p.get_index_weight("000300.SH")
assert pro.calls == ["index_weight"]
assert len(rows) == 1 and rows[0].symbol == "600519.SH"
class TestDelistedStocks:
"""退市股(幸存者偏差修正):list_status 透传、NaN 归一、代码规范过滤。
实测背景:tushare `stock_basic` 不带 list_status 时**只返回在市股票**,
退市股整体缺失(230 只 2019-12 后退市)→ 回测系统性高估收益;
而退市股记录的 industry/area 是 NaN、status 为空,且含 T 前缀异常代码。
"""
def test_list_status_passed_to_api(self) -> None:
pro = FakePro(payload=[])
p = TushareProvider(token="t", pro=pro)
p.get_stock_basic("D")
assert pro.calls == ["stock_basic"]
# tushare 不带 list_status 时只返回在市股票 → 退市股同步必须真的传 "D"
assert pro.call_kwargs[0]["list_status"] == "D"
def test_default_list_status_is_listed_only(self) -> None:
pro = FakePro(payload=[])
p = TushareProvider(token="t", pro=pro)
p.get_stock_basic()
assert pro.calls == ["stock_basic"] # 默认 "L",保持既有行为
assert pro.call_kwargs[0]["list_status"] == "L"
def test_nan_optional_fields_become_none(self) -> None:
stocks = TushareProvider.normalize_stock(
[
{
"ts_code": "000005.SZ",
"name": "ST星源(退)",
"area": float("nan"),
"industry": float("nan"),
"market": float("nan"),
"exchange": "SZSE",
"list_date": "19901210",
"delist_date": "20240426",
}
]
)
s = stocks[0]
assert s.industry is None and s.area is None and s.market is None
assert s.exchange == "SZSE"
assert s.delist_date == date(2024, 4, 26)
def test_missing_status_falls_back_to_query_status(self) -> None:
"""退市表 status 为空:必须按查询的 list_status 兜底,不能一律标成 L。"""
rec = [{"ts_code": "000005.SZ", "name": "ST星源(退)", "list_date": "19901210",
"delist_date": "20240426", "status": None}]
assert TushareProvider.normalize_stock(rec)[0].status == "L" # 既有默认
assert TushareProvider.normalize_stock(rec, default_status="D")[0].status == "D"
def test_abnormal_code_skipped_not_fatal(self) -> None:
"""'T600018.SH'(上港集箱(退),2006 退市)不得让整批退市列表拉取失败。"""
pro = FakePro(
payload=[
{"ts_code": "000005.SZ", "name": "ST星源(退)", "list_date": "19901210",
"delist_date": "20240426"},
{"ts_code": "T600018.SH", "name": "上港集箱(退)", "list_date": "19960101",
"delist_date": "20061020"},
]
)
p = TushareProvider(token="t", pro=pro)
stocks = p.get_stock_basic("D")
assert [s.symbol for s in stocks] == ["000005.SZ"]