Files
qlib/backend/tests/test_tushare_provider.py
T
Simon 2da234220a feat(backend): Phase 1 数据层 — Domain / Provider / Failover 审计 + 持久化 + 同步 CLI
- domain:市场数据实体(Stock / 交易日历 / 日线 / 复权 / 财务含 announce_date)+ Repository 与 MarketDataProvider Protocol
- 数据源:TushareProvider(归一化、重试、鉴权错误归类)、SinaProvider(备用,明确前复权口径与能力边界)、FailoverProvider + SyncLog 审计(禁止静默切换)
- 持久化:SQLAlchemy 2.x Models + Repository 实现(按业务键幂等 upsert、as_of_date 防未来函数过滤)+ Alembic 迁移
- CLI:uv run python -m app.cli.sync {basic|calendar|daily|financial|verify},支持 --resume 断点续传
- 真实 Tushare 验证:stock 5556 / 交易日历 366 / daily+factor 242 / 财务 55;sync_log 审计完整
- 测试:38 passed(domain / provider / failover / repository / 未来函数 / 迁移),ruff clean
2026-09-06 16:59:28 +08:00

151 lines
5.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="")