Files
qlib/backend/app/infrastructure/persistence/sqlalchemy/repositories/strategy_impl.py
T
Simon 9d25d466e5 feat(strategy): M8.3 策略模型 + /api/strategies(命名配置资产,可展开为 ResearchSpec)
- StrategyDefinition:universe/factors/selection/rebalance/costs/portfolio +
  price_adjustment(除 period 外完整策略定义);to_research_spec(period) 展开为标准 Spec
- strategy 表(migration e1f2a3b4c5d6,MySQL 已应用;name 唯一)+ StrategyRepository
- /api/strategies:POST/GET/DELETE + POST /{id}/expand(period+initial_capital → ResearchSpec)
- tests/test_strategies.py(repo CRUD/同名/expand、API CRUD/400/404);全量 pytest 通过
2026-09-09 00:38:09 +08:00

87 lines
3.1 KiB
Python

"""策略 Repository 的 SQLAlchemy 实现(M8.3)。
config 以 JSON 存(StrategyDefinition.model_dump);读取时重建实体。
"""
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 StrategyDefinition
from app.infrastructure.persistence.sqlalchemy.models.strategy import StrategyModel
class SqlAlchemyStrategyRepository:
def __init__(self, session: Session) -> None:
self._session = session
def save(self, definition: StrategyDefinition) -> StrategyDefinition:
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) -> StrategyDefinition | None:
row = self._session.get(StrategyModel, strategy_id)
return _to_entity(row) if row else None
def get_by_name(self, name: str) -> StrategyDefinition | 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[StrategyDefinition]:
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
def _to_entity(row: StrategyModel) -> StrategyDefinition:
data = json.loads(row.config_json)
# 列字段由 DB 行回填,避免与 config_json 重复
for key in ("name", "version", "description", "spec_type"):
data.pop(key, None)
return StrategyDefinition(
id=row.id, name=row.name, version=row.version,
created_at=row.created_at, **data,
)