"""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 # 每次尝试已审计