"""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