"""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 def test_stream_range_many_columns_subset_order_and_null(self, session: Session) -> None: """流式列裁剪:只返回所需数值列、SQL 侧转 float、按 symbol/trade_date 升序。""" repo = SqlAlchemyDailyBarRepository(session) bars = [ self._bar("2024-01-02"), self._bar("2024-01-03"), self._bar("2024-01-04"), ] other = [ DailyBar( symbol="000001.SZ", trade_date=d.trade_date, close=Decimal("9"), volume=Decimal("1"), ) for d in bars ] repo.upsert_many([*bars, *other]) session.commit() rows = list( repo.stream_range_many_columns( ["600519.SH"], date(2024, 1, 2), date(2024, 1, 4), ["close", "volume"] ) ) assert rows == [ ("600519.SH", "2024-01-02", 100.5, 10000.0), ("600519.SH", "2024-01-03", 100.5, 10000.0), ("600519.SH", "2024-01-04", 100.5, 10000.0), ] # NULL 数值 → None;白名单外列报错 null_bar = self._bar("2024-01-02").model_copy(update={"volume": None}) repo.upsert_many([null_bar]) session.commit() rows2 = list( repo.stream_range_many_columns( ["600519.SH"], date(2024, 1, 2), date(2024, 1, 2), ["volume"] ) ) assert rows2 == [("600519.SH", "2024-01-02", None)] with pytest.raises(ValueError): list( repo.stream_range_many_columns( ["600519.SH"], date(2024, 1, 2), date(2024, 1, 4), ["close", "nope"] ) ) 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") def test_source_marker_and_incremental_queries(self, session: Session) -> None: """source 来源标记(默认 tushare)+ 增量判断查询(has_report_period/list_symbol)。""" repo = SqlAlchemyFinancialRepository(session) repo.upsert_many( [ FinancialIndicator( symbol="600519.SH", report_date=date(2024, 6, 30), announce_date=date(2024, 8, 31), eps=Decimal("33.19"), ), # source 默认 tushare FinancialIndicator( symbol="600519.SH", report_date=date(2024, 9, 30), announce_date=date(2024, 10, 30), source="sina", eps=Decimal("48.42"), ), ] ) session.commit() rows = repo.list_symbol("600519.SH") assert {r.source for r in rows} == {"tushare", "sina"} assert rows[0].source == "tushare" # 按 announce_date 升序 assert repo.has_report_period("600519.SH", date(2024, 9, 30)) assert not repo.has_report_period("600519.SH", date(2024, 12, 31)) assert repo.list_symbol("000001.SZ") == [] 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