- 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
192 lines
7.0 KiB
Python
192 lines
7.0 KiB
Python
"""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
|