Files
qlib/backend/tests/test_repositories.py
T
Simon 2da234220a 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
2026-09-06 16:59:28 +08:00

192 lines
7.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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