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.application.services.signal_service import SignalService
|
||||||
from app.domain.repositories.composite import CompositeRepository
|
from app.domain.repositories.composite import CompositeRepository
|
||||||
from app.domain.repositories.factor import FactorRepository
|
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.jobs import ExperimentRepository, JobRepository
|
||||||
from app.domain.repositories.market import (
|
from app.domain.repositories.market import (
|
||||||
AdjustFactorRepository,
|
AdjustFactorRepository,
|
||||||
@@ -32,6 +33,9 @@ from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl impor
|
|||||||
from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import (
|
from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import (
|
||||||
SqlAlchemyFactorRepository,
|
SqlAlchemyFactorRepository,
|
||||||
)
|
)
|
||||||
|
from app.infrastructure.persistence.sqlalchemy.repositories.index_impl import (
|
||||||
|
SqlAlchemyIndexConstituentRepository,
|
||||||
|
)
|
||||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||||
SqlAlchemyAdjustFactorRepository,
|
SqlAlchemyAdjustFactorRepository,
|
||||||
SqlAlchemyDailyBarRepository,
|
SqlAlchemyDailyBarRepository,
|
||||||
@@ -78,6 +82,10 @@ def _chart_service_factory(
|
|||||||
return ChartService(stock_repo, daily_repo, adj_repo)
|
return ChartService(stock_repo, daily_repo, adj_repo)
|
||||||
|
|
||||||
|
|
||||||
|
def _index_repo_factory(session: DbSession) -> IndexConstituentRepository:
|
||||||
|
return SqlAlchemyIndexConstituentRepository(session)
|
||||||
|
|
||||||
|
|
||||||
def _engine_factory() -> QuantEngine:
|
def _engine_factory() -> QuantEngine:
|
||||||
return LocalEngine()
|
return LocalEngine()
|
||||||
|
|
||||||
@@ -86,31 +94,34 @@ def _service_factory(
|
|||||||
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
||||||
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
||||||
engine: Annotated[QuantEngine, Depends(_engine_factory)],
|
engine: Annotated[QuantEngine, Depends(_engine_factory)],
|
||||||
|
index_repo: Annotated[IndexConstituentRepository, Depends(_index_repo_factory)],
|
||||||
) -> ResearchService:
|
) -> ResearchService:
|
||||||
return ResearchService(stock_repo, daily_repo, engine)
|
return ResearchService(stock_repo, daily_repo, engine, index_repo)
|
||||||
|
|
||||||
|
|
||||||
def _replay_service_factory(
|
def _replay_service_factory(
|
||||||
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
||||||
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
||||||
|
index_repo: Annotated[IndexConstituentRepository, Depends(_index_repo_factory)],
|
||||||
) -> ReplayService:
|
) -> ReplayService:
|
||||||
return ReplayService(stock_repo, daily_repo)
|
return ReplayService(stock_repo, daily_repo, index_repo)
|
||||||
|
|
||||||
|
|
||||||
def _signal_service_factory(
|
def _signal_service_factory(
|
||||||
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
||||||
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
||||||
|
index_repo: Annotated[IndexConstituentRepository, Depends(_index_repo_factory)],
|
||||||
) -> SignalService:
|
) -> SignalService:
|
||||||
return SignalService(stock_repo, daily_repo)
|
return SignalService(stock_repo, daily_repo, index_repo)
|
||||||
|
|
||||||
|
|
||||||
def _selection_service_factory(
|
def _selection_service_factory(
|
||||||
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
||||||
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
||||||
financial_repo: Annotated[FinancialRepository, Depends(_financial_repo_factory)],
|
financial_repo: Annotated[FinancialRepository, Depends(_financial_repo_factory)],
|
||||||
|
index_repo: Annotated[IndexConstituentRepository, Depends(_index_repo_factory)],
|
||||||
) -> SelectionService:
|
) -> SelectionService:
|
||||||
|
return SelectionService(stock_repo, daily_repo, financial_repo, index_repo)
|
||||||
return SelectionService(stock_repo, daily_repo, financial_repo)
|
|
||||||
|
|
||||||
|
|
||||||
def _selection_repo_factory(session: DbSession) -> SelectionRepository:
|
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.entities.signal import SignalRules
|
||||||
from app.domain.repositories.market import DailyBarRepository, StockRepository
|
from app.domain.repositories.market import DailyBarRepository, StockRepository
|
||||||
from app.quant.selection import factor_columns
|
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.signal import generate_signals
|
||||||
|
from app.quant.universe import filter_stocks, resolve_members
|
||||||
|
|
||||||
MAX_SYMBOLS = 40
|
MAX_SYMBOLS = 40
|
||||||
MAX_DAYS = 90
|
MAX_DAYS = 90
|
||||||
|
|
||||||
|
|
||||||
class ReplayService:
|
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._stock_repo = stock_repo
|
||||||
self._daily_repo = daily_repo
|
self._daily_repo = daily_repo
|
||||||
|
self._index_repo = index_repo
|
||||||
|
|
||||||
def replay(
|
def replay(
|
||||||
self,
|
self,
|
||||||
@@ -41,7 +48,10 @@ class ReplayService:
|
|||||||
raise ValueError("Bar Replay 需要 universe.symbols 白名单(≤40 只),避免全市场长任务")
|
raise ValueError("Bar Replay 需要 universe.symbols 白名单(≤40 只),避免全市场长任务")
|
||||||
if len(symbols) > MAX_SYMBOLS:
|
if len(symbols) > MAX_SYMBOLS:
|
||||||
raise ValueError(f"Bar Replay 白名单最多 {MAX_SYMBOLS} 只,当前 {len(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:
|
if not stocks:
|
||||||
return ReplayResult(start=start, end=end, top_n=top_n)
|
return ReplayResult(start=start, end=end, top_n=top_n)
|
||||||
|
|
||||||
|
|||||||
@@ -28,7 +28,8 @@ from app.quant.selection import (
|
|||||||
run_condition_selection,
|
run_condition_selection,
|
||||||
run_score_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."
|
_FUNDAMENTAL_PREFIX = "fundamental."
|
||||||
|
|
||||||
@@ -41,14 +42,19 @@ class SelectionService:
|
|||||||
stock_repo: StockRepository,
|
stock_repo: StockRepository,
|
||||||
daily_repo: DailyBarRepository,
|
daily_repo: DailyBarRepository,
|
||||||
financial_repo: FinancialRepository | None = None,
|
financial_repo: FinancialRepository | None = None,
|
||||||
|
index_repo=None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._stock_repo = stock_repo
|
self._stock_repo = stock_repo
|
||||||
self._daily_repo = daily_repo
|
self._daily_repo = daily_repo
|
||||||
self._financial_repo = financial_repo
|
self._financial_repo = financial_repo
|
||||||
|
self._index_repo = index_repo
|
||||||
|
|
||||||
def select(self, query: SelectionQuery) -> SelectionResult:
|
def select(self, query: SelectionQuery) -> SelectionResult:
|
||||||
as_of = query.as_of or date.today()
|
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:
|
if not stocks:
|
||||||
return self._run(query, pd.DataFrame(), stocks, as_of, financial={})
|
return self._run(query, pd.DataFrame(), stocks, as_of, financial={})
|
||||||
symbols = [s.symbol for s in stocks]
|
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.entities.signal import SignalResult, SignalRules
|
||||||
from app.domain.repositories.market import DailyBarRepository, StockRepository
|
from app.domain.repositories.market import DailyBarRepository, StockRepository
|
||||||
from app.quant.selection import factor_columns
|
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.signal import generate_signals
|
||||||
|
from app.quant.universe import filter_stocks, resolve_members
|
||||||
|
|
||||||
|
|
||||||
class SignalService:
|
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._stock_repo = stock_repo
|
||||||
self._daily_repo = daily_repo
|
self._daily_repo = daily_repo
|
||||||
|
self._index_repo = index_repo
|
||||||
|
|
||||||
def signal(self, query: SelectionQuery, rules: SignalRules) -> SignalResult:
|
def signal(self, query: SelectionQuery, rules: SignalRules) -> SignalResult:
|
||||||
as_of = query.as_of or date.today()
|
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:
|
if not stocks:
|
||||||
return generate_signals(pd.DataFrame(), query, rules, as_of)
|
return generate_signals(pd.DataFrame(), query, rules, as_of)
|
||||||
columns = sorted(factor_columns(query))
|
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_st: bool = True
|
||||||
exclude_suspended: bool = True
|
exclude_suspended: bool = True
|
||||||
min_listing_days: int = Field(default=250, ge=0, description="上市至少 N 个自然日")
|
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(
|
symbols: list[str] = Field(
|
||||||
default_factory=list,
|
default_factory=list,
|
||||||
description="白名单(可选):非空时仅这些 symbol 参与选股/回测",
|
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,
|
StockRepository,
|
||||||
)
|
)
|
||||||
from app.quant.engine import QuantEngine
|
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 峰值)
|
# 流式路径每攒多少行落一个 DataFrame 分片(控制 concat 峰值)
|
||||||
_FRAME_CHUNK_ROWS = 50_000
|
_FRAME_CHUNK_ROWS = 50_000
|
||||||
@@ -110,10 +110,12 @@ class ResearchService:
|
|||||||
stock_repo: StockRepository,
|
stock_repo: StockRepository,
|
||||||
daily_repo: DailyBarRepository,
|
daily_repo: DailyBarRepository,
|
||||||
engine: QuantEngine,
|
engine: QuantEngine,
|
||||||
|
index_repo=None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._stock_repo = stock_repo
|
self._stock_repo = stock_repo
|
||||||
self._daily_repo = daily_repo
|
self._daily_repo = daily_repo
|
||||||
self._engine = engine
|
self._engine = engine
|
||||||
|
self._index_repo = index_repo
|
||||||
|
|
||||||
def run_factor_test(self, spec: ResearchSpec, horizon_days: int = 21) -> FactorTestReport:
|
def run_factor_test(self, spec: ResearchSpec, horizon_days: int = 21) -> FactorTestReport:
|
||||||
if spec.type != "factor_test":
|
if spec.type != "factor_test":
|
||||||
@@ -133,7 +135,10 @@ class ResearchService:
|
|||||||
start, end = spec.period
|
start, end = spec.period
|
||||||
# 回测前预留因子 warmup(lookback≤120 交易日,取 300 自然日余量)
|
# 回测前预留因子 warmup(lookback≤120 交易日,取 300 自然日余量)
|
||||||
data_start = start - timedelta(days=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)
|
# 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV)
|
||||||
required = self._engine.required_columns(spec)
|
required = self._engine.required_columns(spec)
|
||||||
return load_daily_df(
|
return load_daily_df(
|
||||||
|
|||||||
@@ -17,17 +17,30 @@ from app.domain.entities.market import Stock
|
|||||||
from app.domain.entities.research import UniverseSpec
|
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(
|
def filter_stocks(
|
||||||
stocks: Sequence[Stock],
|
stocks: Sequence[Stock],
|
||||||
universe: UniverseSpec,
|
universe: UniverseSpec,
|
||||||
as_of: date,
|
as_of: date,
|
||||||
|
members: set[str] | None = None,
|
||||||
) -> list[Stock]:
|
) -> list[Stock]:
|
||||||
"""按股票池口径过滤,返回 as_of 时点应纳入的股票列表。"""
|
"""按股票池口径过滤,返回 as_of 时点应纳入的股票列表。
|
||||||
|
|
||||||
|
members:指数历史成分集合(resolve_members 结果);提供时取交集。
|
||||||
|
"""
|
||||||
symbols = set(universe.symbols) if universe.symbols else None
|
symbols = set(universe.symbols) if universe.symbols else None
|
||||||
out: list[Stock] = []
|
out: list[Stock] = []
|
||||||
for s in stocks:
|
for s in stocks:
|
||||||
if symbols is not None and s.symbol not in symbols:
|
if symbols is not None and s.symbol not in symbols:
|
||||||
continue
|
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:
|
if s.delist_date is not None and s.delist_date < as_of:
|
||||||
continue
|
continue
|
||||||
if universe.exclude_st and s.name and "ST" in s.name.upper():
|
if universe.exclude_st and s.name and "ST" in s.name.upper():
|
||||||
|
|||||||
@@ -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)
|
||||||
Reference in New Issue
Block a user