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
This commit is contained in:
Simon
2026-09-06 16:59:28 +08:00
parent 7a89d97c0b
commit 2da234220a
23 changed files with 2666 additions and 1 deletions
+64
View File
@@ -0,0 +1,64 @@
"""领域实体测试:字段约束与「防未来函数」可见性判断。"""
from __future__ import annotations
from datetime import date
from decimal import Decimal
import pytest
from app.domain.entities.market import (
DailyBar,
FinancialIndicator,
Stock,
SyncLog,
)
from pydantic import ValidationError
class TestStock:
def test_symbol_pattern_enforced(self) -> None:
Stock(symbol="600519.SH", name="贵州茅台", list_date=date(2001, 8, 27))
with pytest.raises(ValidationError):
Stock(symbol="600519", name="x", list_date=date(2001, 1, 1))
with pytest.raises(ValidationError):
Stock(symbol="sh600519", name="x", list_date=date(2001, 1, 1))
class TestDailyBar:
def test_is_complete(self) -> None:
bar = DailyBar(
symbol="600519.SH",
trade_date=date(2024, 1, 2),
open=Decimal("100"),
high=Decimal("101"),
low=Decimal("99"),
close=Decimal("100.5"),
volume=Decimal("10000"),
amount=Decimal("1000000"),
)
assert bar.is_complete
assert not DailyBar(symbol="600519.SH", trade_date=date(2024, 1, 2)).is_complete
class TestFinancialIndicator:
def _fin(self, announce: date) -> FinancialIndicator:
return FinancialIndicator(
symbol="600519.SH",
report_date=date(2024, 6, 30),
announce_date=announce,
eps=Decimal("1.2"),
)
def test_announced_by_after_announce(self) -> None:
fin = self._fin(date(2024, 8, 31))
# 公告日当天已可见;公告前不可见
assert fin.announced_by(date(2024, 8, 31))
assert not fin.announced_by(date(2024, 8, 30))
assert not fin.announced_by(date(2024, 6, 30)) # 报告期不代表公开
class TestSyncLog:
def test_defaults(self) -> None:
log = SyncLog(source="tushare", api="daily", success=True, row_count=5)
assert log.request_time is not None
assert log.data_start is None
+93
View File
@@ -0,0 +1,93 @@
"""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
+62
View File
@@ -0,0 +1,62 @@
"""Alembic 迁移测试:全新数据库 upgrade head 后应包含全部 Phase 1 表。"""
from __future__ import annotations
import sqlite3
from pathlib import Path
from alembic import command
from alembic.config import Config
BACKEND_ROOT = Path(__file__).resolve().parents[1]
def _alembic_config(db_path: Path) -> Config:
cfg = Config(str(BACKEND_ROOT / "alembic.ini"))
cfg.set_main_option("script_location", "app/infrastructure/persistence/migrations")
cfg.set_main_option("sqlalchemy.url", f"sqlite:///{db_path}")
return cfg
def test_upgrade_head_creates_phase1_tables(tmp_path) -> None:
db_path = tmp_path / "fresh.db"
command.upgrade(_alembic_config(db_path), "head")
con = sqlite3.connect(db_path)
try:
tables = {
row[0]
for row in con.execute("select name from sqlite_master where type='table'").fetchall()
}
finally:
con.close()
expected = {
"stock",
"stock_daily",
"adjust_factor",
"trading_calendar",
"financial_indicator",
"sync_log",
"alembic_version",
}
assert expected <= tables
# 关键防未来函数列存在
con = sqlite3.connect(db_path)
try:
fin_cols = {
row[1] for row in con.execute("pragma table_info(financial_indicator)").fetchall()
}
daily_cols = {row[1] for row in con.execute("pragma table_info(stock_daily)").fetchall()}
finally:
con.close()
assert {"report_date", "announce_date"} <= fin_cols
assert {"symbol", "trade_date", "close"} <= daily_cols
def test_upgrade_head_idempotent(tmp_path) -> None:
db_path = tmp_path / "again.db"
cfg = _alembic_config(db_path)
command.upgrade(cfg, "head")
command.upgrade(cfg, "head") # 二次执行不报错
+191
View File
@@ -0,0 +1,191 @@
"""Repository 集成测试:临时 SQLite 上的幂等 upsert / 查询 / 防未来函数过滤。"""
from __future__ import annotations
from datetime import date
from decimal import Decimal
import pytest
from app.domain.entities.market import (
AdjustFactor,
DailyBar,
FinancialIndicator,
Stock,
SyncLog,
TradingCalendar,
)
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.models.market import (
AdjustFactorModel,
FinancialIndicatorModel,
StockDailyModel,
StockModel,
SyncLogModel,
TradingCalendarModel,
)
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyAdjustFactorRepository,
SqlAlchemyDailyBarRepository,
SqlAlchemyFinancialRepository,
SqlAlchemyStockRepository,
SqlAlchemySyncLogRepository,
SqlAlchemyTradingCalendarRepository,
)
from sqlalchemy import create_engine, func, select
from sqlalchemy.orm import Session
@pytest.fixture()
def session(tmp_path) -> Session:
engine = create_engine(f"sqlite:///{tmp_path / 'repo.db'}", future=True)
Base.metadata.create_all(engine)
with Session(engine) as session:
yield session
def _count(session, model) -> int:
return session.scalar(select(func.count()).select_from(model))
class TestStockRepository:
def test_upsert_idempotent_and_update(self, session: Session) -> None:
repo = SqlAlchemyStockRepository(session)
s1 = Stock(symbol="600519.SH", name="贵州茅台", list_date=date(2001, 8, 27))
s2 = Stock(symbol="000001.SZ", name="平安银行", list_date=date(1991, 4, 3))
assert repo.upsert_many([s1, s2]) == 2
session.commit()
assert _count(session, StockModel) == 2
# 幂等:再次 upsert 不新增
repo.upsert_many([s1, s2])
session.commit()
assert _count(session, StockModel) == 2
# 更新既有记录
renamed = s1.model_copy(update={"name": "贵州茅台(更新)"})
repo.upsert_many([renamed])
session.commit()
got = repo.get_by_symbol("600519.SH")
assert got is not None
assert got.name == "贵州茅台(更新)"
class TestDailyBarRepository:
def _bar(self, day: str) -> DailyBar:
return DailyBar(
symbol="600519.SH",
trade_date=date.fromisoformat(day),
open=Decimal("100"),
high=Decimal("101"),
low=Decimal("99"),
close=Decimal("100.5"),
volume=Decimal("10000"),
amount=Decimal("1000000"),
)
def test_upsert_and_get_range(self, session: Session) -> None:
repo = SqlAlchemyDailyBarRepository(session)
bars = [self._bar("2024-01-02"), self._bar("2024-01-03"), self._bar("2024-01-04")]
repo.upsert_many(bars)
session.commit()
assert _count(session, StockDailyModel) == 3
repo.upsert_many([self._bar("2024-01-03")]) # 幂等
session.commit()
assert _count(session, StockDailyModel) == 3
got = repo.get_range("600519.SH", date(2024, 1, 3), date(2024, 1, 4))
assert [b.trade_date.isoformat() for b in got] == ["2024-01-03", "2024-01-04"]
assert repo.latest_date("600519.SH") == date(2024, 1, 4)
assert repo.latest_date("000001.SZ") is None
class TestFinancialRepository:
def _fin(self, announce: str, report: str = "2024-06-30") -> FinancialIndicator:
return FinancialIndicator(
symbol="600519.SH",
report_date=date.fromisoformat(report),
announce_date=date.fromisoformat(announce),
eps=Decimal("1.2"),
)
def test_list_announced_blocks_future(self, session: Session) -> None:
repo = SqlAlchemyFinancialRepository(session)
repo.upsert_many(
[
self._fin("2024-08-15"),
self._fin("2024-08-31"),
self._fin("2024-09-20"),
self._fin("2024-10-30", report="2024-09-30"),
] # Q3 财报
)
session.commit()
# as_of=2024-08-31:只能看到 08-15 与 08-31 两条公告
visible = repo.list_announced("600519.SH", as_of_date=date(2024, 8, 31))
assert len(visible) == 2
assert all(f.announce_date <= date(2024, 8, 31) for f in visible)
assert [f.announce_date.day for f in visible] == [15, 31]
# 报告期约束:只看 Q3 及以后(report_date >= 2024-09-01)
narrowed = repo.list_announced(
"600519.SH", as_of_date=date(2024, 12, 31), report_start=date(2024, 9, 1)
)
assert len(narrowed) == 1
assert narrowed[0].announce_date == date(2024, 10, 30)
def test_upsert_batch_duplicate_key_takes_latest(self, session: Session) -> None:
"""同一批内出现重复幂等键(数据源偶发)不得冲突,后值覆盖。"""
repo = SqlAlchemyFinancialRepository(session)
first = self._fin("2024-08-15")
later = self._fin("2024-08-15").model_copy(update={"eps": Decimal("9.99")})
repo.upsert_many([first, later])
session.commit()
assert _count(session, FinancialIndicatorModel) == 1
got = repo.list_announced("600519.SH", as_of_date=date(2024, 12, 31))
assert len(got) == 1
assert got[0].eps == Decimal("9.99")
class TestSyncLogRepository:
def test_add_and_recent(self, session: Session) -> None:
repo = SqlAlchemySyncLogRepository(session)
repo.add(SyncLog(source="tushare", api="daily", success=True, row_count=3))
repo.add(SyncLog(source="sina", api="daily", success=False, failure_reason="timeout"))
session.commit()
assert _count(session, SyncLogModel) == 2
recent = repo.recent(source="tushare", limit=10)
assert len(recent) == 1
assert recent[0].source == "tushare"
assert recent[0].row_count == 3
class TestOtherRepos:
def test_calendar_and_factor(self, session: Session) -> None:
cal = SqlAlchemyTradingCalendarRepository(session)
cal.upsert_many(
[
TradingCalendar(calendar_date=date(2024, 1, 2)),
TradingCalendar(calendar_date=date(2024, 1, 3), is_open=False),
]
)
session.commit()
assert cal.is_open(date(2024, 1, 2))
assert not cal.is_open(date(2024, 1, 3))
assert len(cal.list_range(date(2024, 1, 1), date(2024, 1, 5))) == 2
assert _count(session, TradingCalendarModel) == 2
adj = SqlAlchemyAdjustFactorRepository(session)
adj.upsert_many(
[
AdjustFactor(
symbol="600519.SH", trade_date=date(2024, 1, 2), factor=Decimal("12.3456")
)
]
)
session.commit()
factors = adj.get_range("600519.SH", date(2024, 1, 1), date(2024, 1, 31))
assert len(factors) == 1
assert float(factors[0].factor) == pytest.approx(12.3456)
assert _count(session, AdjustFactorModel) == 1
+77
View File
@@ -0,0 +1,77 @@
"""新浪 Provider 测试:代码转换、JSONP 解析、能力边界(不触网)。"""
from __future__ import annotations
from datetime import date
import pytest
from app.infrastructure.data_sources.errors import (
DataSourceError,
DataSourceNotSupported,
)
from app.infrastructure.data_sources.sina import SinaProvider, _extract_jsonp, _to_sina_symbol
_KLINE_OK = (
'var data=[{"day":"2024-08-30","open":"1700.0","high":"1720.0","low":"1690.0",'
'"close":"1710.0","volume":"20000"},{"day":"2024-08-31","open":"1710.0",'
'"high":"1725.0","low":"1705.0","close":"1720.0","volume":"18000"}]'
)
class _FakeResp:
def __enter__(self):
return self
def __exit__(self, *exc):
return False
def read(self):
return _KLINE_OK.encode("utf-8")
class _BoomResp(_FakeResp):
def read(self):
raise OSError("socket timeout")
class TestSymbolMap:
def test_mapping(self) -> None:
assert _to_sina_symbol("600519.SH") == "sh600519"
assert _to_sina_symbol("000001.SZ") == "sz000001"
assert _to_sina_symbol("830001.BJ") == "bj830001"
class TestJsonp:
def test_extract(self) -> None:
rows = _extract_jsonp(_KLINE_OK)
assert len(rows) == 2
assert rows[0]["close"] == "1710.0"
def test_bad_payload_raises(self) -> None:
with pytest.raises(DataSourceError, match="无法解析"):
_extract_jsonp("not jsonp")
class TestGetDaily:
def test_ok_with_date_filter(self) -> None:
provider = SinaProvider(urlopen=lambda _url, **kw: _FakeResp())
bars = provider.get_daily("600519.SH", date(2024, 8, 30), date(2024, 8, 30))
assert len(bars) == 1
assert bars[0].close is not None
assert bars[0].symbol == "600519.SH"
def test_network_error_wrapped(self) -> None:
provider = SinaProvider(urlopen=lambda _url, **kw: _BoomResp())
with pytest.raises(DataSourceError, match="sina 请求失败"):
provider.get_daily("600519.SH", date(2024, 8, 1), date(2024, 8, 31))
class TestCapabilities:
def test_not_supported(self) -> None:
provider = SinaProvider()
with pytest.raises(DataSourceNotSupported):
provider.get_stock_basic()
with pytest.raises(DataSourceNotSupported):
provider.get_adjust_factor("600519.SH", date(2024, 1, 1), date(2024, 1, 31))
with pytest.raises(DataSourceNotSupported):
provider.get_financial("600519.SH")
+150
View File
@@ -0,0 +1,150 @@
"""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="")