"""B1-1 指数历史成分测试:Repository 快照与历史成分查询(Survivorship-free)、 Universe.index_code 过滤接线(选股/回测只用 as_of 当日成分)。""" from __future__ import annotations from datetime import date from decimal import Decimal import pytest from app.domain.entities.index import IndexWeight from app.infrastructure.persistence.sqlalchemy.base import Base from app.infrastructure.persistence.sqlalchemy.repositories.index_impl import ( SqlAlchemyIndexConstituentRepository, ) from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker _CODE = "000300.SH" @pytest.fixture() def session(tmp_path): engine = create_engine(f"sqlite:///{tmp_path / 'idx.db'}", future=True) Base.metadata.create_all(engine) Session = sessionmaker(bind=engine, expire_on_commit=False) with Session() as s: yield s def _rows() -> list[IndexWeight]: # 2024-06 期成分:A,C(B 被剔除);2024-11 期成分:A,B(C 新晋替换场景) return [ IndexWeight(index_code=_CODE, index_name="沪深300", trade_date=date(2024, 6, 28), symbol="600000.SH", weight=Decimal("1.0")), IndexWeight(index_code=_CODE, index_name="沪深300", trade_date=date(2024, 6, 28), symbol="600002.SH", weight=Decimal("1.0")), IndexWeight(index_code=_CODE, index_name="沪深300", trade_date=date(2024, 11, 29), symbol="600000.SH", weight=Decimal("1.5")), IndexWeight(index_code=_CODE, index_name="沪深300", trade_date=date(2024, 11, 29), symbol="600001.SH", weight=Decimal("1.0")), ] class TestIndexConstituentRepository: def test_upsert_and_members_at_history(self, session) -> None: repo = SqlAlchemyIndexConstituentRepository(session) assert repo.upsert_many(_rows()) == 4 session.commit() # 2024-07(最近快照 2024-06)→ {A,C};2025(最近 2024-11)→ {A,B} assert repo.members_at(_CODE, date(2024, 7, 15)) == {"600000.SH", "600002.SH"} assert repo.members_at(_CODE, date(2025, 1, 10)) == {"600000.SH", "600001.SH"} # 快照之前 → 空集(不返回未来成分) assert repo.members_at(_CODE, date(2024, 1, 1)) == set() assert repo.latest_date(_CODE) == date(2024, 11, 29) def test_upsert_idempotent(self, session) -> None: repo = SqlAlchemyIndexConstituentRepository(session) repo.upsert_many(_rows()) session.commit() repo.upsert_many([_rows()[2]]) # 重复 session.commit() assert len(repo.members_at(_CODE, date(2024, 12, 31))) == 2 class TestUniverseIndexCodeFilter: def _build(self, tmp_path): engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True) Base.metadata.create_all(engine) Session = sessionmaker(bind=engine, expire_on_commit=False) from app.application.services.selection_service import SelectionService from app.domain.entities.market import Stock from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import ( SqlAlchemyDailyBarRepository, SqlAlchemyStockRepository, ) from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily df = synthetic_daily( {"600000.SH": 0.008, "600001.SH": 0.006, "600002.SH": 0.004}, n=320 ) with Session() as session: SqlAlchemyStockRepository(session).upsert_many( [ Stock(symbol="600000.SH", name="A", list_date=date(1999, 1, 1)), Stock(symbol="600001.SH", name="B", list_date=date(1999, 1, 1)), Stock(symbol="600002.SH", name="C", list_date=date(1999, 1, 1)), ] ) SqlAlchemyDailyBarRepository(session).upsert_many(bars_dataframe_to_daily_bars(df)) SqlAlchemyIndexConstituentRepository(session).upsert_many(_rows()) session.commit() svc = SelectionService( SqlAlchemyStockRepository(session), SqlAlchemyDailyBarRepository(session), index_repo=SqlAlchemyIndexConstituentRepository(session), ) from app.domain.entities.research import UniverseSpec from app.domain.entities.selection import SelectionQuery # as_of=2024-10(成分 {A,C})→ 只从 {A,C} 选,B 绝不进入 q = SelectionQuery( universe=UniverseSpec(exclude_st=False, min_listing_days=0, index_code=_CODE), factors=[{"name": "momentum_60", "weight": 1.0}], top_n=2, as_of=date(2024, 10, 15), ) res = svc.select(q) got = {c.symbol for c in res.candidates} assert got == {"600000.SH", "600002.SH"} # as_of=2025-01(成分 {A,B})→ C 不再可入选 q2 = SelectionQuery( universe=UniverseSpec(exclude_st=False, min_listing_days=0, index_code=_CODE), factors=[{"name": "momentum_60", "weight": 1.0}], top_n=2, as_of=date(2025, 1, 10), ) res2 = svc.select(q2) assert {c.symbol for c in res2.candidates} == {"600000.SH", "600001.SH"} def test_index_code_filter(self, tmp_path) -> None: self._build(tmp_path)