feat(backend): 策略库重构为「选股策略 + 公共配置 + 回测组合」三件套
按用户目标把原来「一个策略 = 全套参数」拆开(已确认的设计决策):
- 公共配置 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)。
This commit is contained in:
@@ -0,0 +1,109 @@
|
||||
"""公共配置 + 回测组合 API 测试(TestClient + 内存 SQLite)。
|
||||
|
||||
覆盖:
|
||||
- /api/config GET 返回默认值、PUT 持久化;
|
||||
- /api/combos CRUD(name 唯一、原地更新保留 created_at、删除);
|
||||
- /api/combos/{id}/run 在策略缺失时提前 400(不等到后台才失败)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from app.api import deps
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.main import app
|
||||
from fastapi.testclient import TestClient
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def client(tmp_path) -> TestClient:
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'combo_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 TestConfigApi:
|
||||
def test_get_returns_defaults(self, client: TestClient) -> None:
|
||||
cfg = client.get("/api/config").json()
|
||||
assert cfg["id"] == "default"
|
||||
assert cfg["price_adjustment"] == "hfq"
|
||||
assert cfg["min_commission"] == 5.0
|
||||
|
||||
def test_put_persists(self, client: TestClient) -> None:
|
||||
body = {
|
||||
"commission_rate": 0.00025, "stamp_tax_rate": 0.0005,
|
||||
"slippage_rate": 0.0008, "min_commission": 3.0,
|
||||
"price_adjustment": "qfq", "benchmark": "000905.SH",
|
||||
}
|
||||
resp = client.put("/api/config", json=body)
|
||||
assert resp.status_code == 200
|
||||
got = client.get("/api/config").json()
|
||||
assert got["commission_rate"] == pytest.approx(0.00025)
|
||||
assert got["price_adjustment"] == "qfq"
|
||||
assert got["benchmark"] == "000905.SH"
|
||||
|
||||
|
||||
class TestCombosApi:
|
||||
def _make_strategy(self, client: TestClient, name: str) -> str:
|
||||
resp = client.post(
|
||||
"/api/strategies",
|
||||
json={"name": name, "factors": [{"name": "dividend_yield", "weight": 1}]},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
return resp.json()["id"]
|
||||
|
||||
def test_crud(self, client: TestClient) -> None:
|
||||
sid = self._make_strategy(client, "高股息")
|
||||
body = {
|
||||
"name": "组合A", "strategy_ids": [sid],
|
||||
"initial_capital": 500000, "hold_count": 10,
|
||||
"hold_min_days": 5, "hold_max_days": 30,
|
||||
"rebalance_freq": "weekly", "period": ["2024-01-01", "2024-06-01"],
|
||||
}
|
||||
created = client.post("/api/combos", json=body)
|
||||
assert created.status_code == 200
|
||||
cid = created.json()["id"]
|
||||
assert cid.startswith("CMB-")
|
||||
assert created.json()["hold_max_days"] == 30
|
||||
|
||||
assert len(client.get("/api/combos").json()) == 1
|
||||
detail = client.get(f"/api/combos/{cid}").json()
|
||||
assert detail["strategy_ids"] == [sid]
|
||||
assert detail["rebalance_freq"] == "weekly"
|
||||
|
||||
# 原地更新保留 id 与 created_at
|
||||
created_at = detail["created_at"]
|
||||
upd = client.put(f"/api/combos/{cid}", json={**body, "name": "组合A", "hold_count": 15})
|
||||
assert upd.status_code == 200
|
||||
assert upd.json()["id"] == cid
|
||||
assert upd.json()["created_at"] == created_at
|
||||
assert upd.json()["hold_count"] == 15
|
||||
|
||||
assert client.delete(f"/api/combos/{cid}").status_code == 200
|
||||
assert client.get(f"/api/combos/{cid}").status_code == 404
|
||||
|
||||
def test_duplicate_name_400(self, client: TestClient) -> None:
|
||||
sid = self._make_strategy(client, "S")
|
||||
body = {"name": "重名", "strategy_ids": [sid], "hold_count": 5,
|
||||
"rebalance_freq": "monthly", "period": ["2024-01-01", "2024-02-01"]}
|
||||
assert client.post("/api/combos", json=body).status_code == 200
|
||||
assert client.post("/api/combos", json=body).status_code == 400
|
||||
|
||||
def test_run_missing_strategy_400(self, client: TestClient) -> None:
|
||||
"""引用不存在的策略 → 提交时即 400,而非等后台 Job 才失败。"""
|
||||
body = {"name": "缺策略组合", "strategy_ids": ["STG-NOT-EXIST"], "hold_count": 5,
|
||||
"rebalance_freq": "monthly", "period": ["2024-01-01", "2024-02-01"]}
|
||||
resp = client.post("/api/combos/run", json=body)
|
||||
assert resp.status_code == 400
|
||||
assert "不存在" in resp.json()["detail"]
|
||||
@@ -0,0 +1,195 @@
|
||||
"""组合回测引擎单测(合成数据,确定性,不依赖数据库/因子注册表)。
|
||||
|
||||
验证三件用户确认的语义:
|
||||
1. Borda 秩和打分:两策略排名不同 → 综合排序可手算预测。
|
||||
2. 持仓天数区间 [Tmin, Tmax]:超 Tmax 强制了结;未满 Tmin 即使掉出 TopN 也暂留。
|
||||
3. 调仓时机 daily/weekly/monthly 产生不同的调仓次数。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, timedelta
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from app.domain.entities.combo import BacktestCombo
|
||||
from app.domain.entities.research import CostSpec
|
||||
from app.quant.combo_engine import HoldingBandRunner, borda_combine
|
||||
|
||||
|
||||
def _business_days(start: date, n: int) -> list[date]:
|
||||
"""生成 n 个连续工作日(跳过周末),用作合成行情索引。"""
|
||||
out: list[date] = []
|
||||
d = start
|
||||
while len(out) < n:
|
||||
if d.weekday() < 5:
|
||||
out.append(d)
|
||||
d += timedelta(days=1)
|
||||
return out
|
||||
|
||||
|
||||
def _flat_close(symbols: list[str], days: list[date], price: float = 100.0) -> pd.DataFrame:
|
||||
"""所有股票恒定价格的面板(收益为 0,便于隔离「选股/调仓」逻辑)。"""
|
||||
idx = pd.to_datetime(days)
|
||||
return pd.DataFrame(price, index=idx, columns=symbols)
|
||||
|
||||
|
||||
# ---------- 1. Borda 秩和 ----------
|
||||
|
||||
|
||||
def test_borda_combine_hand_computed():
|
||||
"""两策略排名不同,综合分 = Σ(1/名次),可手算。"""
|
||||
day = pd.Timestamp("2024-01-02")
|
||||
# 策略1:A > B > C;策略2:C > A > B
|
||||
p1 = pd.DataFrame({"A": [3.0], "B": [2.0], "C": [1.0]}, index=[day])
|
||||
p2 = pd.DataFrame({"A": [2.0], "B": [1.0], "C": [3.0]}, index=[day])
|
||||
combined = borda_combine([p1, p2]).loc[day]
|
||||
# A: 1/1 + 1/2 = 1.5;C: 1/3 + 1/1 = 1.333;B: 1/2 + 1/3 = 0.833
|
||||
assert combined["A"] == pytest.approx(1.5)
|
||||
assert combined["C"] == pytest.approx(1.0 / 3 + 1.0)
|
||||
assert combined["B"] == pytest.approx(1.0 / 2 + 1.0 / 3)
|
||||
order = combined.sort_values(ascending=False).index.tolist()
|
||||
assert order == ["A", "C", "B"] # 并集后统一排序:A、C 进 Top2,B 落选
|
||||
|
||||
|
||||
def test_borda_missing_symbol_contributes_zero():
|
||||
"""某策略面板里没有某股票(NaN)→ 该策略对它贡献 0,但不影响其它策略的贡献。"""
|
||||
day = pd.Timestamp("2024-01-02")
|
||||
p1 = pd.DataFrame({"A": [3.0], "B": [2.0]}, index=[day]) # 策略1 只有 A、B
|
||||
p2 = pd.DataFrame({"A": [1.0], "C": [2.0]}, index=[day]) # 策略2 只有 A、C
|
||||
combined = borda_combine([p1, p2]).loc[day]
|
||||
assert combined["A"] == pytest.approx(1.0 + 1.0 / 2) # 两策略都覆盖 A
|
||||
assert combined["B"] == pytest.approx(1.0 / 2) # 只被策略1 覆盖
|
||||
assert combined["C"] == pytest.approx(1.0) # 只被策略2 覆盖(在其面板里排第 1)
|
||||
|
||||
|
||||
# ---------- 2. 持仓天数区间 ----------
|
||||
|
||||
|
||||
def _make_combo(**overrides) -> BacktestCombo:
|
||||
base = dict(
|
||||
name="t", strategy_ids=["S1"], initial_capital=1_000_000.0,
|
||||
hold_count=1, hold_min_days=0, hold_max_days=None,
|
||||
rebalance_freq="daily", period=(date(2024, 1, 2), date(2024, 1, 31)),
|
||||
)
|
||||
base.update(overrides)
|
||||
return BacktestCombo(**base)
|
||||
|
||||
|
||||
def test_tmax_force_exit_respected():
|
||||
"""恒价 + N=1 + 永远选 A + Tmax=5 + 日频:A 持有超过 5 天即被强制卖出再买回,
|
||||
任何一笔交易的持有天数都不应明显超过 Tmax。"""
|
||||
symbols = ["A", "B", "C"]
|
||||
days = _business_days(date(2024, 1, 2), 30)
|
||||
close = _flat_close(symbols, days)
|
||||
# A 永远最高分 → 永远 Top1
|
||||
score = pd.DataFrame(
|
||||
{"A": [3.0] * len(days), "B": [2.0] * len(days), "C": [1.0] * len(days)},
|
||||
index=pd.to_datetime(days),
|
||||
)
|
||||
combo = _make_combo(hold_count=1, hold_min_days=0, hold_max_days=5, rebalance_freq="daily")
|
||||
runner = HoldingBandRunner(
|
||||
combo=combo, costs=CostSpec(min_commission=0.0), score=score, close=close,
|
||||
)
|
||||
result = runner.run()
|
||||
|
||||
assert result.trades, "应产生交易"
|
||||
# 交易日索引:用引擎同一口径(交易日)验证「任何一笔持仓都不超过 Tmax 个交易日」
|
||||
tday_pos = {pd.Timestamp(d): i for i, d in enumerate(close.index)}
|
||||
def trading_span(a, b):
|
||||
return tday_pos[pd.Timestamp(b)] - tday_pos[pd.Timestamp(a)]
|
||||
for t in result.trades:
|
||||
span = trading_span(t.entry_date, t.exit_date)
|
||||
# 卖出发生在「held > Tmax」的第一个交易日 → 跨度最多 Tmax+1 个交易日
|
||||
assert span <= 5 + 1, f"持仓跨 {span} 个交易日 > Tmax+1,Tmax 安全阀失效:{t}"
|
||||
# 确实反复「卖后再买」—— 证明 Tmax 在强制换手,而不是一直死拿
|
||||
buys = [a for a in result.signal_history if a.signal == "BUY" and a.filled]
|
||||
assert len(buys) >= 4, f"Tmax=5 在 30 个交易日内应触发多次重买,实际仅 {len(buys)} 次"
|
||||
|
||||
|
||||
def test_tmin_protects_against_churn():
|
||||
"""N=1,第 2 天起 B 变成最高分(A 掉出 Top1),但 Tmin=10 → A 在满 10 天前不被卖出。"""
|
||||
symbols = ["A", "B"]
|
||||
days = _business_days(date(2024, 1, 2), 20)
|
||||
close = _flat_close(symbols, days)
|
||||
# 第 0 天 A 最高;第 1 天起 B 最高
|
||||
a_scores = [3.0] + [1.0] * (len(days) - 1)
|
||||
b_scores = [1.0] + [3.0] * (len(days) - 1)
|
||||
score = pd.DataFrame({"A": a_scores, "B": b_scores}, index=pd.to_datetime(days))
|
||||
combo = _make_combo(hold_count=1, hold_min_days=10, hold_max_days=None, rebalance_freq="daily")
|
||||
runner = HoldingBandRunner(
|
||||
combo=combo, costs=CostSpec(min_commission=0.0), score=score, close=close,
|
||||
)
|
||||
result = runner.run()
|
||||
|
||||
# A 应在第 0 天买入
|
||||
a_buys = [a for a in result.signal_history if a.symbol == "A" and a.signal == "BUY" and a.filled]
|
||||
assert a_buys, "A 应在首日买入"
|
||||
a_sells = [t for t in result.trades if t.symbol == "A"]
|
||||
if a_sells:
|
||||
# 若最终卖出,持有天数必须 ≥ Tmin(不能在满 10 天前因掉出 TopN 被卖)
|
||||
for t in a_sells:
|
||||
assert (t.exit_date - t.entry_date).days >= 10, (
|
||||
f"A 仅持 {(t.exit_date - t.entry_date).days} 天就被卖,违反 Tmin=10 保护"
|
||||
)
|
||||
# 关键断言:前 9 个工作日内 A 不应被卖出(Tmin 保护生效)
|
||||
tday_pos = {pd.Timestamp(d): i for i, d in enumerate(close.index)}
|
||||
early_sells = [
|
||||
a for a in result.signal_history
|
||||
if a.symbol == "A" and a.signal == "SELL" and a.filled
|
||||
and tday_pos[pd.Timestamp(a.date)] - tday_pos[pd.Timestamp(days[0])] < 10
|
||||
]
|
||||
assert not early_sells, f"Tmin 保护失效:A 在 10 个交易日内被卖出 {early_sells}"
|
||||
|
||||
|
||||
# ---------- 3. 调仓时机 ----------
|
||||
|
||||
|
||||
def test_rebalance_freq_changes_cadence():
|
||||
"""同一份数据,daily 的调仓日数 > weekly > monthly(用 selection_history 的 distinct 日期数衡量)。"""
|
||||
symbols = ["A", "B", "C"]
|
||||
days = _business_days(date(2024, 1, 2), 60)
|
||||
close = _flat_close(symbols, days)
|
||||
score = pd.DataFrame(
|
||||
{"A": [3.0] * len(days), "B": [2.0] * len(days), "C": [1.0] * len(days)},
|
||||
index=pd.to_datetime(days),
|
||||
)
|
||||
|
||||
def run_with(freq: str) -> int:
|
||||
combo = _make_combo(
|
||||
hold_count=2, rebalance_freq=freq,
|
||||
period=(days[0], days[-1]),
|
||||
)
|
||||
r = HoldingBandRunner(
|
||||
combo=combo, costs=CostSpec(min_commission=0.0), score=score, close=close,
|
||||
).run()
|
||||
return len({p.date for p in r.selection_history})
|
||||
|
||||
daily_n, weekly_n, monthly_n = run_with("daily"), run_with("weekly"), run_with("monthly")
|
||||
assert daily_n > weekly_n > monthly_n, (
|
||||
f"调仓频次应 daily({daily_n}) > weekly({weekly_n}) > monthly({monthly_n})"
|
||||
)
|
||||
|
||||
|
||||
# ---------- 实体校验 ----------
|
||||
|
||||
|
||||
def test_combo_validates_hold_band_and_freq():
|
||||
with pytest.raises(ValueError):
|
||||
BacktestCombo(
|
||||
name="x", strategy_ids=["A"], hold_count=1,
|
||||
hold_min_days=20, hold_max_days=5, # Tmax < Tmin
|
||||
rebalance_freq="monthly", period=(date(2024, 1, 1), date(2024, 2, 1)),
|
||||
)
|
||||
with pytest.raises(ValueError):
|
||||
BacktestCombo(
|
||||
name="x", strategy_ids=["A"], hold_count=1,
|
||||
rebalance_freq="yearly", # 非法频率
|
||||
period=(date(2024, 1, 1), date(2024, 2, 1)),
|
||||
)
|
||||
with pytest.raises(ValueError):
|
||||
BacktestCombo(
|
||||
name="x", strategy_ids=["A", "A"], hold_count=1, # 重复策略
|
||||
rebalance_freq="monthly", period=(date(2024, 1, 1), date(2024, 2, 1)),
|
||||
)
|
||||
@@ -0,0 +1,148 @@
|
||||
"""回测组合服务集成测试(合成数据 + 内存 SQLite,不连真库)。
|
||||
|
||||
验证 ComboService.run 端到端:装配行情 → 多策略 Borda → 持仓区间 runner → BacktestResult。
|
||||
覆盖:
|
||||
- 两策略打分合并后选出并集 TopN;
|
||||
- 持仓天数区间 [Tmin, Tmax] 在真实数据装配路径下生效;
|
||||
- 公共配置的成本/复权被采用并写进 config_snapshot(可复现)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from app.application.services.combo_service import ComboService
|
||||
from app.domain.entities.combo import BacktestCombo, GlobalConfig
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import (
|
||||
DailyBasicModel,
|
||||
StockDailyModel,
|
||||
StockModel,
|
||||
)
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
# 5 只股票,股息率梯度:A 最高 … E 最低
|
||||
_SYMS = ["600000.SH", "600001.SH", "600002.SH", "600003.SH", "600004.SH"]
|
||||
_DV = {"600000.SH": 9.0, "600001.SH": 7.0, "600002.SH": 5.0, "600003.SH": 3.0, "600004.SH": 1.0}
|
||||
_START = date(2024, 1, 1)
|
||||
|
||||
|
||||
def _seed(session: Session) -> None:
|
||||
session.add_all(
|
||||
[
|
||||
StockModel(symbol=s, name=f"股票{s[:6]}", industry="银行", market="主板",
|
||||
area="深圳", list_date=date(2000, 1, 1), status="L")
|
||||
for s in _SYMS
|
||||
]
|
||||
)
|
||||
dates = pd.bdate_range(_START, periods=80)
|
||||
bars, basics = [], []
|
||||
for i, sym in enumerate(_SYMS):
|
||||
price = 10.0 + i
|
||||
for d in dates:
|
||||
price *= 1 + 0.0006 + 0.0002 * i
|
||||
bars.append(StockDailyModel(
|
||||
symbol=sym, trade_date=d.date(), source="tushare", adjust="none",
|
||||
open=Decimal(str(price)), high=Decimal(str(price)),
|
||||
low=Decimal(str(price)), close=Decimal(str(price)),
|
||||
volume=Decimal("1000000"), amount=Decimal(str(price * 1e6)),
|
||||
))
|
||||
basics.append(DailyBasicModel(
|
||||
symbol=sym, trade_date=d.date(), source="tushare",
|
||||
close=Decimal(str(price)), dv_ratio=Decimal(str(_DV[sym])),
|
||||
dv_ttm=Decimal(str(_DV[sym])), pe=Decimal("8"), pb=Decimal("1"),
|
||||
total_mv=Decimal("1e11"),
|
||||
))
|
||||
session.add_all(bars + basics)
|
||||
session.commit()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def service(tmp_path) -> ComboService:
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'combo.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
session = Session(engine)
|
||||
_seed(session)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyDailyBarRepository,
|
||||
SqlAlchemyDailyBasicRepository,
|
||||
SqlAlchemyStockRepository,
|
||||
)
|
||||
yield ComboService(
|
||||
SqlAlchemyStockRepository(session),
|
||||
SqlAlchemyDailyBarRepository(session),
|
||||
basic_repo=SqlAlchemyDailyBasicRepository(session),
|
||||
)
|
||||
session.close()
|
||||
|
||||
|
||||
def _strategy(sid: str, name: str, *, conditions=None) -> SelectionStrategy:
|
||||
return SelectionStrategy(
|
||||
id=sid, name=name,
|
||||
factors=[{"name": "dividend_yield", "weight": 1.0}],
|
||||
conditions=conditions or [],
|
||||
)
|
||||
|
||||
|
||||
def test_combo_run_produces_backtest_result_with_config_snapshot(service: ComboService) -> None:
|
||||
"""两策略(一个带 dv_ratio 条件、一个不带)→ 组合跑出 BacktestResult,
|
||||
且 config_snapshot 固化了当时的成本/复权与策略定义(可复现)。"""
|
||||
combo = BacktestCombo(
|
||||
id="CMB-T1", name="双策略高股息", strategy_ids=["S1", "S2"],
|
||||
initial_capital=1_000_000.0, hold_count=2, hold_min_days=0, hold_max_days=None,
|
||||
rebalance_freq="monthly", period=(date(2024, 1, 2), date(2024, 4, 15)),
|
||||
)
|
||||
strategies = [
|
||||
_strategy("S1", "纯高股息"),
|
||||
_strategy("S2", "高股息+过滤", conditions=[{"field": "dv_ratio", "op": "lte", "value": 30}]),
|
||||
]
|
||||
config = GlobalConfig(commission_rate=0.0003, stamp_tax_rate=0.0005,
|
||||
slippage_rate=0.001, min_commission=5.0, price_adjustment="hfq")
|
||||
|
||||
result = service.run(combo, strategies, config)
|
||||
|
||||
assert result.summary.initial_capital == 1_000_000.0
|
||||
assert result.equity_curve, "应产出净值曲线"
|
||||
assert result.trades or result.positions, "应有成交或持仓"
|
||||
# 可复现快照:含组合参数 + 两策略定义 + 当时成本/复权
|
||||
snap = result.config_snapshot
|
||||
assert snap["combo"]["hold_count"] == 2
|
||||
assert {s["id"] for s in snap["strategies"]} == {"S1", "S2"}
|
||||
assert snap["costs"]["min_commission"] == 5.0
|
||||
assert snap["price_adjustment"] == "hfq"
|
||||
|
||||
|
||||
def test_hold_max_days_limits_holding_in_real_run(service: ComboService) -> None:
|
||||
"""日频 + Tmax=8:任何一笔交易的持有交易日数不超过 Tmax+1。"""
|
||||
combo = BacktestCombo(
|
||||
id="CMB-T2", name="短持", strategy_ids=["S1"],
|
||||
hold_count=1, hold_min_days=0, hold_max_days=8,
|
||||
rebalance_freq="daily", period=(date(2024, 1, 2), date(2024, 4, 15)),
|
||||
)
|
||||
strategies = [_strategy("S1", "纯高股息")]
|
||||
config = GlobalConfig(min_commission=0.0, price_adjustment="none")
|
||||
result = service.run(combo, strategies, config)
|
||||
|
||||
assert result.trades, "日频短持应产生多次换手"
|
||||
# 用结果的 signal_history 重建交易日序列来按交易日计跨度
|
||||
trade_dates = sorted({pd.Timestamp(p.date) for p in result.equity_curve})
|
||||
pos = {d: i for i, d in enumerate(trade_dates)}
|
||||
for t in result.trades:
|
||||
span = pos[pd.Timestamp(t.exit_date)] - pos[pd.Timestamp(t.entry_date)]
|
||||
assert span <= 8 + 1, f"持仓跨 {span} 个交易日 > Tmax+1:{t}"
|
||||
|
||||
|
||||
def test_missing_strategy_raises_clear_error(service: ComboService) -> None:
|
||||
"""组合引用了 S-GONE,但只传入了 S-OTHER → 明确报出缺失的 id(不静默跳过)。"""
|
||||
combo = BacktestCombo(
|
||||
id="CMB-T3", name="缺策略", strategy_ids=["S-GONE"],
|
||||
hold_count=1, rebalance_freq="monthly", period=(date(2024, 1, 2), date(2024, 2, 1)),
|
||||
)
|
||||
other = _strategy("S-OTHER", "别的")
|
||||
with pytest.raises(ValueError, match="S-GONE"):
|
||||
service.run(combo, [other], GlobalConfig())
|
||||
@@ -1,13 +1,15 @@
|
||||
"""M8.3 策略测试:repo CRUD + /api/strategies + expand 为 ResearchSpec。"""
|
||||
"""选股策略(SelectionStrategy)测试:repo CRUD + /api/strategies + 说明生成。
|
||||
|
||||
2026-09 重构后:策略库只存「选股条件组合」(股票池 + 因子 + 条件),
|
||||
不再持有 selection/rebalance/costs/portfolio/区间 —— 那些移到回测组合与公共配置。
|
||||
旧的 `to_research_spec` / `/expand` 已移除。
|
||||
"""
|
||||
|
||||
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.domain.entities.strategy import SelectionStrategy
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
|
||||
SqlAlchemyStrategyRepository,
|
||||
@@ -18,12 +20,12 @@ from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
|
||||
def _st() -> StrategyDefinition:
|
||||
return StrategyDefinition(
|
||||
def _st() -> SelectionStrategy:
|
||||
return SelectionStrategy(
|
||||
name="质量成长动量",
|
||||
description="ROE+动量(演示)",
|
||||
factors=[{"name": "momentum_60", "weight": 1.0}],
|
||||
selection=SelectionSpec(top_n=10),
|
||||
conditions=[{"field": "dv_ratio", "op": "lte", "value": 30}],
|
||||
)
|
||||
|
||||
|
||||
@@ -43,7 +45,9 @@ class TestStrategyRepository:
|
||||
session.commit()
|
||||
got = repo.get("STG-T1")
|
||||
assert got is not None and got.name == "质量成长动量"
|
||||
assert got.selection.top_n == 10
|
||||
# 选股策略只保留选股相关字段
|
||||
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
|
||||
@@ -57,13 +61,16 @@ class TestStrategyRepository:
|
||||
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"
|
||||
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()
|
||||
@@ -83,41 +90,43 @@ def client(tmp_path):
|
||||
|
||||
|
||||
class TestStrategiesApi:
|
||||
def test_crud_and_expand(self, client) -> None:
|
||||
def test_crud_and_describe(self, client) -> None:
|
||||
body = {
|
||||
"name": "演示策略",
|
||||
"description": "动量",
|
||||
"factors": [{"name": "momentum_60", "weight": 1}],
|
||||
"selection": {"top_n": 10},
|
||||
"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-")
|
||||
|
||||
assert len(client.get("/api/strategies").json()) == 1
|
||||
# 回读不含回测参数字段
|
||||
detail = client.get(f"/api/strategies/{sid}").json()
|
||||
assert detail["name"] == "演示策略"
|
||||
assert "selection" not in detail and "costs" not in detail
|
||||
|
||||
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 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_and_bad_period(self, client) -> None:
|
||||
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
|
||||
sid = client.get("/api/strategies").json()[0]["id"]
|
||||
bad = client.post(
|
||||
|
||||
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-06-01", "2024-01-01"]},
|
||||
json={"period": ["2024-01-01", "2024-06-01"]},
|
||||
)
|
||||
assert bad.status_code == 400
|
||||
assert resp.status_code in (404, 405)
|
||||
|
||||
@@ -24,7 +24,7 @@ from app.domain.entities.research import (
|
||||
SelectionSpec,
|
||||
UniverseSpec,
|
||||
)
|
||||
from app.domain.entities.strategy import StrategyDefinition
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.main import app
|
||||
from app.quant.factors import FactorDef
|
||||
@@ -301,17 +301,19 @@ class TestStepsAndWarnings:
|
||||
|
||||
|
||||
class TestDefinitionInput:
|
||||
def test_strategy_definition_input_uses_placeholder_period(self) -> None:
|
||||
st = StrategyDefinition(
|
||||
def test_selection_strategy_describes_only_selection(self) -> None:
|
||||
"""选股策略说明只讲「怎么选」,不假装知道回测参数(重构后无 period/costs/selection)。"""
|
||||
st = SelectionStrategy(
|
||||
name="演示",
|
||||
factors=[FactorSpec(name="momentum_60", weight=1.0)],
|
||||
selection=SelectionSpec(top_n=10),
|
||||
conditions=[{"field": "dv_ratio", "op": "lte", "value": 30}],
|
||||
)
|
||||
doc = describe_strategy(st)
|
||||
assert doc.summary and doc.formula and doc.steps
|
||||
assert "选股策略" in doc.summary
|
||||
assert "momentum_60" in doc.formula
|
||||
assert "1900" not in doc.formula, "占位区间不得泄漏到展示文本"
|
||||
assert any("无回测区间" in w for w in doc.warnings)
|
||||
# 如实声明回测参数不在策略内
|
||||
assert any("回测组合" in w for w in doc.warnings)
|
||||
|
||||
def test_real_spec_has_no_placeholder_warning(self) -> None:
|
||||
doc = describe_strategy(_spec())
|
||||
@@ -325,13 +327,13 @@ class TestDefinitionInput:
|
||||
describe_strategy(spec)
|
||||
assert spec.model_dump() == before
|
||||
|
||||
st = StrategyDefinition(name="演示", factors=[FactorSpec(name="momentum_60")])
|
||||
st = SelectionStrategy(name="演示", factors=[FactorSpec(name="momentum_60")])
|
||||
before_st = st.model_dump()
|
||||
describe_strategy(st)
|
||||
assert st.model_dump() == before_st
|
||||
|
||||
def test_unsupported_input_raises_clear_type_error(self) -> None:
|
||||
with pytest.raises(TypeError, match="ResearchSpec 或 StrategyDefinition"):
|
||||
with pytest.raises(TypeError, match="ResearchSpec 或 SelectionStrategy"):
|
||||
describe_strategy({"factors": []}) # type: ignore[arg-type]
|
||||
|
||||
|
||||
@@ -391,7 +393,7 @@ class TestStrategyDocApi:
|
||||
assert resp.status_code == 200
|
||||
doc = resp.json()
|
||||
assert doc["summary"]
|
||||
assert any("无回测区间" in w for w in doc["warnings"])
|
||||
assert "选股策略" in doc["summary"]
|
||||
assert client.get("/api/strategies/STG-NOT-EXIST/describe").status_code == 404
|
||||
|
||||
def test_post_fills_empty_description(self, client: TestClient) -> None:
|
||||
@@ -428,8 +430,9 @@ class TestStrategyDocApi:
|
||||
def test_auto_description_fits_column_width(self, client: TestClient) -> None:
|
||||
"""自动说明必须落在 `StrategyModel.description = String(300)` 之内。
|
||||
|
||||
超长在 SQLite(测试库)不会报错、到 MySQL 严格模式会 Data too long,
|
||||
因此这里显式断言列宽;截断必须带省略号(显式标记,不静默改短)。
|
||||
重构后选股策略的说明只讲「怎么选」,天然简洁(不再拼回测公式),
|
||||
即使挂满全部因子也远低于列宽 —— 这里断言「一定放得下」即可;
|
||||
截断分支(超长带省略号)由 ResearchSpec 路径保留,选股策略触达不到。
|
||||
"""
|
||||
from app.api.strategies import _DESCRIPTION_MAX_CHARS
|
||||
from app.quant.factors import list_factors
|
||||
@@ -440,10 +443,8 @@ class TestStrategyDocApi:
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
desc = resp.json()["description"]
|
||||
assert len(desc) <= _DESCRIPTION_MAX_CHARS
|
||||
assert desc.endswith("…"), "超长被截断时必须显式带省略号"
|
||||
assert len(desc) <= _DESCRIPTION_MAX_CHARS, "选股策略说明也必须落在列宽内"
|
||||
|
||||
# 常规(单因子)说明远短于列宽:不应被截断
|
||||
normal = client.post(
|
||||
"/api/strategies",
|
||||
json={"name": "单因子策略", "factors": [{"name": "dividend_yield", "weight": 1}]},
|
||||
@@ -460,7 +461,6 @@ class TestStrategyUpdateApi:
|
||||
"name": name,
|
||||
"description": "初始说明",
|
||||
"factors": [{"name": "momentum_60", "weight": 1}],
|
||||
"selection": {"top_n": 10},
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
@@ -478,14 +478,12 @@ class TestStrategyUpdateApi:
|
||||
"name": "策略A",
|
||||
"description": "改后的说明",
|
||||
"factors": [{"name": "momentum_20", "weight": 2}],
|
||||
"selection": {"top_n": 5},
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["id"] == sid, "原地更新必须保持 id 不变(不新建)"
|
||||
assert body["created_at"] == created_at, "PUT 不得刷新创建时间"
|
||||
assert body["selection"]["top_n"] == 5
|
||||
assert body["factors"][0]["name"] == "momentum_20"
|
||||
assert body["description"] == "改后的说明"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user