Files
qlib/backend/tests/test_tushare_provider.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

227 lines
8.0 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}