Files
qlib/backend/tests/test_tushare_provider.py
T
Simon 93e32f4e63 feat(data): B1-2 指数成分同步(Provider + CLI sync index_weight)
- MarketDataProvider.get_index_weight(协议);Tushare 实现 normalize_index_weight +
  get_index_weight(ts_code=... 全历史成分权重);Sina 抛 DataSourceNotSupported;
  FailoverProvider 代理并审计每次尝试
- CLI:sync index_weight --code 000300.SH(拉取→幂等落库 index_weight→打印最新快照;
  失败走 sync_log 审计并返回非零)
- tests:Tushare 映射与调用(FakePro)、Failover 主源单源语义(新浪不支持被审计);
  全量 pytest 通过
2026-09-09 07:29:26 +08:00

260 lines
9.4 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 测试:归一化、重试与鉴权错误归类(用 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] = []
def __getattr__(self, api: str):
def _run(**kwargs):
self.calls.append(api)
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"