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:
+16
-5
@@ -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))
|
||||
|
||||
@@ -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
|
||||
@@ -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 参与选股/回测",
|
||||
|
||||
@@ -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: ...
|
||||
+46
@@ -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)
|
||||
@@ -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(
|
||||
|
||||
@@ -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():
|
||||
|
||||
Reference in New Issue
Block a user