- 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 通过
127 lines
4.7 KiB
Python
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 # 每次尝试已审计
|