按用户目标把原来「一个策略 = 全套参数」拆开(已确认的设计决策):
- 公共配置 GlobalConfig(全局唯一):佣金/印花税/滑点/最低佣金/复权口径/基准
- 选股策略 SelectionStrategy(原 StrategyDefinition 改名):只剩股票池+因子+条件,
不再持有 selection/rebalance/costs/portfolio/区间/资金
- 回测组合 BacktestCombo:引用若干选股策略 + 回测时才定的参数
(起始资金、持仓数 N、持仓天数区间 [Tmin,Tmax]、调仓时机 日/周/月、区间)
引擎(app/quant/combo_engine.py,新增):
- 多策略打分 = 并集 + Borda 秩和(各策略 1/名次 求和;不假设不同策略分值可比,
能容纳各策略股票池不同);抽出纯函数 borda_combine 便于单测
- 持仓天数区间 [Tmin,Tmax]:Tmax **每个交易日**强制了结(安全阀,月频下也不超期);
Tmin 仅在调仓日保护(掉出 TopN 但未满 Tmin 暂留,防频繁换手);调仓日为增量调仓
(只卖超期/掉队且满 Tmin 的,从 TopN 补买至 N 只,不主动减持以尊重 Tmin)
- 调仓时机 daily/weekly/monthly(local_engine.rebalance_dates 新增日频分支)
- 产出与旧 runner 同构的 BacktestResult,前端可视化无需改动;config_snapshot 固化
ComboRunSpec(组合+当时各策略定义+当时成本/复权)保证可复现
数据层:
- 新表 global_config(默认行:万三/hfq/最低佣金5元)、backtest_combo
- 迁移 b4c5d6e7f8a9:建两表 + 把存量 strategy.config_json 的回测参数键剥掉、
spec_type 收敛为 selection(已在真实 MariaDB 验证:STG-16BFBF08 清洗后只剩
universe/factors/conditions)
- 仓储 SqlAlchemyGlobalConfigRepository / SqlAlchemyComboRepository + Protocol
API:
- /api/config GET/PUT;/api/combos CRUD + /{id}/run + /run(kind=combo 异步 Job)
- job_executor 新增 combo 分支:取齐策略+读公共配置→ComboService.run,归档 kind
记 backtest(结果结构相同)
- /api/strategies 切到 SelectionStrategy,移除已废弃的 /{id}/expand
- strategy_doc.describe_strategy 支持 SelectionStrategy(只讲「怎么选」,如实声明
资金/持仓/调仓/成本/区间在回测组合里定)
旧的 ResearchSpec + /api/backtests 保留(因子测试与既有契约自检仍用),
作为底层 escape hatch;用户产品路径改为回测组合。
测试:新增 test_combo_engine(6)/test_combo_service(3)/test_combo_api(5),
改写 test_strategies/test_strategy_doc 适配新模型。全量 403 passed(原 388)。
133 lines
5.2 KiB
Python
133 lines
5.2 KiB
Python
"""选股策略(SelectionStrategy)测试:repo CRUD + /api/strategies + 说明生成。
|
|
|
|
2026-09 重构后:策略库只存「选股条件组合」(股票池 + 因子 + 条件),
|
|
不再持有 selection/rebalance/costs/portfolio/区间 —— 那些移到回测组合与公共配置。
|
|
旧的 `to_research_spec` / `/expand` 已移除。
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import pytest
|
|
from app.api import deps
|
|
from app.domain.entities.strategy import SelectionStrategy
|
|
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() -> SelectionStrategy:
|
|
return SelectionStrategy(
|
|
name="质量成长动量",
|
|
description="ROE+动量(演示)",
|
|
factors=[{"name": "momentum_60", "weight": 1.0}],
|
|
conditions=[{"field": "dv_ratio", "op": "lte", "value": 30}],
|
|
)
|
|
|
|
|
|
@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.factors[0].name == "momentum_60"
|
|
assert len(got.conditions) == 1 and got.conditions[0].field == "dv_ratio"
|
|
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_no_backtest_params_in_entity(self) -> None:
|
|
"""选股策略实体不应再有回测执行参数字段(重构的核心约束)。"""
|
|
st = _st()
|
|
dumped = st.model_dump()
|
|
for forbidden in (
|
|
"selection", "rebalance", "costs", "portfolio",
|
|
"initial_capital", "period", "price_adjustment",
|
|
"selection_interval_months", "rebalance_interval_months",
|
|
):
|
|
assert forbidden not in dumped, f"选股策略不应含回测参数字段 {forbidden}"
|
|
|
|
|
|
@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_describe(self, client) -> None:
|
|
body = {
|
|
"name": "演示策略",
|
|
"description": "动量",
|
|
"factors": [{"name": "momentum_60", "weight": 1}],
|
|
"conditions": [{"field": "dv_ratio", "op": "lte", "value": 30}],
|
|
}
|
|
created = client.post("/api/strategies", json=body)
|
|
assert created.status_code == 200
|
|
sid = created.json()["id"]
|
|
assert sid.startswith("STG-")
|
|
# 回读不含回测参数字段
|
|
detail = client.get(f"/api/strategies/{sid}").json()
|
|
assert detail["name"] == "演示策略"
|
|
assert "selection" not in detail and "costs" not in detail
|
|
|
|
assert len(client.get("/api/strategies").json()) == 1
|
|
|
|
# 说明生成:选股策略走专用路径,不假装知道回测参数
|
|
doc = client.get(f"/api/strategies/{sid}/describe").json()
|
|
assert "选股策略" in doc["summary"]
|
|
assert any("回测组合" in w for w in doc["warnings"])
|
|
|
|
assert client.delete(f"/api/strategies/{sid}").status_code == 200
|
|
assert client.get(f"/api/strategies/{sid}").status_code == 404
|
|
|
|
def test_duplicate_name_400(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
|
|
|
|
def test_expand_endpoint_removed(self, client) -> None:
|
|
"""/expand 已随重构移除(回测改由「回测组合」驱动,不再从单策略展开 ResearchSpec)。"""
|
|
body = {"name": "B", "factors": [{"name": "momentum_60", "weight": 1}]}
|
|
sid = client.post("/api/strategies", json=body).json()["id"]
|
|
resp = client.post(
|
|
f"/api/strategies/{sid}/expand",
|
|
json={"period": ["2024-01-01", "2024-06-01"]},
|
|
)
|
|
assert resp.status_code in (404, 405)
|