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

172 lines
5.9 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