- 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 通过
124 lines
4.4 KiB
Python
124 lines
4.4 KiB
Python
"""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
|