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:
Simon
2026-09-30 21:43:28 +08:00
parent 50a1030afa
commit 40bd603b44
25 changed files with 2250 additions and 174 deletions
+109
View File
@@ -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"]
+195
View File
@@ -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)),
)
+148
View File
@@ -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())
+42 -33
View File
@@ -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)
+15 -17
View File
@@ -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"] == "改后的说明"