From 9cc4bfccac23214023f065067ae92451ec2815c0 Mon Sep 17 00:00:00 2001 From: Simon Date: Wed, 9 Sep 2026 07:27:13 +0800 Subject: [PATCH] =?UTF-8?q?feat(universe):=20B1-1=20=E6=8C=87=E6=95=B0?= =?UTF-8?q?=E5=8E=86=E5=8F=B2=E6=88=90=E5=88=86=EF=BC=88index=5Fweight?= =?UTF-8?q?=EF=BC=89+=20Universe=20=E6=8C=89=20as=5Fof=20=E6=88=90?= =?UTF-8?q?=E5=88=86=E8=BF=87=E6=BB=A4?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - index_weight 表(migration f5e0d1c2b3a4,MySQL 已应用;index_code+date+symbol 唯一) + IndexWeight 实体 + IndexConstituentRepository(members_at:取 <=as_of 最近一期快照, Survivorship-free / 无未来成分;latest_date) - UniverseSpec.index_code + universe.filter_stocks members 交集 + resolve_members; Research/Selection/Signal/Replay 服务注入 index repo(历史成分过滤,选股/回测共用) - tests/test_index_universe.py:快照历史成分(成分变更不入早期结果)、幂等、 空快照期空集、index_code 过滤下 as_of 一致性;全量 pytest 通过 --- backend/app/api/deps.py | 21 +++- .../application/services/replay_service.py | 16 ++- .../application/services/selection_service.py | 10 +- .../application/services/signal_service.py | 16 ++- backend/app/domain/entities/index.py | 21 ++++ backend/app/domain/entities/research.py | 3 + backend/app/domain/repositories/index.py | 18 +++ .../f5e0d1c2b3a4_index_weight_table.py | 46 +++++++ .../persistence/sqlalchemy/models/index.py | 31 +++++ .../sqlalchemy/repositories/index_impl.py | 70 +++++++++++ backend/app/quant/service.py | 9 +- backend/app/quant/universe.py | 15 ++- backend/tests/test_index_universe.py | 119 ++++++++++++++++++ 13 files changed, 379 insertions(+), 16 deletions(-) create mode 100644 backend/app/domain/entities/index.py create mode 100644 backend/app/domain/repositories/index.py create mode 100644 backend/app/infrastructure/persistence/migrations/versions/f5e0d1c2b3a4_index_weight_table.py create mode 100644 backend/app/infrastructure/persistence/sqlalchemy/models/index.py create mode 100644 backend/app/infrastructure/persistence/sqlalchemy/repositories/index_impl.py create mode 100644 backend/tests/test_index_universe.py diff --git a/backend/app/api/deps.py b/backend/app/api/deps.py index d5d22cd..e5ab7e0 100644 --- a/backend/app/api/deps.py +++ b/backend/app/api/deps.py @@ -16,6 +16,7 @@ from app.application.services.selection_service import SelectionService from app.application.services.signal_service import SignalService from app.domain.repositories.composite import CompositeRepository from app.domain.repositories.factor import FactorRepository +from app.domain.repositories.index import IndexConstituentRepository from app.domain.repositories.jobs import ExperimentRepository, JobRepository from app.domain.repositories.market import ( AdjustFactorRepository, @@ -32,6 +33,9 @@ from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl impor from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import ( SqlAlchemyFactorRepository, ) +from app.infrastructure.persistence.sqlalchemy.repositories.index_impl import ( + SqlAlchemyIndexConstituentRepository, +) from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import ( SqlAlchemyAdjustFactorRepository, SqlAlchemyDailyBarRepository, @@ -78,6 +82,10 @@ def _chart_service_factory( return ChartService(stock_repo, daily_repo, adj_repo) +def _index_repo_factory(session: DbSession) -> IndexConstituentRepository: + return SqlAlchemyIndexConstituentRepository(session) + + def _engine_factory() -> QuantEngine: return LocalEngine() @@ -86,31 +94,34 @@ def _service_factory( stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)], daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)], engine: Annotated[QuantEngine, Depends(_engine_factory)], + index_repo: Annotated[IndexConstituentRepository, Depends(_index_repo_factory)], ) -> ResearchService: - return ResearchService(stock_repo, daily_repo, engine) + return ResearchService(stock_repo, daily_repo, engine, index_repo) def _replay_service_factory( stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)], daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)], + index_repo: Annotated[IndexConstituentRepository, Depends(_index_repo_factory)], ) -> ReplayService: - return ReplayService(stock_repo, daily_repo) + return ReplayService(stock_repo, daily_repo, index_repo) def _signal_service_factory( stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)], daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)], + index_repo: Annotated[IndexConstituentRepository, Depends(_index_repo_factory)], ) -> SignalService: - return SignalService(stock_repo, daily_repo) + return SignalService(stock_repo, daily_repo, index_repo) def _selection_service_factory( stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)], daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)], financial_repo: Annotated[FinancialRepository, Depends(_financial_repo_factory)], + index_repo: Annotated[IndexConstituentRepository, Depends(_index_repo_factory)], ) -> SelectionService: - - return SelectionService(stock_repo, daily_repo, financial_repo) + return SelectionService(stock_repo, daily_repo, financial_repo, index_repo) def _selection_repo_factory(session: DbSession) -> SelectionRepository: diff --git a/backend/app/application/services/replay_service.py b/backend/app/application/services/replay_service.py index ab3c55d..fdc7f80 100644 --- a/backend/app/application/services/replay_service.py +++ b/backend/app/application/services/replay_service.py @@ -16,17 +16,24 @@ from app.domain.entities.selection import SelectionQuery from app.domain.entities.signal import SignalRules from app.domain.repositories.market import DailyBarRepository, StockRepository from app.quant.selection import factor_columns -from app.quant.service import filter_stocks, load_daily_df +from app.quant.service import load_daily_df from app.quant.signal import generate_signals +from app.quant.universe import filter_stocks, resolve_members MAX_SYMBOLS = 40 MAX_DAYS = 90 class ReplayService: - def __init__(self, stock_repo: StockRepository, daily_repo: DailyBarRepository) -> None: + def __init__( + self, + stock_repo: StockRepository, + daily_repo: DailyBarRepository, + index_repo=None, + ) -> None: self._stock_repo = stock_repo self._daily_repo = daily_repo + self._index_repo = index_repo def replay( self, @@ -41,7 +48,10 @@ class ReplayService: raise ValueError("Bar Replay 需要 universe.symbols 白名单(≤40 只),避免全市场长任务") if len(symbols) > MAX_SYMBOLS: raise ValueError(f"Bar Replay 白名单最多 {MAX_SYMBOLS} 只,当前 {len(symbols)}") - stocks = filter_stocks(self._stock_repo.list(), query.universe, as_of=start) + stocks = filter_stocks( + self._stock_repo.list(), query.universe, as_of=start, + members=resolve_members(self._index_repo, query.universe, start), + ) if not stocks: return ReplayResult(start=start, end=end, top_n=top_n) diff --git a/backend/app/application/services/selection_service.py b/backend/app/application/services/selection_service.py index 0478c32..915f6c5 100644 --- a/backend/app/application/services/selection_service.py +++ b/backend/app/application/services/selection_service.py @@ -28,7 +28,8 @@ from app.quant.selection import ( run_condition_selection, run_score_selection, ) -from app.quant.service import filter_stocks, load_daily_df +from app.quant.service import load_daily_df +from app.quant.universe import filter_stocks, resolve_members _FUNDAMENTAL_PREFIX = "fundamental." @@ -41,14 +42,19 @@ class SelectionService: stock_repo: StockRepository, daily_repo: DailyBarRepository, financial_repo: FinancialRepository | None = None, + index_repo=None, ) -> None: self._stock_repo = stock_repo self._daily_repo = daily_repo self._financial_repo = financial_repo + self._index_repo = index_repo def select(self, query: SelectionQuery) -> SelectionResult: as_of = query.as_of or date.today() - stocks = filter_stocks(self._stock_repo.list(), query.universe, as_of=as_of) + stocks = filter_stocks( + self._stock_repo.list(), query.universe, as_of=as_of, + members=resolve_members(self._index_repo, query.universe, as_of), + ) if not stocks: return self._run(query, pd.DataFrame(), stocks, as_of, financial={}) symbols = [s.symbol for s in stocks] diff --git a/backend/app/application/services/signal_service.py b/backend/app/application/services/signal_service.py index 4c89219..27045ed 100644 --- a/backend/app/application/services/signal_service.py +++ b/backend/app/application/services/signal_service.py @@ -13,18 +13,28 @@ from app.domain.entities.selection import SelectionQuery from app.domain.entities.signal import SignalResult, SignalRules from app.domain.repositories.market import DailyBarRepository, StockRepository from app.quant.selection import factor_columns -from app.quant.service import filter_stocks, load_daily_df +from app.quant.service import load_daily_df from app.quant.signal import generate_signals +from app.quant.universe import filter_stocks, resolve_members class SignalService: - def __init__(self, stock_repo: StockRepository, daily_repo: DailyBarRepository) -> None: + def __init__( + self, + stock_repo: StockRepository, + daily_repo: DailyBarRepository, + index_repo=None, + ) -> None: self._stock_repo = stock_repo self._daily_repo = daily_repo + self._index_repo = index_repo def signal(self, query: SelectionQuery, rules: SignalRules) -> SignalResult: as_of = query.as_of or date.today() - stocks = filter_stocks(self._stock_repo.list(), query.universe, as_of=as_of) + stocks = filter_stocks( + self._stock_repo.list(), query.universe, as_of=as_of, + members=resolve_members(self._index_repo, query.universe, as_of), + ) if not stocks: return generate_signals(pd.DataFrame(), query, rules, as_of) columns = sorted(factor_columns(query)) diff --git a/backend/app/domain/entities/index.py b/backend/app/domain/entities/index.py new file mode 100644 index 0000000..12f4210 --- /dev/null +++ b/backend/app/domain/entities/index.py @@ -0,0 +1,21 @@ +"""指数及其历史成分(v3 §9/§30:Survivorship-free Universe 的数据基础)。 + +index_weight:指数在某交易日的成分快照(来自指数权重表,每期含当时成分与权重)。 +历史成分语义:as_of 某日的成分 = 该日(<=as_of 最近一期)快照中的股票 —— +禁止用今天的成分回测过去(未来函数/幸存者偏差红线)。 +""" + +from __future__ import annotations + +from datetime import date +from decimal import Decimal + +from pydantic import BaseModel, Field + + +class IndexWeight(BaseModel): + index_code: str = Field(pattern=r"^\d{6}\.(SH|SZ|CSI|CI)$", description="如 000300.SH / 000905.SH") + index_name: str | None = None + trade_date: date # 该快照对应交易日(成分时点) + symbol: str = Field(pattern=r"^\d{6}\.(SH|SZ|BJ)$") + weight: Decimal | None = None diff --git a/backend/app/domain/entities/research.py b/backend/app/domain/entities/research.py index 40e167f..2d5bf53 100644 --- a/backend/app/domain/entities/research.py +++ b/backend/app/domain/entities/research.py @@ -26,6 +26,9 @@ class UniverseSpec(BaseModel): exclude_st: bool = True exclude_suspended: bool = True min_listing_days: int = Field(default=250, ge=0, description="上市至少 N 个自然日") + index_code: str | None = Field( + default=None, description="指数成分过滤(如 000300.SH):按 as_of 当日历史成分(v3 §9)" + ) symbols: list[str] = Field( default_factory=list, description="白名单(可选):非空时仅这些 symbol 参与选股/回测", diff --git a/backend/app/domain/repositories/index.py b/backend/app/domain/repositories/index.py new file mode 100644 index 0000000..04e24f0 --- /dev/null +++ b/backend/app/domain/repositories/index.py @@ -0,0 +1,18 @@ +"""指数成分 Repository Protocol(B1)。""" + +from __future__ import annotations + +from collections.abc import Sequence +from datetime import date +from typing import Protocol + +from app.domain.entities.index import IndexWeight + + +class IndexConstituentRepository(Protocol): + def upsert_many(self, rows: Sequence[IndexWeight]) -> int: ... + + def members_at(self, index_code: str, as_of: date) -> set[str]: + """as_of 当日成分:取 <=as_of 最近一期快照的股票集合(历史成分语义)。""" + + def latest_date(self, index_code: str) -> date | None: ... diff --git a/backend/app/infrastructure/persistence/migrations/versions/f5e0d1c2b3a4_index_weight_table.py b/backend/app/infrastructure/persistence/migrations/versions/f5e0d1c2b3a4_index_weight_table.py new file mode 100644 index 0000000..f80c18f --- /dev/null +++ b/backend/app/infrastructure/persistence/migrations/versions/f5e0d1c2b3a4_index_weight_table.py @@ -0,0 +1,46 @@ +"""index_weight 表(B1 指数历史成分) + +Revision ID: f5e0d1c2b3a4 +Revises: e1f2a3b4c5d6 +Create Date: 2026-09-09 + +""" + +from __future__ import annotations + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +revision: str = "f5e0d1c2b3a4" +down_revision: str | None = "e1f2a3b4c5d6" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.create_table( + "index_weight", + sa.Column( + "id", + sa.BigInteger().with_variant(sa.Integer(), "sqlite"), + autoincrement=True, + nullable=False, + ), + sa.Column("index_code", sa.String(length=12), nullable=False), + sa.Column("index_name", sa.String(length=64), nullable=True), + sa.Column("trade_date", sa.Date(), nullable=False), + sa.Column("symbol", sa.String(length=12), nullable=False), + sa.Column("weight", sa.Numeric(precision=10, scale=6), nullable=True), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint("index_code", "trade_date", "symbol", name="uq_idxw_code_date_sym"), + ) + op.create_index("ix_index_weight_index_code", "index_weight", ["index_code"]) + op.create_index("ix_index_weight_trade_date", "index_weight", ["trade_date"]) + + +def downgrade() -> None: + op.drop_index("ix_index_weight_trade_date", table_name="index_weight") + op.drop_index("ix_index_weight_index_code", table_name="index_weight") + op.drop_table("index_weight") diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/index.py b/backend/app/infrastructure/persistence/sqlalchemy/models/index.py new file mode 100644 index 0000000..960352f --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/index.py @@ -0,0 +1,31 @@ +"""指数历史成分表(B1)。 + +index_weight:指数成分快照(index_code, trade_date, symbol 唯一)。 +历史成分查询(members_at)取 <= as_of 最近一期快照 —— Survivorship-free Universe。 +""" + +from __future__ import annotations + +from datetime import date +from decimal import Decimal + +from sqlalchemy import BigInteger, Date, Integer, Numeric, String, UniqueConstraint +from sqlalchemy.orm import Mapped, mapped_column + +from app.infrastructure.persistence.sqlalchemy.base import Base + +PK_INT = BigInteger().with_variant(Integer, "sqlite") + + +class IndexWeightModel(Base): + __tablename__ = "index_weight" + __table_args__ = ( + UniqueConstraint("index_code", "trade_date", "symbol", name="uq_idxw_code_date_sym"), + ) + + id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True) + index_code: Mapped[str] = mapped_column(String(12), index=True) + index_name: Mapped[str | None] = mapped_column(String(64), nullable=True) + trade_date: Mapped[date] = mapped_column(Date, index=True) + symbol: Mapped[str] = mapped_column(String(12)) + weight: Mapped[Decimal | None] = mapped_column(Numeric(10, 6), nullable=True) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/index_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/index_impl.py new file mode 100644 index 0000000..6519fe3 --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/index_impl.py @@ -0,0 +1,70 @@ +"""指数成分 Repository 的 SQLAlchemy 实现(B1)。""" + +from __future__ import annotations + +from collections.abc import Sequence +from datetime import date + +from sqlalchemy import select +from sqlalchemy.orm import Session + +from app.domain.entities.index import IndexWeight +from app.infrastructure.persistence.sqlalchemy.models.index import IndexWeightModel + + +class SqlAlchemyIndexConstituentRepository: + def __init__(self, session: Session) -> None: + self._session = session + + def upsert_many(self, rows: Sequence[IndexWeight]) -> int: + if not rows: + return 0 + existing = { + (r.index_code, r.trade_date, r.symbol): r + for r in self._session.scalars( + select(IndexWeightModel).where( + IndexWeightModel.index_code.in_({x.index_code for x in rows}) + ) + ) + } + changed = 0 + for row in rows: + key = (row.index_code, row.trade_date, row.symbol) + model = existing.get(key) + if model is None: + self._session.add(IndexWeightModel(**row.model_dump())) + changed += 1 + else: + for k, v in row.model_dump(exclude={"index_code", "trade_date", "symbol"}).items(): + setattr(model, k, v) + changed += 1 + return changed + + def latest_date(self, index_code: str) -> date | None: + return self._session.scalar( + select(IndexWeightModel.trade_date) + .where(IndexWeightModel.index_code == index_code) + .order_by(IndexWeightModel.trade_date.desc()) + .limit(1) + ) + + def members_at(self, index_code: str, as_of: date) -> set[str]: + """<= as_of 最近一期快照的成分(历史成分;无快照 → 空集)。""" + latest = self._session.scalar( + select(IndexWeightModel.trade_date) + .where( + IndexWeightModel.index_code == index_code, + IndexWeightModel.trade_date <= as_of, + ) + .order_by(IndexWeightModel.trade_date.desc()) + .limit(1) + ) + if latest is None: + return set() + rows = self._session.scalars( + select(IndexWeightModel.symbol).where( + IndexWeightModel.index_code == index_code, + IndexWeightModel.trade_date == latest, + ) + ).all() + return set(rows) diff --git a/backend/app/quant/service.py b/backend/app/quant/service.py index 40a0964..16d425b 100644 --- a/backend/app/quant/service.py +++ b/backend/app/quant/service.py @@ -25,7 +25,7 @@ from app.domain.repositories.market import ( StockRepository, ) from app.quant.engine import QuantEngine -from app.quant.universe import filter_stocks # noqa: F401 —— 选股/回测共用范围过滤 +from app.quant.universe import filter_stocks, resolve_members # noqa: F401 —— 范围过滤 # 流式路径每攒多少行落一个 DataFrame 分片(控制 concat 峰值) _FRAME_CHUNK_ROWS = 50_000 @@ -110,10 +110,12 @@ class ResearchService: stock_repo: StockRepository, daily_repo: DailyBarRepository, engine: QuantEngine, + index_repo=None, ) -> None: self._stock_repo = stock_repo self._daily_repo = daily_repo self._engine = engine + self._index_repo = index_repo def run_factor_test(self, spec: ResearchSpec, horizon_days: int = 21) -> FactorTestReport: if spec.type != "factor_test": @@ -133,7 +135,10 @@ class ResearchService: start, end = spec.period # 回测前预留因子 warmup(lookback≤120 交易日,取 300 自然日余量) data_start = start - timedelta(days=300) - stocks = filter_stocks(self._stock_repo.list(), spec.universe, as_of=start) + stocks = filter_stocks( + self._stock_repo.list(), spec.universe, as_of=start, + members=resolve_members(self._index_repo, spec.universe, start), + ) # 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV) required = self._engine.required_columns(spec) return load_daily_df( diff --git a/backend/app/quant/universe.py b/backend/app/quant/universe.py index 9a31ca5..5bbf18e 100644 --- a/backend/app/quant/universe.py +++ b/backend/app/quant/universe.py @@ -17,17 +17,30 @@ from app.domain.entities.market import Stock from app.domain.entities.research import UniverseSpec +def resolve_members(index_repo, universe: UniverseSpec, as_of: date) -> set[str] | None: + """若 universe 指定指数成分 → 取 as_of 当日历史成分;否则 None(不过滤)。""" + if index_repo is None or not universe.index_code: + return None + return index_repo.members_at(universe.index_code, as_of) + + def filter_stocks( stocks: Sequence[Stock], universe: UniverseSpec, as_of: date, + members: set[str] | None = None, ) -> list[Stock]: - """按股票池口径过滤,返回 as_of 时点应纳入的股票列表。""" + """按股票池口径过滤,返回 as_of 时点应纳入的股票列表。 + + members:指数历史成分集合(resolve_members 结果);提供时取交集。 + """ symbols = set(universe.symbols) if universe.symbols else None out: list[Stock] = [] for s in stocks: if symbols is not None and s.symbol not in symbols: continue + if members is not None and s.symbol not in members: + continue if s.delist_date is not None and s.delist_date < as_of: continue if universe.exclude_st and s.name and "ST" in s.name.upper(): diff --git a/backend/tests/test_index_universe.py b/backend/tests/test_index_universe.py new file mode 100644 index 0000000..0f2080b --- /dev/null +++ b/backend/tests/test_index_universe.py @@ -0,0 +1,119 @@ +"""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)