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:
@@ -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
|
||||
Reference in New Issue
Block a user