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 通过
This commit is contained in:
Simon
2026-09-09 07:27:13 +08:00
parent 03fb463216
commit 9cc4bfccac
13 changed files with 379 additions and 16 deletions
+16 -5
View File
@@ -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:
@@ -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)
@@ -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]
@@ -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))
+21
View File
@@ -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
+3
View File
@@ -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 参与选股/回测",
+18
View File
@@ -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: ...
@@ -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")
@@ -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)
@@ -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)
+7 -2
View File
@@ -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(
+14 -1
View File
@@ -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():