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

127 lines
4.7 KiB
Python

"""Failover 审计测试:主源失败 → 备用源兜底,每次尝试留 SyncLog。"""
from __future__ import annotations
from datetime import date
import pytest
from app.domain.entities.market import SyncLog
from app.infrastructure.data_sources.errors import (
DataSourceError,
DataSourceNotSupported,
)
from app.infrastructure.data_sources.failover import FailoverProvider
class PrimaryProvider:
name = "tushare"
def __init__(self, *, fail: bool = False, fail_message: str = "boom") -> None:
self.fail = fail
self.fail_message = fail_message
def get_daily(self, symbol, start, end):
if self.fail:
raise DataSourceError(self.fail_message)
return ["bar-ok"]
class FallbackProvider:
name = "sina"
def __init__(self, *, fail: bool = False, unsupported: bool = False) -> None:
self.fail = fail
self.unsupported = unsupported
def get_daily(self, symbol, start, end):
if self.unsupported:
raise DataSourceNotSupported("新浪不支持日线兜底")
if self.fail:
raise DataSourceError("sina down")
return ["bar-sina"]
def _logs() -> list[SyncLog]:
collected: list[SyncLog] = []
return collected
def _make(primary, fallback, collected) -> FailoverProvider:
return FailoverProvider(primary, fallback, audit=collected.append)
class TestFailover:
def test_primary_success_no_fallback(self) -> None:
collected: list[SyncLog] = []
f = _make(PrimaryProvider(), FallbackProvider(), collected)
assert f.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 5)) == ["bar-ok"]
assert len(collected) == 1
assert collected[0].source == "tushare"
assert collected[0].success is True
assert collected[0].row_count == 1
assert collected[0].data_start == date(2024, 1, 1)
def test_primary_fails_fallback_succeeds(self) -> None:
collected: list[SyncLog] = []
f = _make(PrimaryProvider(fail=True), FallbackProvider(), collected)
assert f.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 5)) == ["bar-sina"]
assert [log.source for log in collected] == ["tushare", "sina"]
assert collected[0].success is False
assert "boom" in (collected[0].failure_reason or "")
assert collected[1].success is True
def test_fallback_unsupported_raises_with_audit(self) -> None:
collected: list[SyncLog] = []
f = _make(PrimaryProvider(fail=True), FallbackProvider(unsupported=True), collected)
with pytest.raises(DataSourceError, match="备用源不支持"):
f.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 5))
assert len(collected) == 2
assert collected[1].success is False
def test_both_fail_raises_with_audit(self) -> None:
collected: list[SyncLog] = []
f = _make(PrimaryProvider(fail=True), FallbackProvider(fail=True), collected)
with pytest.raises(DataSourceError, match="主备数据源均失败"):
f.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 5))
assert len(collected) == 2
def test_no_fallback_raises(self) -> None:
collected: list[SyncLog] = []
f = FailoverProvider(PrimaryProvider(fail=True), None, audit=collected.append)
with pytest.raises(DataSourceError, match="无备用源"):
f.get_daily("600519.SH", date(2024, 1, 1), date(2024, 1, 5))
assert len(collected) == 1
class TestIndexFailover:
def test_primary_only_index(self) -> None:
"""指数成分:主源成功即用(新浪不支持会自动尝试并审计失败 → 结果仍来自主源)。"""
from datetime import date
from decimal import Decimal
from app.domain.entities.index import IndexWeight
from app.infrastructure.data_sources.failover import FailoverProvider
from app.infrastructure.data_sources.sina import SinaProvider
from app.infrastructure.data_sources.tushare import TushareProvider
class _P(TushareProvider):
name = "tushare"
def __init__(self): # 不触网
pass
def get_index_weight(self, index_code):
return [IndexWeight(index_code=index_code, index_name="沪深300",
trade_date=date(2024, 6, 28), symbol="600519.SH",
weight=Decimal("1"))]
logs: list[SyncLog] = []
def audit(log: SyncLog):
logs.append(log)
f = FailoverProvider(_P(), SinaProvider(), audit=audit)
rows = f.get_index_weight("000300.SH")
assert len(rows) == 1 and rows[0].symbol == "600519.SH"
assert len(logs) >= 1 # 每次尝试已审计