Files
tick-stock-panel/backend/app/api/backtest.py
T
im47cn ec3309163b feat(optimizer): 参数网格搜索优化器 (并行回测 + 目标排序 + SSE + 前端面板) (#82)
* feat(optimizer): 参数网格搜索优化器核心 + PanelCache 线程安全

PR2a 第一部分 (后端核心, 无 API/前端):

app/backtest/optimizer.py:
- expand_param_grid: 校验(类型/范围/选项) + 笛卡尔积展开; 支持显式候选值列表
  与 {min,max,step} 范围两种写法; GRID_MAX_COMBINATIONS=2000 硬上限防爆炸。
- StrategyOptimizer.optimize: ThreadPoolExecutor 并行跑各参数组回测, 按目标指标
  排序返回最优 + 全排名。objective 统一转越大越好(min 类目标取负), None/inf/
  失败组沉底。支持进度回调 (done/total/best) 与 cancel_event。
- 可选目标含 PR1 新增的 sortino/mc_maxdd_*。

engine.py PanelCache 线程安全:
- get_or_compute 整体加锁 — 并行优化器对同一 symbols/日期跑几十组参数时, 首个
  线程 compute 面板、其余等待后命中缓存, 同一面板只 scan_parquet+compute_all 一次
  (关键性能: 参数不影响面板 key)。max_size 2->4, ttl 180->900s 适配长任务。

测试 24 例: 网格展开/校验(未知参数/越界/选项/空/组合爆炸) + 编排(排序/每组一次/
min方向/失败None沉底/取消/进度/非法目标)。用假 service 注入受控 stats 验证。

* fix(optimizer): 子代理审查修复 — 异常隔离/展示符号/浮点端点/kwargs 校验

两份子代理审查发现 4 个真实 bug:

[高] 单组异常拖垮整批: _run_one 未捕获 service.run 异常, fut.result() 会
re-raise 冲出 as_completed → 整个网格搜索崩溃、已完成结果全丢。加了并行后
这个概率显著上升。用 try/except 隔离, 该组记为 error 继续。

[高] min 方向 best_score 符号错误: 内部把 min 目标取负做排序键, 但直接把
取负值当 best_score 返回 → 用户看到 avg_holding_days=-3.0 (负天数)。根因是
排序键与展示值混用。改为分离: 内部 _sort (取负空间, 不外露) + objective_raw
(原始展示值)。best_score/进度回调均用原始值。

[中] 浮点累加丢端点: v += step 累积误差使 0.1 步长丢失 max 端点
(0.1+0.1+0.1=0.30000004 > 0.3+1e-9)。改整数计数 lo + i*step。

[中] backtest_kwargs 非法/冲突 key: 展开传给 StrategyBacktestConfig 时若含
非法或保留字段 (symbols 等) 会在 worker 抛 TypeError, 被上述 #1 放大成崩溃。
入口加白名单校验, 提前明确报错。

测试新增 8 例: 符号还原/max_drawdown 负值排序/异常隔离/kwargs 非法+保留/
base_params 合并/浮点端点/去重。原 test_min_direction 只验 params 放过了符号
bug, 现补 best_score 断言。

* feat(optimizer): 参数优化 API SSE 端点 + 前端优化器面板

PR2a 完成 (API + 前端), 接上后端核心:

后端 api/backtest.py:
- GET /optimize/stream: 复用 _BacktestJob SSE 框架, 后台线程跑 StrategyOptimizer,
  progress_cb 推 done/total/best_score, 完成推 best_params/results 排名。param_grid
  走 JSON 字符串查询参数 (EventSource 仅支持 GET, 同 params/overrides 惯例)。
- POST /optimize/cancel: 从 query string 复原同一 job_key, cancel_event 停止。
- backtest_kwargs 透传各回测参数 (matching/fees/mode/...) 到每组回测。

前端:
- lib/optimizerTask.ts: SSE 客户端 (镜像 backtestTask), 进度/结果/重连/取消。
- pages/backtest/StrategyOptimizer.tsx: 配置面板 (选策略 → 勾选可扫参数设
  min/max/step, bool/select 自动全扫; 优化目标下拉; 日期; 组合数实时预估 + 2000
  上限提示) + 结果面板 (最优参数高亮 + 排名表: objective/夏普/索提诺/收益/回撤/
  胜率/交易数)。
- Backtest.tsx: 新增 '参数优化' 第三 tab。

测试 test_optimizer_api.py 3 例: job_key 确定性 + 区分 grid/objective + stream 与
cancel 复原同一 key (守护 PR3 C1 类失配)。前端 tsc 无新增类型错误。

注: SSE 端到端需真实日K数据, 本地 mode=none 无法验证实际回测; 但 SSE 管线镜像
已测的 strategy_stream, 优化器核心 24 单测覆盖。

* fix(optimizer): 子代理审查修复 — API 取消/空网格/方向对齐 + 前端重连/切换/展示

两份子代理审查(API + 前端)发现的真实问题:

后端 API:
- [中] param_grid 为 null/[]/'' 等合法 JSON 但非网格对象时, 原逻辑跳过线程却不置
  job.done -> event_generator 永久空转、job 挂死在表中。改为非空 dict 校验拦下。
- [中] 取消后前端收到 done 而非取消提示: 优化器把 cancel 当每组失败正常返回 dict,
  done 分支照推'完成'。改为 done 分支先检查 cancel_event, 分流为'优化已取消'。
- [低] direction 空串边界: stream 侧 '' 与 cancel 侧 or None 口径不一致致 job_key
  失配(cancel 失效)。stream 加 direction = direction or None 对齐。

前端:
- [中] tryReconnectOptimize 是死代码(无调用方)违反 NO DEAD CODE: 接入 useEffect
  挂载恢复(镜像 StrategyBacktest 的 tryReconnect), 刷新/切页后恢复未完成优化。
- [中] 切策略后旧结果残留错配(参数列是旧策略): onSelectStrategy 加 clearOptimize。
- [低] 排名表 slice(50) 静默截断: 加'仅显示前50/共N组'提示。
- [低] objective_raw 裸数与 best_score 精度不一: 统一 toFixed(3)。

后端 76 测试通过; 前端 tsc 无新增类型错误。

* refactor(optimizer): job_key 回吐 — 消除 cancel 两侧重算的脆弱契约

采纳子代理审查建议, 从结构上根除整类 job_key 失配 bug:

之前 cancel 需从 query string 逐字段重算 job_key, 必须与 stream 侧完全一致 —
任何默认值/None-空串/类型转换漂移都静默导致取消失效 (PR3 C1、本轮 direction
空串失配都是这个结构的产物)。

改为: stream 首个 SSE 事件 (event: job) 回吐后端算出的 job_key, 前端存下,
cancel 直接原样传回按 key 查表。cancel 侧不再重算, 契约漂移无从发生。

- 后端 optimize_stream: 首事件 yield event: job {key}; optimize_cancel 简化为
  body.job_key 直接查 _running_jobs (删除 40 行 qs 重算)。
- 前端 optimizerTask: 监听 job 事件存 currentJobKey + localStorage; stopOptimize
  改传 {job_key}; done/error/cancel 清理 key。
- 测试: 原 stream/cancel qs 对齐测试已无意义, 改为验证 cancel 按回吐 key 查表
  (命中/已完成/未知 key 三态), 用轻量 fake Request 直调 endpoint。

后端 76 测试通过; 前端 tsc 无新增错误。策略回测路径未动 (已合并 + C1 测试守护),
本重构仅限本 PR 新增的优化器路径。

* chore(optimizer): 移除冗余 PanelCache 改动 — main 已独立实现线程安全

rebase 到 main 时发现上游已独立给 PanelCache 加锁 (且 compute 放锁外, 比本 PR
原方案更优), 并新增 asset_type 维度。故本 PR 的 PanelCache 改动 (加锁 + size/ttl
bump + docstring) 全部冗余且 docstring 已与 main 实际锁行为不符, 回退到 main 版本。
优化器共享单一 panel key, main 的 PanelCache 已完全够用。

至此本 PR 零 engine.py 改动。

* fix(optimizer): 处理 #82 review 的 6 处问题

作者 review #82 提出的阻塞/改进项, 逐条修复:

[阻塞1] 前端构建失败: EmptyState 只接受 title/hint, 误用了 description ->
npm run build (tsc -b) 报 TS2322。改为 hint, build 通过。

[阻塞2] 优化没用用户当前策略配置: optimize API/前端只传 strategy_id/param_grid/
objective/日期/mode, 未传 params/overrides。补齐 —— API 新增 params(base_params)/
overrides 两个 query 参数并纳入 job_key; 前端把选中策略的 params_defaults 作为未扫描
参数固定值, buildDefaultOverrides(strategy) 让 basic_filter/信号/风控按当前策略参与。
抽 lib/strategyOverrides.ts 共享 (与策略回测页同口径, 避免重复)。

[阻塞3] 切策略丢失运行中任务控制权: 原 clearOptimize 只清前端状态, 不关 SSE/不 cancel
后端/不清 localStorage -> 后端继续跑但 Stop 消失。改为: 有任务在跑时切策略先 stopOptimize
(真正 cancel + 关连接 + 清存储)。

[阻塞4] 停止按钮竞态: job_key 只在收到首个 job 事件后才有, 刚点开始就点停止时前端还没
key, cancel 落空。改为: stopOptimize 标记 cancelRequested, 有 key 立即 POST cancel,
无 key 则保持 SSE 等 job 事件到达时补发 cancel 再关 (关 SSE 不停后端 daemon 线程, 必须
真 POST); 加 5s 兜底。

[改进5] SSE 断线健壮性: 无 data 断线原全靠浏览器自动重连无上限。加 MAX_RECONNECT=5,
超限置 error 停 pending。

[改进6] 组合数校验与后端不一致: 前端 round((hi-lo)/step)+1 会把 min=0/max=1/step=0.6
显示为可运行 3 组, 但后端生成末值 1.2>max 报错。新增 sweepError 与后端 _candidates_for
同口径校验 (步长不整除/越界), 前端提前拦并禁用运行。

后端 156 测试通过; 前端 npm run build 通过。
2026-07-10 11:38:30 +08:00

722 lines
27 KiB
Python

"""回测 API — 信号回测 + 因子回测 + 策略回测。"""
from __future__ import annotations
import asyncio
import json
import queue
import threading
from dataclasses import asdict
from datetime import date, timedelta
from typing import Literal
from fastapi import APIRouter, HTTPException, Request
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, Field
from app.config import settings
from app.services.backtest import (
BacktestConfig,
BacktestService,
VectorbtUnavailable,
is_available,
)
router = APIRouter(prefix="/api/backtest", tags=["backtest"])
FACTOR_DEFAULT_DAYS = 180
STRATEGY_DEFAULT_DAYS = 365 * 3
BACKTEST_MAX_SERVER_DAYS = 186
FACTOR_MAX_SYMBOLS = 1000
BACKTEST_SERVER_GUARD_MESSAGE = (
"当前服务器内存约 1.8GB,回测区间最多支持 6 个月;"
"更长周期容易触发 OOM,建议在 8GB 以上内存环境或本机运行。"
)
def _get_engine(request: Request):
"""获取或创建 BacktestEngine (单例,PanelCache 跨请求生效)。"""
from app.backtest.engine import BacktestEngine
engine = getattr(request.app.state, "backtest_engine", None)
if engine is None:
engine = BacktestEngine(request.app.state.repo)
request.app.state.backtest_engine = engine
return engine
def _resolve_start(req: BaseModel, end: date, default_days: int) -> date:
"""未传 start 使用默认区间;显式传 null/空值表示全部历史。"""
start = getattr(req, "start")
if start is not None:
return start
if "start" in req.model_fields_set:
return date(1900, 1, 1)
return end - timedelta(days=default_days)
def _guard_server_backtest_range(start: date, end: date):
if not settings.backtest_range_guard:
return
days = (end - start).days + 1
if days > BACKTEST_MAX_SERVER_DAYS:
raise HTTPException(status_code=400, detail=BACKTEST_SERVER_GUARD_MESSAGE)
# ================================================================
# 状态
# ================================================================
@router.get("/status")
def status():
"""前端可用此接口判断回测页是否要灰显。"""
return {"available": True}
# ================================================================
# 信号回测 (现有接口,保持不变)
# ================================================================
class BacktestRequest(BaseModel):
symbols: list[str] = Field(..., min_length=1)
start: date | None = None
end: date | None = None
entries: list[str] = []
exits: list[str] = []
stop_loss_pct: float | None = None
max_hold_days: int | None = None
fees_pct: float = 0.0002
slippage_bps: float = 5
matching: Literal["close_t", "open_t+1"] = "close_t"
asset_type: str = "stock"
@router.post("/run")
def run(req: BacktestRequest, request: Request):
"""信号回测 — 现有接口,向后兼容。"""
repo = request.app.state.repo
svc = BacktestService(repo)
end = req.end or date.today()
start = req.start or (end - timedelta(days=365 * 3))
cfg = BacktestConfig(
symbols=req.symbols,
start=start,
end=end,
entries=req.entries,
exits=req.exits,
stop_loss_pct=req.stop_loss_pct,
max_hold_days=req.max_hold_days,
fees_pct=req.fees_pct,
slippage_bps=req.slippage_bps,
matching=req.matching,
asset_type=req.asset_type,
)
try:
result = svc.run(cfg)
except VectorbtUnavailable as e:
raise HTTPException(status_code=503, detail=str(e)) from e
return asdict(result)
# ================================================================
# 因子回测
# ================================================================
class FactorColumnsResponse(BaseModel):
columns: list[dict]
@router.get("/factor/columns")
def factor_columns():
"""返回可用的因子列列表。"""
from app.backtest.factor import FACTOR_COLUMNS
return {"columns": FACTOR_COLUMNS}
class FactorBacktestRequest(BaseModel):
factor_name: str
symbols: list[str] | None = None
start: date | None = None
end: date | None = None
n_groups: int = 5
rebalance: Literal["daily", "weekly", "monthly"] = "monthly"
weight: Literal["equal", "factor_weight"] = "equal"
fees_pct: float = 0.0002
slippage_bps: float = 5.0
asset_type: str = "stock"
@router.post("/factor/run")
def factor_run(req: FactorBacktestRequest, request: Request):
"""因子回测 — IC/IR 分析 + 分层回测。"""
from app.backtest.factor import FactorBacktestService, FactorConfig
engine = _get_engine(request)
svc = FactorBacktestService(engine)
end = req.end or date.today()
start = _resolve_start(req, end, STRATEGY_DEFAULT_DAYS)
_guard_server_backtest_range(start, end)
symbols = req.symbols if req.symbols else None
if symbols is not None and len(symbols) > FACTOR_MAX_SYMBOLS:
raise HTTPException(
status_code=400,
detail=f"指定标的最多支持 {FACTOR_MAX_SYMBOLS} 只,请缩小标的范围。",
)
cfg = FactorConfig(
factor_name=req.factor_name,
symbols=symbols,
start=start,
end=end,
n_groups=req.n_groups,
rebalance=req.rebalance,
weight=req.weight,
fees_pct=req.fees_pct,
slippage_bps=req.slippage_bps,
asset_type=req.asset_type,
)
result = svc.run(cfg)
return asdict(result)
# ================================================================
# 策略回测
# ================================================================
class StrategyBacktestRequest(BaseModel):
strategy_id: str
symbols: list[str] | None = None
start: date | None = None
end: date | None = None
params: dict | None = None
overrides: dict | None = None
# matching 向后兼容; 显式传 entry_fill/exit_fill 时以二者为准。
matching: Literal["close_t", "open_t+1"] = "open_t+1"
entry_fill: Literal["close_t", "open_t+1"] | None = None
exit_fill: Literal["close_t", "open_t+1"] | None = None
fees_pct: float = 0.0002
commission_pct: float | None = None
stamp_tax_pct: float | None = None
slippage_bps: float = 5.0
max_positions: int = 10
max_exposure_pct: float = 1.0
initial_capital: float = 1_000_000.0
position_sizing: Literal["equal", "score_weight"] = "equal"
mode: Literal["position", "full"] = "position"
holding_days: int = 5
asset_type: str = "stock"
@router.post("/strategy/run")
def strategy_run(req: StrategyBacktestRequest, request: Request):
"""策略回测 — 复用 StrategyDef 体系做全周期回测。"""
from app.backtest.strategy import StrategyBacktestService, StrategyBacktestConfig
engine = _get_engine(request)
strategy_engine = request.app.state.strategy_engine
svc = StrategyBacktestService(engine, strategy_engine)
end = req.end or date.today()
start = _resolve_start(req, end, FACTOR_DEFAULT_DAYS)
_guard_server_backtest_range(start, end)
cfg = StrategyBacktestConfig(
strategy_id=req.strategy_id,
symbols=req.symbols if req.symbols else None,
start=start,
end=end,
params=req.params,
overrides=req.overrides,
matching=req.matching,
entry_fill=req.entry_fill,
exit_fill=req.exit_fill,
fees_pct=req.fees_pct,
commission_pct=req.commission_pct,
stamp_tax_pct=req.stamp_tax_pct,
slippage_bps=req.slippage_bps,
max_positions=req.max_positions,
max_exposure_pct=req.max_exposure_pct,
initial_capital=req.initial_capital,
position_sizing=req.position_sizing,
mode=req.mode,
holding_days=req.holding_days,
asset_type=req.asset_type,
)
result = svc.run(cfg)
return asdict(result)
# ── SSE 流式回测 (实时进度 + 可取消 + 支持重连) ───────────────────
import time
import hashlib
class _BacktestJob:
"""单个回测任务的状态, 存模块级供重连使用。"""
__slots__ = ("key", "cancel_event", "progress", "result", "error", "done", "finish_ts")
def __init__(self, key: str):
self.key = key
self.cancel_event = threading.Event()
self.progress: list[dict] = [] # 进度历史 (新连接可回放)
self.result = None # 完成后的结果
self.error: str | None = None
self.done = False
self.finish_ts: float = 0.0
# 模块级任务表: key -> _BacktestJob
_running_jobs: dict[str, _BacktestJob] = {}
_jobs_lock = threading.Lock()
_JOB_TTL = 300 # 完成后保留 5 分钟
# 并发回测上限: 多个重回测同时跑会 OOM (服务器内存约 1.8GB)。用信号量限并发,
# 超出的任务在 _run_backtest 里排队, SSE 连接照常保持, run 一开始就有进度。
_backtest_semaphore = threading.Semaphore(2)
def _cleanup_stale_jobs():
"""清理过期任务 (完成超过 TTL 的)。全程持 _jobs_lock: 迭代+pop 与其他访问互斥。"""
now = time.time()
with _jobs_lock:
stale = [k for k, j in _running_jobs.items() if j.done and now - j.finish_ts > _JOB_TTL]
for k in stale:
_running_jobs.pop(k, None)
def _make_job_key(
strategy_id: str, symbols: str | None, start: str | None, end: str | None,
matching: str, entry_fill: str | None, exit_fill: str | None,
fees_pct: float, slippage_bps: float,
max_positions: int, max_exposure_pct: float, initial_capital: float, position_sizing: str,
params: str | None, overrides: str | None,
mode: str = "position", holding_days: int = 5,
commission_pct: float | None = None, stamp_tax_pct: float | None = None,
asset_type: str = "stock",
) -> str:
raw = f"{strategy_id}|{symbols}|{start}|{end}|{matching}|{entry_fill}|{exit_fill}|{fees_pct}|{slippage_bps}|{max_positions}|{max_exposure_pct}|{initial_capital}|{position_sizing}|{params}|{overrides}|{mode}|{holding_days}|{commission_pct}|{stamp_tax_pct}|{asset_type}"
return hashlib.md5(raw.encode()).hexdigest()[:12]
@router.get("/strategy/stream")
async def strategy_stream(
request: Request,
strategy_id: str,
symbols: str | None = None,
start: str | None = None,
end: str | None = None,
matching: str = "open_t+1",
entry_fill: str | None = None,
exit_fill: str | None = None,
fees_pct: float = 0.0002,
commission_pct: float | None = None,
stamp_tax_pct: float | None = None,
slippage_bps: float = 5.0,
max_positions: int = 10,
max_exposure_pct: float = 1.0,
initial_capital: float = 1_000_000.0,
position_sizing: str = "equal",
params: str | None = None,
overrides: str | None = None,
mode: str = "position",
holding_days: int = 5,
asset_type: str = "stock",
):
"""SSE 流式策略回测: 实时推送进度, 完成后推送结果, 支持重连 (刷新/切页后恢复)。
- 相同参数的任务只启动一次, 多次连接订阅同一个任务
- 断开连接不会取消任务 (除非显式调用 cancel)
- 结果保留 5 分钟供重连
事件类型:
- progress: {day, total, date, equity}
- done: {result} (完整回测结果)
- error: {message}
"""
from app.backtest.strategy import StrategyBacktestService, StrategyBacktestConfig
engine = _get_engine(request)
strategy_engine = request.app.state.strategy_engine
svc = StrategyBacktestService(engine, strategy_engine)
end_date = date.fromisoformat(end) if end else date.today()
if start:
start_date = date.fromisoformat(start)
else:
# 空 start = 全部历史: 用本地最早日K日期, 查不到再回退到默认窗口
earliest = request.app.state.repo.earliest_daily_date()
start_date = earliest or (end_date - timedelta(days=FACTOR_DEFAULT_DAYS))
# 服务端范围保护
guard_violated = False
if settings.backtest_range_guard:
days = (end_date - start_date).days + 1
if days > BACKTEST_MAX_SERVER_DAYS:
guard_violated = True
job_key = _make_job_key(
strategy_id, symbols, start, end,
matching, entry_fill, exit_fill,
fees_pct, slippage_bps, max_positions, max_exposure_pct, initial_capital, position_sizing,
params, overrides,
mode, holding_days,
commission_pct, stamp_tax_pct,
asset_type=asset_type,
)
_cleanup_stale_jobs()
# 获取或创建任务
with _jobs_lock:
job = _running_jobs.get(job_key)
if job is None:
job = _BacktestJob(job_key)
_running_jobs[job_key] = job
is_new = True
else:
is_new = False
async def event_generator():
# 范围保护: 直接报错
if guard_violated:
yield f"event: error\ndata: {json.dumps({'message': BACKTEST_SERVER_GUARD_MESSAGE}, ensure_ascii=False)}\n\n"
return
# 如果是新任务, 启动回测线程
if is_new and not job.done:
cfg = StrategyBacktestConfig(
strategy_id=strategy_id,
symbols=[s.strip() for s in symbols.split(",") if s.strip()] if symbols else None,
start=start_date,
end=end_date,
params=json.loads(params) if params else None,
overrides=json.loads(overrides) if overrides else None,
matching=matching,
entry_fill=entry_fill,
exit_fill=exit_fill,
fees_pct=fees_pct,
commission_pct=commission_pct,
stamp_tax_pct=stamp_tax_pct,
slippage_bps=slippage_bps,
max_positions=int(max_positions),
max_exposure_pct=float(max_exposure_pct),
initial_capital=float(initial_capital),
position_sizing=position_sizing,
mode=mode,
holding_days=int(holding_days),
asset_type=asset_type,
)
def _run_backtest():
# 信号量限并发: 超额任务在此阻塞排队, 不并发吃满内存 (等待期间 cancel_event
# 仍可置位, svc.run 会据此提前返回 cancelled)。持槽跑完在 finally 释放。
_backtest_semaphore.acquire()
try:
result = svc.run(cfg, lambda d: job.progress.append(d), job.cancel_event)
job.result = result
job.done = True
job.finish_ts = time.time()
except Exception as e:
job.error = str(e)
job.done = True
job.finish_ts = time.time()
finally:
_backtest_semaphore.release()
# 启动后台线程 (不阻塞事件循环)
threading.Thread(target=_run_backtest, daemon=True).start()
# 订阅进度: 用读指针读 job.progress 列表 (多连接互不干扰)
cursor = 0
tick = 0
try:
while True:
# 已完成: 推送最终结果/错误并退出
if job.done:
if job.error:
yield f"event: error\ndata: {json.dumps({'message': job.error}, ensure_ascii=False)}\n\n"
elif job.result is not None:
r = job.result
if hasattr(r, "error") and r.error == "cancelled":
yield f"event: error\ndata: {json.dumps({'message': '回测已取消'}, ensure_ascii=False)}\n\n"
elif hasattr(r, "error") and r.error:
yield f"event: error\ndata: {json.dumps({'message': r.error}, ensure_ascii=False)}\n\n"
else:
yield f"event: done\ndata: {json.dumps(asdict(r), ensure_ascii=False, default=str)}\n\n"
return
# 断开检测: 每 4 轮检查一次 (降低 GIL 抢占频率)
tick += 1
if tick % 4 == 0 and await request.is_disconnected():
break
# 推送新进度 (从 cursor 开始读)
prog_list = job.progress
while cursor < len(prog_list):
msg = prog_list[cursor]
cursor += 1
yield f"event: progress\ndata: {json.dumps(msg, ensure_ascii=False, default=str)}\n\n"
await asyncio.sleep(0.5)
except asyncio.CancelledError:
raise
return StreamingResponse(event_generator(), media_type="text/event-stream")
@router.post("/strategy/cancel")
async def strategy_cancel(request: Request):
"""取消正在运行的回测任务 (前端传 query string, 后端算 job_key)。"""
body = await request.json()
qs = body.get("qs", "")
# 解析 qs 得到参数
from urllib.parse import parse_qs
p = parse_qs(qs)
def _get(key: str, default: str = "") -> str:
return p.get(key, [default])[0]
def _get_opt_float(key: str) -> float | None:
# 可选成本参数: 缺省或空串 → None (与 stream 侧 float | None 口径一致, 保证 job_key 对齐)。
v = _get(key)
return float(v) if v else None
job_key = _make_job_key(
_get("strategy_id"),
_get("symbols") or None,
_get("start") or None,
_get("end") or None,
_get("matching", "open_t+1"),
_get("entry_fill") or None,
_get("exit_fill") or None,
float(_get("fees_pct", "0.0002")),
float(_get("slippage_bps", "5")),
int(_get("max_positions", "10")),
float(_get("max_exposure_pct", "1")),
float(_get("initial_capital", "1000000")),
_get("position_sizing", "equal"),
_get("params") or None,
_get("overrides") or None,
_get("mode", "position"),
int(_get("holding_days", "5")),
commission_pct=_get_opt_float("commission_pct"),
stamp_tax_pct=_get_opt_float("stamp_tax_pct"),
asset_type=_get("asset_type", "stock"),
)
# 持锁读任务表: 与 _cleanup_stale_jobs 的 pop、stream 的写入互斥
with _jobs_lock:
job = _running_jobs.get(job_key)
if job and not job.done:
job.cancel_event.set()
return {"ok": True}
return {"ok": False, "message": "任务不存在或已完成"}
# ══════════════════════════════════════════════════════════════
# 参数网格优化器 — 复用 _BacktestJob SSE 框架 (多组参数并行回测 + 排序)
# ══════════════════════════════════════════════════════════════
# 透传给每组回测的 StrategyBacktestConfig 字段 (作为 backtest_kwargs)。
_OPT_BT_FIELDS = [
"matching", "fees_pct", "commission_pct", "stamp_tax_pct", "slippage_bps",
"max_positions", "max_exposure_pct", "initial_capital", "position_sizing",
"mode", "holding_days",
]
def _make_opt_job_key(strategy_id, symbols, start, end, param_grid, objective, direction, bt_sig, params=None, overrides=None) -> str:
raw = f"OPT|{strategy_id}|{symbols}|{start}|{end}|{param_grid}|{objective}|{direction}|{bt_sig}|{params}|{overrides}"
return hashlib.md5(raw.encode()).hexdigest()[:12]
def _opt_backtest_kwargs(
matching, fees_pct, commission_pct, stamp_tax_pct, slippage_bps,
max_positions, max_exposure_pct, initial_capital, position_sizing, mode, holding_days,
) -> dict:
return {
"matching": matching,
"fees_pct": fees_pct,
"commission_pct": commission_pct,
"stamp_tax_pct": stamp_tax_pct,
"slippage_bps": slippage_bps,
"max_positions": int(max_positions),
"max_exposure_pct": float(max_exposure_pct),
"initial_capital": float(initial_capital),
"position_sizing": position_sizing,
"mode": mode,
"holding_days": int(holding_days),
}
@router.get("/optimize/stream")
async def optimize_stream(
request: Request,
strategy_id: str,
param_grid: str, # JSON: {param_id: [values] | {min,max,step}}
objective: str = "sortino",
direction: str | None = None,
max_workers: int = 4,
params: str | None = None, # JSON: 未扫描参数固定为用户当前值 (base_params)
overrides: str | None = None, # JSON: 策略当前的 basic_filter/signals/风控等覆盖
symbols: str | None = None,
start: str | None = None,
end: str | None = None,
matching: str = "open_t+1",
fees_pct: float = 0.0002,
commission_pct: float | None = None,
stamp_tax_pct: float | None = None,
slippage_bps: float = 5.0,
max_positions: int = 10,
max_exposure_pct: float = 1.0,
initial_capital: float = 1_000_000.0,
position_sizing: str = "equal",
mode: str = "position",
holding_days: int = 5,
):
"""SSE 流式参数优化: 并行跑各参数组回测, 按 objective 排序。
事件类型:
- progress: {type: "optimizer_progress", done, total, best_score}
- done: {result} (含 best_params / results 排名)
- error: {message}
"""
from app.backtest.optimizer import OptimizeConfig, StrategyOptimizer
from app.backtest.strategy import StrategyBacktestService
engine = _get_engine(request)
strategy_engine = request.app.state.strategy_engine
svc = StrategyBacktestService(engine, strategy_engine)
end_date = date.fromisoformat(end) if end else date.today()
if start:
start_date = date.fromisoformat(start)
else:
earliest = request.app.state.repo.earliest_daily_date()
start_date = earliest or (end_date - timedelta(days=FACTOR_DEFAULT_DAYS))
guard_violated = False
if settings.backtest_range_guard and (end_date - start_date).days + 1 > BACKTEST_MAX_SERVER_DAYS:
guard_violated = True
# 空串归一为 None, 与 cancel 侧 `_get("direction") or None` 口径一致, 避免 job_key 失配。
direction = direction or None
bt_kwargs = _opt_backtest_kwargs(
matching, fees_pct, commission_pct, stamp_tax_pct, slippage_bps,
max_positions, max_exposure_pct, initial_capital, position_sizing, mode, holding_days,
)
bt_sig = "|".join(f"{k}={bt_kwargs[k]}" for k in _OPT_BT_FIELDS)
job_key = _make_opt_job_key(strategy_id, symbols, start, end, param_grid, objective, direction, bt_sig, params, overrides)
_cleanup_stale_jobs()
with _jobs_lock:
job = _running_jobs.get(job_key)
if job is None:
job = _BacktestJob(job_key)
_running_jobs[job_key] = job
is_new = True
else:
is_new = False
async def event_generator():
# 首个事件回吐 job_key, 前端存下供 cancel 直接引用 (消除两侧重算契约)。
yield f"event: job\ndata: {json.dumps({'key': job_key}, ensure_ascii=False)}\n\n"
if guard_violated:
yield f"event: error\ndata: {json.dumps({'message': BACKTEST_SERVER_GUARD_MESSAGE}, ensure_ascii=False)}\n\n"
return
if is_new and not job.done:
try:
grid = json.loads(param_grid)
except (json.JSONDecodeError, TypeError):
grid = None
# grid 必须是非空 dict; null/[]/"" 等合法 JSON 但结构错误也在此拦下,
# 否则会跳过线程启动却不置 done -> event_generator 永久空转、job 挂死。
if not isinstance(grid, dict) or not grid:
job.error = "param_grid 必须是非空的参数网格对象"
job.done = True
job.finish_ts = time.time()
grid = None
if grid is not None:
# 未扫描参数固定为用户当前值 (base_params); overrides 让策略的 basic_filter/
# 信号/风控按用户当前配置参与, 保证优化的就是用户实际回测的策略。
try:
base_params = json.loads(params) if params else {}
except (json.JSONDecodeError, TypeError):
base_params = {}
try:
ov = json.loads(overrides) if overrides else None
except (json.JSONDecodeError, TypeError):
ov = None
ocfg = OptimizeConfig(
strategy_id=strategy_id,
symbols=[s.strip() for s in symbols.split(",") if s.strip()] if symbols else None,
start=start_date,
end=end_date,
param_grid=grid,
objective=objective,
direction=direction,
max_workers=int(max_workers),
base_params=base_params if isinstance(base_params, dict) else {},
overrides=ov if isinstance(ov, dict) else None,
backtest_kwargs=bt_kwargs,
)
def _run_opt():
try:
opt = StrategyOptimizer(svc, strategy_engine)
job.result = opt.optimize(ocfg, lambda d: job.progress.append(d), job.cancel_event)
job.done = True
job.finish_ts = time.time()
except Exception as e:
job.error = str(e)
job.done = True
job.finish_ts = time.time()
threading.Thread(target=_run_opt, daemon=True).start()
cursor = 0
tick = 0
try:
while True:
if job.done:
if job.error:
yield f"event: error\ndata: {json.dumps({'message': job.error}, ensure_ascii=False)}\n\n"
elif job.cancel_event.is_set():
# 取消时优化器把每组记为 cancelled 并正常返回, 需在此分流为取消提示而非"完成"。
yield f"event: error\ndata: {json.dumps({'message': '优化已取消'}, ensure_ascii=False)}\n\n"
elif job.result is not None:
yield f"event: done\ndata: {json.dumps(job.result, ensure_ascii=False, default=str)}\n\n"
return
tick += 1
if tick % 4 == 0 and await request.is_disconnected():
break
while cursor < len(job.progress):
msg = job.progress[cursor]
cursor += 1
yield f"event: progress\ndata: {json.dumps(msg, ensure_ascii=False, default=str)}\n\n"
await asyncio.sleep(0.5)
except asyncio.CancelledError:
raise
return StreamingResponse(event_generator(), media_type="text/event-stream")
@router.post("/optimize/cancel")
async def optimize_cancel(request: Request):
"""取消优化任务 — 前端传 stream 首事件回吐的 job_key, 后端直接查表。
不再让 cancel 侧重算 job_key: 两侧重算必须逐字段一致的脆弱契约(PR3 C1 / direction
空串失配都源于此)在此彻底消除。stream 首个 SSE 事件把后端算出的 key 回吐给前端,
cancel 原样传回即可。
"""
body = await request.json()
job_key = body.get("job_key", "")
job = _running_jobs.get(job_key)
if job and not job.done:
job.cancel_event.set()
return {"ok": True}
return {"ok": False, "message": "任务不存在或已完成"}