Files
qlib/backend/tests/test_index_universe.py
T
Simon 9cc4bfccac feat(universe): B1-1 指数历史成分(index_weight)+ Universe 按 as_of 成分过滤
- 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 通过
2026-09-09 07:27:13 +08:00

120 lines
5.2 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.
"""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)