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:
@@ -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