Files
qlib/backend/app/infrastructure/persistence/sqlalchemy/repositories/index_impl.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

71 lines
2.4 KiB
Python

"""指数成分 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)