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 通过
This commit is contained in:
@@ -22,6 +22,7 @@ from app.domain.repositories.market import (
|
||||
)
|
||||
from app.domain.repositories.selection import SelectionRepository
|
||||
from app.domain.repositories.signal import SignalRepository
|
||||
from app.domain.repositories.strategy import StrategyRepository
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl import (
|
||||
SqlAlchemyCompositeRepository,
|
||||
)
|
||||
@@ -39,6 +40,9 @@ from app.infrastructure.persistence.sqlalchemy.repositories.selection_impl impor
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.signal_impl import (
|
||||
SqlAlchemySignalRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
|
||||
SqlAlchemyStrategyRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.session import get_session
|
||||
from app.quant.engine import LocalEngine, QuantEngine
|
||||
from app.quant.service import ResearchService
|
||||
@@ -102,6 +106,10 @@ def _signal_repo_factory(session: DbSession) -> SignalRepository:
|
||||
return SqlAlchemySignalRepository(session)
|
||||
|
||||
|
||||
def _strategy_repo_factory(session: DbSession) -> StrategyRepository:
|
||||
return SqlAlchemyStrategyRepository(session)
|
||||
|
||||
|
||||
StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)]
|
||||
DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)]
|
||||
EngineDep = Annotated[QuantEngine, Depends(_engine_factory)]
|
||||
@@ -112,6 +120,7 @@ FactorRepoDep = Annotated[FactorRepository, Depends(_factor_repo_factory)]
|
||||
CompositeRepoDep = Annotated[CompositeRepository, Depends(_composite_repo_factory)]
|
||||
SignalRepoDep = Annotated[SignalRepository, Depends(_signal_repo_factory)]
|
||||
SignalServiceDep = Annotated[SignalService, Depends(_signal_service_factory)]
|
||||
StrategyRepoDep = Annotated[StrategyRepository, Depends(_strategy_repo_factory)]
|
||||
|
||||
|
||||
def _job_repo_factory(session: DbSession):
|
||||
|
||||
@@ -19,6 +19,7 @@ from app.api import (
|
||||
selections,
|
||||
signals,
|
||||
stocks,
|
||||
strategies,
|
||||
)
|
||||
|
||||
api_router = APIRouter()
|
||||
@@ -29,6 +30,7 @@ api_router.include_router(composites.router)
|
||||
api_router.include_router(research.router)
|
||||
api_router.include_router(selections.router)
|
||||
api_router.include_router(signals.router)
|
||||
api_router.include_router(strategies.router)
|
||||
api_router.include_router(jobs.router)
|
||||
api_router.include_router(experiments.router)
|
||||
api_router.include_router(agent.router)
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
"""策略 API(M8.3):/api/strategies CRUD + 展开为 ResearchSpec。
|
||||
|
||||
POST /api/strategies 保存策略(name 唯一)
|
||||
GET /api/strategies 列表
|
||||
GET /api/strategies/{id}
|
||||
DELETE /api/strategies/{id}
|
||||
POST /api/strategies/{id}/expand body: {period:[start,end], initial_capital?} → ResearchSpec
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.api.deps import DbSession, StrategyRepoDep
|
||||
from app.application.services.job_executor import new_id
|
||||
from app.domain.entities.research import ResearchSpec
|
||||
from app.domain.entities.strategy import StrategyDefinition
|
||||
|
||||
router = APIRouter(prefix="/strategies", tags=["strategies"])
|
||||
|
||||
|
||||
class ExpandRequest(BaseModel):
|
||||
period: tuple[date, date]
|
||||
initial_capital: float = Field(default=1_000_000.0, gt=0)
|
||||
|
||||
|
||||
@router.post("", response_model=StrategyDefinition, summary="保存策略")
|
||||
def create_strategy(
|
||||
definition: StrategyDefinition,
|
||||
strategy_repo: StrategyRepoDep,
|
||||
session: DbSession,
|
||||
) -> StrategyDefinition:
|
||||
try:
|
||||
saved = strategy_repo.save(definition.model_copy(update={"id": new_id("STG")}))
|
||||
session.commit()
|
||||
return saved
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.get("", response_model=list[StrategyDefinition], summary="策略列表")
|
||||
def list_strategies(strategy_repo: StrategyRepoDep) -> list[StrategyDefinition]:
|
||||
return strategy_repo.list()
|
||||
|
||||
|
||||
@router.get("/{strategy_id}", response_model=StrategyDefinition, summary="读取策略")
|
||||
def get_strategy(strategy_id: str, strategy_repo: StrategyRepoDep) -> StrategyDefinition:
|
||||
row = strategy_repo.get(strategy_id)
|
||||
if row is None:
|
||||
raise HTTPException(status_code=404, detail=f"策略 {strategy_id} 不存在")
|
||||
return row
|
||||
|
||||
|
||||
@router.delete("/{strategy_id}", summary="删除策略")
|
||||
def delete_strategy(
|
||||
strategy_id: str,
|
||||
strategy_repo: StrategyRepoDep,
|
||||
session: DbSession,
|
||||
) -> dict:
|
||||
if not strategy_repo.delete(strategy_id):
|
||||
raise HTTPException(status_code=404, detail=f"策略 {strategy_id} 不存在")
|
||||
session.commit()
|
||||
return {"deleted": strategy_id}
|
||||
|
||||
|
||||
@router.post("/{strategy_id}/expand", response_model=ResearchSpec, summary="展开为研究 Spec")
|
||||
def expand_strategy(
|
||||
strategy_id: str,
|
||||
req: ExpandRequest,
|
||||
strategy_repo: StrategyRepoDep,
|
||||
) -> ResearchSpec:
|
||||
row = strategy_repo.get(strategy_id)
|
||||
if row is None:
|
||||
raise HTTPException(status_code=404, detail=f"策略 {strategy_id} 不存在")
|
||||
if req.period[0] >= req.period[1]:
|
||||
raise HTTPException(status_code=400, detail="period 必须满足 start < end")
|
||||
return row.to_research_spec(period=req.period, initial_capital=req.initial_capital)
|
||||
@@ -0,0 +1,53 @@
|
||||
"""策略领域实体(M8.3,v2 §17/§5.3)。
|
||||
|
||||
Strategy = 完整策略定义(universe + factors + selection + rebalance + costs +
|
||||
portfolio,除回测区间 period 外),保存为命名资产;回测时补 period 展开为
|
||||
ResearchSpec(v2 §18 Research Specification 为统一契约,策略是其持久化形态)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.domain.entities.research import (
|
||||
CostSpec,
|
||||
FactorSpec,
|
||||
PortfolioSpec,
|
||||
SelectionSpec,
|
||||
UniverseSpec,
|
||||
)
|
||||
|
||||
|
||||
class StrategyDefinition(BaseModel):
|
||||
id: str = ""
|
||||
name: str = Field(min_length=1, max_length=64)
|
||||
description: str = ""
|
||||
spec_type: str = Field(default="backtest", pattern="^(backtest|factor_test)$")
|
||||
universe: UniverseSpec = UniverseSpec()
|
||||
price_adjustment: str = Field(default="none", pattern="^(none|qfq)$")
|
||||
factors: list[FactorSpec] = Field(min_length=1)
|
||||
selection: SelectionSpec = SelectionSpec()
|
||||
rebalance: str = Field(default="monthly", pattern="^(weekly|monthly)$")
|
||||
costs: CostSpec = CostSpec()
|
||||
portfolio: PortfolioSpec = PortfolioSpec()
|
||||
version: str = "1"
|
||||
created_at: datetime | None = None
|
||||
|
||||
def to_research_spec(self, period: tuple[date, date], initial_capital: float | None = None):
|
||||
"""补全回测区间/资金后展开为标准 ResearchSpec(可在 Job/回测执行)。"""
|
||||
from app.domain.entities.research import ResearchSpec
|
||||
|
||||
return ResearchSpec(
|
||||
type=self.spec_type,
|
||||
universe=self.universe,
|
||||
price_adjustment=self.price_adjustment,
|
||||
factors=self.factors,
|
||||
selection=self.selection,
|
||||
rebalance=self.rebalance,
|
||||
period=period,
|
||||
costs=self.costs,
|
||||
portfolio=self.portfolio,
|
||||
initial_capital=initial_capital if initial_capital else 1_000_000.0,
|
||||
)
|
||||
@@ -0,0 +1,20 @@
|
||||
"""策略 Repository Protocol(M8.3)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.strategy import StrategyDefinition
|
||||
|
||||
|
||||
class StrategyRepository(Protocol):
|
||||
def save(self, definition: StrategyDefinition) -> StrategyDefinition:
|
||||
"""新建(name 冲突抛 ValueError)。"""
|
||||
|
||||
def get(self, strategy_id: str) -> StrategyDefinition | None: ...
|
||||
|
||||
def get_by_name(self, name: str) -> StrategyDefinition | None: ...
|
||||
|
||||
def list(self) -> list[StrategyDefinition]: ...
|
||||
|
||||
def delete(self, strategy_id: str) -> bool: ...
|
||||
+38
@@ -0,0 +1,38 @@
|
||||
"""strategy 表(M8.3 策略持久化)
|
||||
|
||||
Revision ID: e1f2a3b4c5d6
|
||||
Revises: d8e0b2f3c4d5
|
||||
Create Date: 2026-09-09
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "e1f2a3b4c5d6"
|
||||
down_revision: str | None = "d8e0b2f3c4d5"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"strategy",
|
||||
sa.Column("id", sa.String(length=32), nullable=False),
|
||||
sa.Column("name", sa.String(length=64), nullable=False),
|
||||
sa.Column("description", sa.String(length=300), nullable=False),
|
||||
sa.Column("spec_type", sa.String(length=16), nullable=False),
|
||||
sa.Column("config_json", sa.Text(), nullable=False),
|
||||
sa.Column("version", sa.String(length=16), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("name", name="uq_strategy_name"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("strategy")
|
||||
@@ -30,3 +30,6 @@ from app.infrastructure.persistence.sqlalchemy.models.signal import ( # noqa: F
|
||||
SignalEventModel,
|
||||
SignalSnapshotModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.strategy import ( # noqa: F401
|
||||
StrategyModel,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
"""strategy 表(M8.3)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
|
||||
class StrategyModel(Base):
|
||||
__tablename__ = "strategy"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
name: Mapped[str] = mapped_column(String(64), unique=True)
|
||||
description: Mapped[str] = mapped_column(String(300), default="")
|
||||
spec_type: Mapped[str] = mapped_column(String(16), default="backtest")
|
||||
config_json: Mapped[str] = mapped_column(Text)
|
||||
version: Mapped[str] = mapped_column(String(16), default="1")
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
@@ -0,0 +1,86 @@
|
||||
"""策略 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,
|
||||
)
|
||||
@@ -0,0 +1,123 @@
|
||||
"""M8.3 策略测试:repo CRUD + /api/strategies + expand 为 ResearchSpec。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
import pytest
|
||||
from app.api import deps
|
||||
from app.domain.entities.research import SelectionSpec
|
||||
from app.domain.entities.strategy import StrategyDefinition
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
|
||||
SqlAlchemyStrategyRepository,
|
||||
)
|
||||
from app.main import app
|
||||
from fastapi.testclient import TestClient
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
|
||||
def _st() -> StrategyDefinition:
|
||||
return StrategyDefinition(
|
||||
name="质量成长动量",
|
||||
description="ROE+动量(演示)",
|
||||
factors=[{"name": "momentum_60", "weight": 1.0}],
|
||||
selection=SelectionSpec(top_n=10),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def session(tmp_path):
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'st.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
Session = sessionmaker(bind=engine, expire_on_commit=False)
|
||||
with Session() as s:
|
||||
yield s
|
||||
|
||||
|
||||
class TestStrategyRepository:
|
||||
def test_save_get_list_delete(self, session) -> None:
|
||||
repo = SqlAlchemyStrategyRepository(session)
|
||||
repo.save(_st().model_copy(update={"id": "STG-T1"}))
|
||||
session.commit()
|
||||
got = repo.get("STG-T1")
|
||||
assert got is not None and got.name == "质量成长动量"
|
||||
assert got.selection.top_n == 10
|
||||
assert len(repo.list()) == 1
|
||||
assert repo.get_by_name("质量成长动量") is not None
|
||||
assert repo.delete("STG-T1") is True
|
||||
session.commit()
|
||||
assert repo.get("STG-T1") is None
|
||||
|
||||
def test_duplicate_name(self, session) -> None:
|
||||
repo = SqlAlchemyStrategyRepository(session)
|
||||
repo.save(_st().model_copy(update={"id": "STG-A"}))
|
||||
session.commit()
|
||||
with pytest.raises(ValueError):
|
||||
repo.save(_st().model_copy(update={"id": "STG-B"}))
|
||||
|
||||
def test_expand_to_research_spec(self, session) -> None:
|
||||
st = _st().model_copy(update={"id": "STG-E"})
|
||||
spec = st.to_research_spec((date(2024, 1, 1), date(2024, 6, 1)))
|
||||
assert spec.type == "backtest"
|
||||
assert spec.period == (date(2024, 1, 1), date(2024, 6, 1))
|
||||
assert spec.factors[0].name == "momentum_60"
|
||||
assert spec.price_adjustment == "none"
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def client(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)
|
||||
|
||||
def _session_override():
|
||||
with Session() as s:
|
||||
yield s
|
||||
|
||||
app.dependency_overrides[deps.get_session] = _session_override
|
||||
with TestClient(app) as c:
|
||||
yield c
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
|
||||
class TestStrategiesApi:
|
||||
def test_crud_and_expand(self, client) -> None:
|
||||
body = {
|
||||
"name": "演示策略",
|
||||
"description": "动量",
|
||||
"factors": [{"name": "momentum_60", "weight": 1}],
|
||||
"selection": {"top_n": 10},
|
||||
}
|
||||
created = client.post("/api/strategies", json=body)
|
||||
assert created.status_code == 200
|
||||
sid = created.json()["id"]
|
||||
assert sid.startswith("STG-")
|
||||
|
||||
assert len(client.get("/api/strategies").json()) == 1
|
||||
detail = client.get(f"/api/strategies/{sid}").json()
|
||||
assert detail["name"] == "演示策略"
|
||||
|
||||
resp = client.post(
|
||||
f"/api/strategies/{sid}/expand",
|
||||
json={"period": ["2024-01-01", "2024-06-01"]},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
spec = resp.json()
|
||||
assert spec["type"] == "backtest"
|
||||
assert spec["factors"][0]["name"] == "momentum_60"
|
||||
|
||||
assert client.delete(f"/api/strategies/{sid}").status_code == 200
|
||||
assert client.get(f"/api/strategies/{sid}").status_code == 404
|
||||
|
||||
def test_duplicate_and_bad_period(self, client) -> None:
|
||||
body = {"name": "A", "factors": [{"name": "momentum_60", "weight": 1}]}
|
||||
assert client.post("/api/strategies", json=body).status_code == 200
|
||||
assert client.post("/api/strategies", json=body).status_code == 400
|
||||
sid = client.get("/api/strategies").json()[0]["id"]
|
||||
bad = client.post(
|
||||
f"/api/strategies/{sid}/expand",
|
||||
json={"period": ["2024-06-01", "2024-01-01"]},
|
||||
)
|
||||
assert bad.status_code == 400
|
||||
Reference in New Issue
Block a user