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.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))
+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_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 参与选股/回测",
+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, 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(
+14 -1
View File
@@ -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():
+119
View File
@@ -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)