"""选股策略 Repository 的 SQLAlchemy 实现。 config_json 只存选股相关字段(universe/factors/conditions);读取时重建 SelectionStrategy。 2026-09 重构:策略库不再持有回测执行参数(selection/rebalance/costs/portfolio/区间), 旧行若残留这些键,读出时由 Pydantic 的 extra 忽略策略丢弃(见 _to_entity)。 """ from __future__ import annotations import json from datetime import datetime from sqlalchemy import select from sqlalchemy.orm import Session from app.domain.entities.strategy import SelectionStrategy from app.infrastructure.persistence.sqlalchemy.models.strategy import StrategyModel class SqlAlchemyStrategyRepository: def __init__(self, session: Session) -> None: self._session = session def save(self, definition: SelectionStrategy) -> SelectionStrategy: if not definition.id: raise ValueError("需要 id(由调用方生成)") dup = self._session.scalar( select(StrategyModel).where(StrategyModel.name == definition.name).limit(1) ) if dup is not None and dup.id != definition.id: raise ValueError(f"策略名已存在:{definition.name}") now = definition.created_at or datetime.now() row = self._session.get(StrategyModel, definition.id) config_json = json.dumps(definition.model_dump(exclude={"id", "created_at"}), ensure_ascii=False) if row is None: self._session.add( StrategyModel( id=definition.id, name=definition.name, description=definition.description, spec_type=definition.spec_type, config_json=config_json, version=definition.version, created_at=now, ) ) else: row.name = definition.name row.description = definition.description row.spec_type = definition.spec_type row.config_json = config_json row.version = definition.version self._session.flush() return definition def get(self, strategy_id: str) -> SelectionStrategy | None: row = self._session.get(StrategyModel, strategy_id) return _to_entity(row) if row else None def get_by_name(self, name: str) -> SelectionStrategy | None: row = self._session.scalar( select(StrategyModel).where(StrategyModel.name == name).limit(1) ) return _to_entity(row) if row else None def list(self) -> list[SelectionStrategy]: rows = self._session.scalars( select(StrategyModel).order_by(StrategyModel.name) ).all() return [_to_entity(r) for r in rows] def delete(self, strategy_id: str) -> bool: row = self._session.get(StrategyModel, strategy_id) if row is None: return False self._session.delete(row) return True # 旧 strategy.config_json 可能残留的回测执行参数字段(重构前写入)—— 读出时丢弃, # 因为 SelectionStrategy 不再承载它们(已迁到回测组合 / 公共配置)。 _LEGACY_BACKTEST_KEYS = frozenset({ "selection", "rebalance", "costs", "portfolio", "price_adjustment", "selection_interval_months", "rebalance_interval_months", }) def _to_entity(row: StrategyModel) -> SelectionStrategy: data = json.loads(row.config_json) # 列字段由 DB 行回填,避免与 config_json 重复。 # description 必须一并回填:它是列字段(String(300)),save() 会写入, # 但这里若只从 config_json 里 pop 掉却不回填,读回的策略说明会恒为空串 # (读写不对称:保存的说明看不到,策略库/编辑页都拿不到)。 for key in ("name", "version", "description", "spec_type"): data.pop(key, None) # 丢弃旧行的回测参数字段(Pydantic 默认 forbid extra 会因这些键报错) for key in _LEGACY_BACKTEST_KEYS: data.pop(key, None) return SelectionStrategy( id=row.id, name=row.name, version=row.version, description=row.description, created_at=row.created_at, **data, )