"""回测组合 + 公共配置 Repository 的 SQLAlchemy 实现(2026-09 重构)。""" from __future__ import annotations import json from datetime import datetime from sqlalchemy import select from sqlalchemy.orm import Session from app.domain.entities.combo import BacktestCombo, GlobalConfig from app.infrastructure.persistence.sqlalchemy.models.combo import ( BacktestComboModel, GlobalConfigModel, ) DEFAULT_CONFIG_ID = "default" class SqlAlchemyGlobalConfigRepository: def __init__(self, session: Session) -> None: self._session = session def get(self) -> GlobalConfig: row = self._session.get(GlobalConfigModel, DEFAULT_CONFIG_ID) if row is None: return GlobalConfig() # 未配置过 → 返回默认值(不写库,由调用方决定是否 save) return GlobalConfig( id=row.id, commission_rate=float(row.commission_rate), stamp_tax_rate=float(row.stamp_tax_rate), slippage_rate=float(row.slippage_rate), min_commission=float(row.min_commission), price_adjustment=row.price_adjustment, benchmark=row.benchmark, updated_at=row.updated_at, ) def save(self, config: GlobalConfig) -> GlobalConfig: row = self._session.get(GlobalConfigModel, DEFAULT_CONFIG_ID) if row is None: row = GlobalConfigModel(id=DEFAULT_CONFIG_ID) self._session.add(row) row.commission_rate = config.commission_rate row.stamp_tax_rate = config.stamp_tax_rate row.slippage_rate = config.slippage_rate row.min_commission = config.min_commission row.price_adjustment = config.price_adjustment row.benchmark = config.benchmark row.updated_at = datetime.now() self._session.flush() return config.model_copy(update={"id": DEFAULT_CONFIG_ID, "updated_at": row.updated_at}) class SqlAlchemyComboRepository: def __init__(self, session: Session) -> None: self._session = session def save(self, combo: BacktestCombo) -> BacktestCombo: if not combo.id: raise ValueError("需要 id(由调用方生成)") dup = self._session.scalar( select(BacktestComboModel).where(BacktestComboModel.name == combo.name).limit(1) ) if dup is not None and dup.id != combo.id: raise ValueError(f"回测组合名已存在:{combo.name}") now = combo.created_at or datetime.now() row = self._session.get(BacktestComboModel, combo.id) strategy_ids_json = json.dumps(combo.strategy_ids, ensure_ascii=False) if row is None: self._session.add( BacktestComboModel( id=combo.id, name=combo.name, description=combo.description, strategy_ids_json=strategy_ids_json, initial_capital=combo.initial_capital, hold_count=combo.hold_count, hold_min_days=combo.hold_min_days, hold_max_days=combo.hold_max_days, rebalance_freq=combo.rebalance_freq, start_date=combo.period[0], end_date=combo.period[1], version=combo.version, created_at=now, ) ) else: row.name = combo.name row.description = combo.description row.strategy_ids_json = strategy_ids_json row.initial_capital = combo.initial_capital row.hold_count = combo.hold_count row.hold_min_days = combo.hold_min_days row.hold_max_days = combo.hold_max_days row.rebalance_freq = combo.rebalance_freq row.start_date = combo.period[0] row.end_date = combo.period[1] row.version = combo.version self._session.flush() return combo def get(self, combo_id: str) -> BacktestCombo | None: row = self._session.get(BacktestComboModel, combo_id) return _to_combo(row) if row else None def list(self) -> list[BacktestCombo]: rows = self._session.scalars( select(BacktestComboModel).order_by(BacktestComboModel.created_at.desc()) ).all() return [_to_combo(r) for r in rows] def delete(self, combo_id: str) -> bool: row = self._session.get(BacktestComboModel, combo_id) if row is None: return False self._session.delete(row) return True def _to_combo(row: BacktestComboModel) -> BacktestCombo: return BacktestCombo( id=row.id, name=row.name, description=row.description, strategy_ids=json.loads(row.strategy_ids_json), initial_capital=float(row.initial_capital), hold_count=row.hold_count, hold_min_days=row.hold_min_days, hold_max_days=row.hold_max_days, rebalance_freq=row.rebalance_freq, period=(row.start_date, row.end_date), version=row.version, created_at=row.created_at, )