Files
easy_tdx_max/src/easy_tdx/web/routers/backtest.py
T
Justin Gu e374a0da28 release: v1.32.6 — 两周改动深度审查全面修复(回测口径三件套/LLM 安全加固/涨停价舍入/时区统一/缓存与竞态等 58 处)
对 v1.21→v1.32.5 的 249 文件 4.2 万行改动做六路专项审查,本轮落地全部发现:

回测正确性:组合收益 fillna(0) 虚增、轮动停牌日过期价成交、单标的 WF 逐窗指标
被预热区稀释(三件套均带先红后绿回归);worst_drawdown 方向、grading 容错、
组合体检品种费率、寻优端点费率透传。

安全:LLM api_url 仅 http/https 且禁 userinfo(封死 file:// 读取与 Key 外送链)、
错误响应不回显原始 body、响应体 2MB 上限、配置原子写、坏配置字段级防御。

数据:涨跌停价整数分币舍入(67/318/90 个价位错 1 分漏判清零)、交易时段/采样/
provisional 统一沪时区、warehouse 增量缺口自动全量重拉、provisional 定点转正、
baostock 真故障抛错 + W/M 去 tradestatus(实测服务端报错,周月兜底此前从未工作)
+ 指数 vol 股→手(实测锚定)、ccpm 结构变更抛错。

Web API:缓存键补 count/vipdoc、NaN 清洗先于缓存、count>800 分页取全量、
submit 透传真实状态、pending 不再被淘汰成幽灵、watchlist/server 入参约束。

公式:FILTER 去副作用、0-1 值域误判收严、递归深度上限、REF 负移位显式禁止。

前端:4 处请求竞态序号守卫、Sparkline viewBox、北交所 market=2 映射、
空数据缓存死角、AI 弹窗卸载中止轮询、量能/资金日历口径修正。

CLI/CI:warehouse sync 失败 exit 1、参数校验干净报错、release 真实发布 SHA256、
CI 超时与缓存、spec 补 baostock 前提。

约 60 条回归测试先红后绿;pytest 1820 全过,ruff/mypy/vue-tsc/node --test 全绿。
2026-09-06 22:16:48 +08:00

1422 lines
55 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""回测路由:策略枚举、同步回测、后台任务回测、任务轮询。
设计要点:
- 回测是纯计算(不依赖行情连接的 lifespan),因此**不注入 tdx_client**——
只有「按标的取行情」才需要 client,且必须在 async 上下文里取好数据后再
交给后台线程跑回测(``get_security_bars`` 是 async,不能跨线程调用)。
- 后台任务用 :class:`~easy_tdx.web.task_runner.BacktestTaskRunner`,结果
线程安全,重启即丢。
- 同步回测仅支持内联 OHLCV(前端已有数据),避免长任务阻塞 event loop。
"""
from __future__ import annotations
from typing import Any
import pandas as pd
from fastapi import APIRouter, Depends
from easy_tdx.web.backtest_schemas import (
BacktestRequest,
BacktestResultResponse,
MultiSeedRequest,
MultiStrategyBacktestRequest,
OptimizeAllBacktestRequest,
OptimizeAllRankEntry,
OptimizeAllResult,
OptimizeBacktestRequest,
PortfolioBacktestRequest,
RotationBacktestRequest,
SignalScanRequest,
StrategySchemaResponse,
TaskListResponse,
TaskStateResponse,
TaskSubmitResponse,
TaskSummary,
serialize_result,
)
from easy_tdx.web.deps import get_client
from easy_tdx.web.task_runner import get_runner
router = APIRouter(tags=["backtest"])
# 标准 TdxClient 单次 get_security_bars 取数上限(协议约束,服务器对更大
# 请求静默截断)。所有按标的取数路径都必须经 _fetch_bars_paged 翻页。
_BARS_PAGE_SIZE = 800
# ── 策略枚举 ───────────────────────────────────────────────────────────────────
@router.get("/backtest/strategies", response_model=StrategySchemaResponse)
async def list_strategies() -> StrategySchemaResponse:
"""枚举所有预置策略及其参数 schema(供前端动态渲染策略选择 + 参数表单)。"""
from easy_tdx.backtest.strategies import get_registry
entries = get_registry().all()
schemas = [e.to_schema() for e in entries]
return StrategySchemaResponse(strategies=schemas, count=len(schemas))
# ── 同步回测(内联数据) ───────────────────────────────────────────────────────
@router.post("/backtest/run", response_model=BacktestResultResponse)
async def run_backtest(req: BacktestRequest) -> BacktestResultResponse:
"""同步回测(仅支持内联 OHLCV 数据)。
适用于单标的快速回测(<3s)。需要取行情或长任务请用 ``/backtest/run/async``。
"""
if req.ohlcv is None:
raise ValueError(
"同步回测(/backtest/run)必须提供 ohlcv 内联数据;取行情请用 /backtest/run/async"
)
df = _ohlcv_to_df(req.ohlcv)
result_dict = _run_backtest(df, req)
return BacktestResultResponse(**result_dict)
# ── 后台任务回测 ───────────────────────────────────────────────────────────────
@router.post("/backtest/run/async", response_model=TaskSubmitResponse, status_code=202)
async def run_backtest_async(
req: BacktestRequest,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交后台回测任务,立即返回 task_id。
支持内联数据或按标的取行情。取行情在 async 上下文完成(client 是 async 的),
之后回测在后台线程执行。通过 ``GET /backtest/tasks/{task_id}`` 轮询结果。
"""
# 1. 取数据(async 上下文内完成)
if req.ohlcv is not None:
df = _ohlcv_to_df(req.ohlcv)
bars_desc = f"{len(df)} 根"
elif req.symbol is not None:
df = await _fetch_bars(client, req.symbol, req.category, req.count)
bars_desc = f"{req.symbol} {req.category}×{req.count}"
else:
# BacktestRequest 校验器已保证二者至少其一,此处不可达
raise ValueError("必须提供 ohlcv 或 symbol")
# 2. 捕获回测所需的不可变快照(避免闭包捕获可变 req)
snapshot = req.model_copy()
description = f"{snapshot.strategy} | {bars_desc}"
# 3. 提交后台任务
runner = get_runner()
task_id = runner.submit(lambda: _run_backtest(df, snapshot), description=description)
# 提交瞬间通常是 pending/running;极快任务可能已 done/failed,如实上报
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
@router.get("/backtest/tasks", response_model=TaskListResponse)
async def list_tasks(limit: int = 20) -> TaskListResponse:
"""列出最近 N 个任务摘要(按最近使用倒序,不含完整 result)。
供对比页选择要对比的 task;选中后再逐个调 /tasks/{task_id} 拉详情。
"""
import time
runner = get_runner()
states = runner.list_recent(limit)
summaries = [
TaskSummary(
task_id=s.task_id,
status=s.status,
description=s.description,
created_at=s.created_at,
elapsed=(s.finished_at or time.time()) - (s.started_at or s.created_at),
)
for s in states
]
return TaskListResponse(tasks=summaries, count=len(summaries))
@router.get("/backtest/tasks/{task_id}", response_model=TaskStateResponse)
async def get_task(task_id: str) -> TaskStateResponse:
"""查询后台回测任务状态。done 时 result 字段含完整回测结果。"""
state = get_runner().peek(task_id)
if state is None:
# 未知 task → 404(通过 ValueError 走 400 handler;这里用 KeyError 由
# 调用方判定。为保持语义清晰,统一抛 ValueError → HTTP 400
raise ValueError(f"未知任务 '{task_id}'")
return TaskStateResponse(
task_id=state.task_id,
status=state.status,
result=state.result,
error=state.error,
description=state.description,
elapsed=(state.finished_at or _now()) - (state.started_at or state.created_at),
)
@router.get("/backtest/tasks/{task_id}/export")
async def export_task(task_id: str, format: str = "json") -> Any:
"""导出已完成任务的回测结果(JSON 全量 / CSV 主表)。
``format=json``:完整 result 字典(performance + equity_curve + trades +
config 等)。``format=csv``:导出结果中的主表——优先 trades(成交明细),
其次 ranking(寻优排名)、equity_curve(资金曲线);都缺时导出
performance 键值对。
"""
import csv
import io
import json as _json
from fastapi.responses import Response
state = get_runner().peek(task_id)
if state is None:
raise ValueError(f"未知任务 '{task_id}'")
if state.status != "done" or state.result is None:
raise ValueError(f"任务 '{task_id}' 尚未完成(status={state.status}),无法导出")
result = state.result
fmt = format.lower()
if fmt not in ("json", "csv"):
raise ValueError(f"不支持的导出格式 '{format}'(可选 json / csv")
if fmt == "json":
payload = _json.dumps(result, ensure_ascii=False, default=str)
return Response(
content=payload,
media_type="application/json",
headers={"Content-Disposition": f'attachment; filename="backtest-{task_id[:8]}.json"'},
)
# CSV:挑主表
rows: list[dict[str, Any]] | None = None
label = "metrics"
for key, name in (("trades", "trades"), ("ranking", "ranking"), ("equity_curve", "equity")):
val = result.get(key)
if isinstance(val, list) and val and isinstance(val[0], dict):
rows = val
label = name
break
if rows is not None:
buf = io.StringIO()
dict_writer = csv.DictWriter(buf, fieldnames=list(rows[0].keys()), extrasaction="ignore")
dict_writer.writeheader()
for r in rows:
dict_writer.writerow(r)
content = buf.getvalue()
else:
perf = result.get("performance")
if not isinstance(perf, dict):
raise ValueError(f"任务 '{task_id}' 的结果不含可导出的表格数据")
buf = io.StringIO()
list_writer = csv.writer(buf)
list_writer.writerow(["metric", "value"])
for k, v in perf.items():
list_writer.writerow([k, v])
content = buf.getvalue()
label = "performance"
return Response(
content=content,
media_type="text/csv; charset=utf-8",
headers={
"Content-Disposition": f'attachment; filename="backtest-{task_id[:8]}-{label}.csv"'
},
)
# ── 组合回测 ───────────────────────────────────────────────────────────────────
@router.post("/backtest/portfolio/run/async", response_model=TaskSubmitResponse, status_code=202)
async def run_portfolio_backtest_async(
req: PortfolioBacktestRequest,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交组合(多标的)回测后台任务。
逐个标的取行情(async),组装 StockData 列表后提交后台任务跑
PortfolioBacktestEngine。通过 GET /backtest/tasks/{task_id} 轮询结果。
"""
# 1. 逐个标的取行情(async 上下文内)
stock_data_list = await _fetch_portfolio_bars(
client, req.stocks, req.category, req.start_date, req.end_date
)
if not stock_data_list:
raise ValueError("所有标的均未取到有效行情数据")
# 2. 捕获不可变快照
snapshot = req.model_copy()
description = f"{snapshot.strategy} | {len(stock_data_list)}只标的"
# 3. 提交后台任务
runner = get_runner()
task_id = runner.submit(
lambda: _run_portfolio_backtest(stock_data_list, snapshot),
description=description,
)
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
# ── 多策略组合回测(资金分仓) ───────────────────────────────────────────────
@router.post(
"/backtest/multi-strategy/run/async", response_model=TaskSubmitResponse, status_code=202
)
async def run_multi_strategy_backtest_async(
req: MultiStrategyBacktestRequest,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交多策略组合回测后台任务(资金分仓 / 并行制)。
勾选 N 个策略,各自在原标的(取最新行情)上独立回测,各拿总资金 1/N。
单个策略取数失败则跳过(不中断整组),全部失败返回 400。结果为
MultiStrategyResult(结构同 PortfolioResult),通过 GET /backtest/tasks/{task_id} 轮询。
"""
slots = await _fetch_multi_strategy_bars(client, req.items)
if not slots:
raise ValueError("所有策略槽位均未取到有效行情数据")
snapshot = req.model_copy()
description = f"多策略组合 | {len(slots)}个策略"
runner = get_runner()
task_id = runner.submit(
lambda: _run_multi_strategy_backtest(slots, snapshot),
description=description,
)
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
@router.post(
"/backtest/multi-strategy/wf/run/async", response_model=TaskSubmitResponse, status_code=202
)
async def run_multi_strategy_walkforward_async(
req: MultiStrategyBacktestRequest,
n_windows: int = 7,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交多策略组合级 Walk-Forward 样本外验证后台任务。
逐槽位取行情后,按全部槽位日期并集切窗(预热区 + N 个连续测试窗),
每窗各槽位独立回测并合成组合窗内净值。结果为 ``{"walkforward": {...}}``
(与单标的 WF 同构),通过 GET /backtest/tasks/{task_id} 轮询。
"""
slots = await _fetch_multi_strategy_bars(client, req.items)
if not slots:
raise ValueError("所有策略槽位均未取到有效行情数据")
snapshot = req.model_copy()
description = f"多策略组合WF | {len(slots)}个策略 × {n_windows}窗"
runner = get_runner()
task_id = runner.submit(
lambda: _run_multi_strategy_walkforward(slots, snapshot, n_windows),
description=description,
)
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
@router.post(
"/backtest/multi-strategy/evaluate/run/async",
response_model=TaskSubmitResponse,
status_code=202,
)
async def run_multi_strategy_evaluate_async(
req: MultiStrategyBacktestRequest,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交多策略组合级一条龙评估后台任务:组合回测 + 组合 WF + 跨槽位适配性
体检 + 综合评分 + 组合评级 + 等权买入持有基准对比。
结果结构见 ``easy_tdx.backtest.benchmark.evaluate_multi`` 文档(与
单标的 evaluate_strategy 同构),通过 GET /backtest/tasks/{task_id} 轮询。
"""
slots = await _fetch_multi_strategy_bars(client, req.items)
if not slots:
raise ValueError("所有策略槽位均未取到有效行情数据")
snapshot = req.model_copy()
description = f"多策略组合一条龙 | {len(slots)}个策略"
runner = get_runner()
task_id = runner.submit(
lambda: _run_multi_strategy_evaluate(slots, snapshot),
description=description,
)
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
@router.post("/backtest/optimize/run/async", response_model=TaskSubmitResponse, status_code=202)
async def run_optimize_async(
req: OptimizeBacktestRequest,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交参数网格寻优后台任务。
在单个标的上对策略参数做网格搜索。数据获取支持内联 ohlcv 或按 symbol 取行情。
通过 GET /backtest/tasks/{task_id} 轮询结果。
"""
# 1. 取数据
if req.ohlcv is not None:
df = _ohlcv_to_df(req.ohlcv)
desc_bars = f"{len(df)} 根"
elif req.symbol is not None:
df = await _fetch_bars(client, req.symbol, req.category, 800)
desc_bars = f"{req.symbol}"
if req.start_date or req.end_date:
df = _filter_df_by_date(df, req.start_date, req.end_date)
else:
raise ValueError("必须提供 ohlcv 或 symbol")
# 2. 捕获快照
snapshot = req.model_copy()
grid_size = 1
for vals in snapshot.param_grid.values():
grid_size *= len(vals)
description = f"{snapshot.strategy} 寻优 | {desc_bars} | {grid_size}点"
# 3. 提交后台任务
runner = get_runner()
task_id = runner.submit(
lambda: _run_optimize(df, snapshot),
description=description,
)
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
# ── 一键寻优所有策略 ───────────────────────────────────────────────────────────
@router.post("/backtest/optimize-all/run/async", response_model=TaskSubmitResponse, status_code=202)
async def run_optimize_all_async(
req: OptimizeAllBacktestRequest,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交「一键寻优所有策略」后台任务。
在单个标的上,对所有策略的预设参数网格(见 presets.STRATEGY_PRESETS)依次
做网格寻优,取各策略最优点汇总成全局排名。数据获取支持内联 ohlcv 或按
symbol 取行情。通过 GET /backtest/tasks/{task_id} 轮询结果。
"""
# 1. 取数据
if req.ohlcv is not None:
df = _ohlcv_to_df(req.ohlcv)
desc_bars = f"{len(df)} 根"
elif req.symbol is not None:
df = await _fetch_bars(client, req.symbol, req.category, 800)
desc_bars = f"{req.symbol}"
if req.start_date or req.end_date:
df = _filter_df_by_date(df, req.start_date, req.end_date)
else:
raise ValueError("必须提供 ohlcv 或 symbol")
# 2. 捕获快照
snapshot = req.model_copy()
description = f"一键寻优全部策略 | {desc_bars}"
# 3. 提交后台任务
runner = get_runner()
task_id = runner.submit(
lambda: _run_optimize_all(df, snapshot),
description=description,
)
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
# ── 信号雷达(一键扫描已保存策略)────────────────────────────────────────────
@router.post("/backtest/signal-scan/run/async", response_model=TaskSubmitResponse, status_code=202)
async def run_signal_scan_async(
req: SignalScanRequest,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交「信号雷达」后台任务:扫描策略库全部已保存策略的最近买卖信号。
single/portfolio/multi 统一展开成"策略×标的"子任务,按 (symbol, category)
去重取最近 800 根 K 线(async 上下文内完成),后台线程内逐条跑信号流程
(与回测引擎同口径,含仓位跟踪)。只扫信号、不重跑回测、不改业绩快照。
结果为 SignalScanResult,通过 GET /backtest/tasks/{task_id} 轮询。
"""
from easy_tdx.web.signal_scan import expand_targets, fetch_scan_bars, run_scan
from easy_tdx.web.strategy_store import get_store
records = get_store().list_all()
if not records:
raise ValueError("策略库为空,请先在回测页保存策略")
targets = expand_targets(records)
bars = await fetch_scan_bars(client, targets)
description = (
f"信号扫描 | {len(records)}条策略 · {len(targets)}个子任务 · 窗口{req.window_bars}根"
)
runner = get_runner()
task_id = runner.submit(
lambda: run_scan(bars, targets, req.window_bars),
description=description,
)
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
# ── Walk-Forward / 一条龙评估(v1.25 防过拟合链)──────────────────────────────
@router.post("/backtest/wf/run/async", response_model=TaskSubmitResponse, status_code=202)
async def run_walkforward_async(
req: BacktestRequest,
n_windows: int = 7,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交 Walk-Forward 样本外验证后台任务(默认 7 窗、每窗独立开仓)。
数据来源同 /backtest/run/async(内联 ohlcv 或按 symbol 取行情)。
结果含逐窗收益、盈利窗占比 consistency、连乘收益等,通过
GET /backtest/tasks/{task_id} 轮询。
"""
df = await _resolve_df(client, req)
snapshot = req.model_copy()
description = f"{snapshot.strategy} WF验证 | {len(df)}根 | {snapshot.symbol or '内联数据'}"
runner = get_runner()
task_id = runner.submit(
lambda: _run_walkforward(df, snapshot, n_windows), description=description
)
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
@router.post("/backtest/evaluate/run/async", response_model=TaskSubmitResponse, status_code=202)
async def run_evaluate_async(
req: BacktestRequest,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交一条龙评估后台任务:回测 + WF + 适配性体检 + 综合评分 + S-D 评级
+ 买入持有基准对比(excess_return)。
结果结构见 ``easy_tdx.backtest.benchmark.evaluate_strategy`` 文档,
通过 GET /backtest/tasks/{task_id} 轮询;可用
GET /backtest/tasks/{task_id}/export?format=json 导出。
"""
df = await _resolve_df(client, req)
snapshot = req.model_copy()
description = f"{snapshot.strategy} 一条龙评估 | {snapshot.symbol or '内联数据'}"
runner = get_runner()
task_id = runner.submit(lambda: _run_evaluate(df, snapshot), description=description)
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
# ── 组合级 Walk-Forward / 一条龙评估(对齐单标的防过拟合链)──────────────────
@router.post("/backtest/portfolio/wf/run/async", response_model=TaskSubmitResponse, status_code=202)
async def run_portfolio_walkforward_async(
req: PortfolioBacktestRequest,
n_windows: int = 7,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交组合级 Walk-Forward 样本外验证后台任务。
逐个标的取行情后,按全部标的日期并集切窗(预热区 + N 个连续测试窗),
每窗各标的独立回测并合成组合窗内净值。结果为
``{"walkforward": {...}}``(与单标的 WF 同构),通过
GET /backtest/tasks/{task_id} 轮询。
"""
stock_data_list = await _fetch_portfolio_bars(
client, req.stocks, req.category, req.start_date, req.end_date
)
if not stock_data_list:
raise ValueError("所有标的均未取到有效行情数据")
snapshot = req.model_copy()
description = f"{snapshot.strategy} 组合WF | {len(stock_data_list)}只标的 × {n_windows}窗"
runner = get_runner()
task_id = runner.submit(
lambda: _run_portfolio_walkforward(stock_data_list, snapshot, n_windows),
description=description,
)
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
@router.post(
"/backtest/portfolio/evaluate/run/async", response_model=TaskSubmitResponse, status_code=202
)
async def run_portfolio_evaluate_async(
req: PortfolioBacktestRequest,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交组合级一条龙评估后台任务:组合回测 + 组合 WF + 跨标的适配性体检
+ 综合评分 + 组合评级 + 等权买入持有基准对比。
结果结构见 ``easy_tdx.backtest.benchmark.evaluate_portfolio`` 文档(与
单标的 evaluate_strategy 同构),通过 GET /backtest/tasks/{task_id} 轮询。
"""
stock_data_list = await _fetch_portfolio_bars(
client, req.stocks, req.category, req.start_date, req.end_date
)
if not stock_data_list:
raise ValueError("所有标的均未取到有效行情数据")
snapshot = req.model_copy()
description = f"{snapshot.strategy} 组合一条龙 | {len(stock_data_list)}只标的"
runner = get_runner()
task_id = runner.submit(
lambda: _run_portfolio_evaluate(stock_data_list, snapshot),
description=description,
)
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
async def _resolve_df(client: Any, req: BacktestRequest) -> pd.DataFrame:
"""内联 ohlcv 或按 symbol 取行情(/backtest/run/async 同逻辑的复用封装)。"""
if req.ohlcv is not None:
return _ohlcv_to_df(req.ohlcv)
if req.symbol is not None:
return await _fetch_bars(client, req.symbol, req.category, max(req.count, 800))
raise ValueError("必须提供 ohlcv 或 symbol")
@router.post("/backtest/multiseed/run/async", response_model=TaskSubmitResponse, status_code=202)
async def run_multiseed_async(
req: MultiSeedRequest,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交多 seed 随机抽样验证后台任务(v1.25 晋级门槛)。
股票池逐标的取行情(async),后台线程跑 MultiSeedValidator:多 seed
抽样 → 跨样本稳定性 → 四项晋级门槛(正收益比例/平均夏普/平均交易数/
平均收益)。结果含 per_seed_positive_ratio 稳定性列,通过
GET /backtest/tasks/{task_id} 轮询。
"""
stock_dfs: dict[str, pd.DataFrame] = {}
for symbol in req.stocks:
try:
page = await _fetch_bars_paged(client, symbol, req.category, req.count)
except Exception: # noqa: BLE001 — 单标的失败跳过
continue
if len(page) >= 30:
stock_dfs[symbol] = page
if len(stock_dfs) < 2:
raise ValueError(f"股票池有效标的不足 2 只({len(stock_dfs)}/{len(req.stocks)}")
snapshot = req.model_copy()
description = f"{snapshot.strategy} 多seed验证 | {len(stock_dfs)}只池×{snapshot.n_seeds}seed"
runner = get_runner()
task_id = runner.submit(lambda: _run_multiseed(stock_dfs, snapshot), description=description)
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
def _run_multiseed(stock_dfs: dict[str, pd.DataFrame], req: MultiSeedRequest) -> dict[str, Any]:
"""执行多 seed 验证(后台线程内调用)。"""
from easy_tdx.backtest.validation import MultiSeedValidator
validator = MultiSeedValidator(
strategy=_build_strategy_from(req.strategy, req.params),
stock_dfs=stock_dfs,
n_seeds=req.n_seeds,
sample_size=req.sample_size,
gates=req.gates,
cash=req.cash,
commission=req.commission,
min_commission=req.min_commission,
stamp_tax=req.stamp_tax,
slippage=req.slippage,
execution=req.execution,
auto_fees=req.auto_fees,
)
return validator.run().to_dict()
def _build_strategy_from(strategy_name: str, params: dict[str, Any]) -> Any:
"""按名 + 参数构造策略实例(registry KeyError → ValueError)。"""
from easy_tdx.backtest.strategies import get_registry
try:
entry = get_registry().get(strategy_name)
except KeyError as exc:
raise ValueError(str(exc)) from exc
return entry.build(params)
def _build_strategy(req: BacktestRequest) -> Any:
"""解析策略实例(registry KeyError → ValueError → HTTP 400)。"""
from easy_tdx.backtest.strategies import get_registry
try:
entry = get_registry().get(req.strategy)
except KeyError as exc:
raise ValueError(str(exc)) from exc
return entry.build(req.params)
def _run_walkforward(df: pd.DataFrame, req: BacktestRequest, n_windows: int = 7) -> dict[str, Any]:
"""执行 Walk-Forward 验证并附带常规回测绩效(后台线程内调用)。"""
from easy_tdx.backtest.walkforward import WalkForwardEngine
wf = WalkForwardEngine(
strategy=_build_strategy(req),
n_windows=n_windows,
cash=req.cash,
commission=req.commission,
min_commission=req.min_commission,
stamp_tax=req.stamp_tax,
slippage=req.slippage,
execution=req.execution,
symbol=req.symbol,
auto_fees=req.auto_fees,
).run(df)
return {"walkforward": wf.to_dict()}
def _run_evaluate(df: pd.DataFrame, req: BacktestRequest) -> dict[str, Any]:
"""执行一条龙评估(后台线程内调用)。"""
from easy_tdx.backtest.benchmark import evaluate_strategy
return evaluate_strategy(
strategy=_build_strategy(req),
df=df,
cash=req.cash,
commission=req.commission,
min_commission=req.min_commission,
stamp_tax=req.stamp_tax,
slippage=req.slippage,
execution=req.execution,
symbol=req.symbol,
auto_fees=req.auto_fees,
)
def _run_portfolio_walkforward(
stock_data_list: list[Any], req: PortfolioBacktestRequest, n_windows: int = 7
) -> dict[str, Any]:
"""执行组合级 Walk-Forward 验证(后台线程内调用)。"""
from easy_tdx.backtest.strategies import get_registry
from easy_tdx.backtest.walkforward import PortfolioWalkForwardEngine
try:
entry = get_registry().get(req.strategy)
except KeyError as exc:
raise ValueError(str(exc)) from exc
wf = PortfolioWalkForwardEngine(
strategy=entry.build(req.params),
stocks=stock_data_list,
n_windows=n_windows,
total_cash=req.cash,
commission=req.commission,
min_commission=req.min_commission,
stamp_tax=req.stamp_tax,
slippage=req.slippage,
execution=req.execution,
auto_fees=req.auto_fees,
).run()
return {"walkforward": wf.to_dict()}
def _run_portfolio_evaluate(
stock_data_list: list[Any], req: PortfolioBacktestRequest
) -> dict[str, Any]:
"""执行组合级一条龙评估(后台线程内调用)。"""
from easy_tdx.backtest.benchmark import evaluate_portfolio
from easy_tdx.backtest.strategies import get_registry
try:
entry = get_registry().get(req.strategy)
except KeyError as exc:
raise ValueError(str(exc)) from exc
return evaluate_portfolio(
strategy=entry.build(req.params),
stocks=stock_data_list,
total_cash=req.cash,
commission=req.commission,
min_commission=req.min_commission,
stamp_tax=req.stamp_tax,
slippage=req.slippage,
execution=req.execution,
auto_fees=req.auto_fees,
)
# ── 轮动组合回测(v1.27)─────────────────────────────────────────────────────
@router.post("/backtest/rotation/run/async", response_model=TaskSubmitResponse, status_code=202)
async def run_rotation_async(
req: RotationBacktestRequest,
client: Any = Depends(get_client),
) -> TaskSubmitResponse:
"""提交轮动组合回测后台任务(v1.27)。
按打分排名定期换仓:固定槽位等额、跌出排名自动卖出补位、日/周/月刷新、
可选槽内止盈止损。打分支持内置动量(``score="momentum"`` + ``period``
或通达信公式数值输出(``score="formula"`` + ``formula_text`` + ``score_col``)。
"""
stock_dfs: dict[str, pd.DataFrame] = {}
for symbol in req.stocks:
try:
page = await _fetch_bars_paged(client, symbol, req.category, req.count)
except Exception: # noqa: BLE001 — 单标的失败跳过
continue
if page is not None and len(page) >= 30:
stock_dfs[symbol] = page
if len(stock_dfs) < 2:
raise ValueError(f"股票池有效标的不足 2 只({len(stock_dfs)}/{len(req.stocks)}")
snapshot = req.model_copy()
description = f"轮动组合 | {len(stock_dfs)}× {snapshot.slots}× {snapshot.refresh}"
runner = get_runner()
task_id = runner.submit(lambda: _run_rotation(stock_dfs, snapshot), description=description)
state = runner.get(task_id)
return TaskSubmitResponse(task_id=task_id, status=state.status)
def _run_rotation(
stock_dfs: dict[str, pd.DataFrame], req: RotationBacktestRequest
) -> dict[str, Any]:
"""执行轮动回测(后台线程内调用)。"""
from easy_tdx.backtest.rotation import RotationEngine, formula_score, momentum_score
if req.score == "formula":
if not req.formula_text:
raise ValueError("score=formula 需要提供 formula_text")
score_fn = formula_score(req.formula_text, req.score_col)
else:
score_fn = momentum_score(req.period)
engine = RotationEngine(
stock_dfs=stock_dfs,
score_fn=score_fn,
slots=req.slots,
refresh=req.refresh,
cash=req.cash,
commission=req.commission,
min_commission=req.min_commission,
stamp_tax=req.stamp_tax,
stop_loss=req.stop_loss,
take_profit=req.take_profit,
)
out = engine.run().to_dict()
# 组合评级(净值曲线口径)
from easy_tdx.backtest.grading import grade_portfolio_equity
out["grade"] = grade_portfolio_equity(out["equity_curve"]).to_dict()
return out
# ── 内部实现 ───────────────────────────────────────────────────────────────────
def _run_backtest(df: pd.DataFrame, req: BacktestRequest) -> dict[str, Any]:
"""执行回测并返回清洗后的结果字典(后台线程内调用)。"""
from easy_tdx.backtest import BacktestEngine
from easy_tdx.backtest.strategies import get_registry
# 解析策略 + 校验参数(registry 抛 KeyError,统一转 ValueError → HTTP 400
try:
entry = get_registry().get(req.strategy)
except KeyError as exc:
raise ValueError(str(exc)) from exc
strategy = entry.build(req.params)
engine = BacktestEngine(
strategy=strategy,
cash=req.cash,
commission=req.commission,
min_commission=req.min_commission,
stamp_tax=req.stamp_tax,
slippage=req.slippage,
execution=req.execution,
symbol=req.symbol,
auto_fees=req.auto_fees,
)
result = engine.run(df)
out = serialize_result(result)
# v1.25:评级/评分后端化——REST 直接输出 S-D 档位与 0-100 综合分
from easy_tdx.backtest.grading import grade_performance
from easy_tdx.backtest.scoring import score_strategy
out["grade"] = grade_performance(dict(result.performance)).to_dict()
out["score"] = score_strategy(dict(result.performance)).to_dict()
return out
def _ohlcv_to_df(records: list[dict[str, Any]]) -> pd.DataFrame:
"""把内联 OHLCV 记录列表转为 DataFrame,校验必需列并把 datetime 转为真正的时间类型。
StrategyDataProxy 依赖 datetime 列为 datetime64/pandas Timestamp 才能正确
编码为 YYYYMMDD 整数;若内联数据传字符串日期,这里负责转换。
"""
required = {"datetime", "open", "high", "low", "close", "vol", "amount"}
df = pd.DataFrame(records)
missing = required - set(df.columns)
if missing:
raise ValueError(f"ohlcv 缺少必需列: {sorted(missing)};需要 {sorted(required)}")
if len(df) < 2:
raise ValueError(f"ohlcv 至少需要 2 根 K 线,当前 {len(df)} 根")
# 确保 datetime 是真正的时间类型(容忍字符串/数值输入)
if not pd.api.types.is_datetime64_any_dtype(df["datetime"]):
df["datetime"] = pd.to_datetime(df["datetime"], errors="coerce")
return df
def _normalize_bars_dt(df: pd.DataFrame) -> pd.DataFrame:
"""把取到的 K 线规范化为引擎可直接消费的列布局(返回新 df 或原 df)。
引擎(StrategyDataProxy / PortfolioTracker):时间列必须叫 ``datetime``
(``date`` 列会被当成数值列强转 float 而报错),类型接受 int YYYYMMDD
或 datetime64。真实 TDX 日线返回 int ``date`` 列、分钟线返回 ``datetime``
而 E2E mock 返回字符串 ``date``——这里统一:改名 ``date``→``datetime``、
字符串/对象类型 coerce 成 datetime64、删除遗留的 ``date`` 冗余列。
"""
if "datetime" not in df.columns and "date" in df.columns:
df = df.copy()
df["datetime"] = df["date"]
dt = df["datetime"]
if dt.dtype.kind not in "iu" and not pd.api.types.is_datetime64_any_dtype(dt):
df = df.copy()
df["datetime"] = pd.to_datetime(df["datetime"], errors="coerce")
if "date" in df.columns:
df = df.drop(columns=["date"])
return df
async def _fetch_bars_paged(client: Any, symbol: str, category: str, count: int) -> pd.DataFrame:
"""按 800/页翻页取最多 ``count`` 根 K 线,返回时间升序 DataFrame。
TDX 协议单次 get_security_bars 最多返回 800 根:count>800 的单次调用会被
服务器静默截断(multiseed / rotation / formula 曾各自单页取数,悄悄少
数据)。本辅助按 start=0,800,1600… 翻页拼接,页间按时间升序排序;
末页不足 800 根视为数据起点,提前停止。列结构与 get_security_bars
原始输出一致(日线 ``date`` / 分钟 ``datetime``),不做改名/类型规整。
"""
from easy_tdx.web.convert import category_from_str, market_from_str
market_str, code = symbol.split(":", 1)
market = market_from_str(market_str)
cat = category_from_str(category)
frames: list[pd.DataFrame] = []
fetched = 0
while fetched < count:
page_size = min(_BARS_PAGE_SIZE, count - fetched)
page_df = await client.get_security_bars(market, code, cat, fetched, page_size)
if page_df is None or len(page_df) == 0:
break
frames.append(page_df)
fetched += len(page_df)
if len(page_df) < page_size:
break # 数据起点
if not frames:
return pd.DataFrame()
df = pd.concat(frames, ignore_index=True)
dt_col = "datetime" if "datetime" in df.columns else "date"
if dt_col in df.columns:
# 页间天然逆序(page0=最新一页),拼接后按时间升序
df = df.sort_values(dt_col).reset_index(drop=True)
return df
async def _fetch_bars(client: Any, symbol: str, category: str, count: int) -> pd.DataFrame:
"""按标的取 K 线(async,必须在 event loop 内调用)。"""
df = await _fetch_bars_paged(client, symbol, category, count)
if len(df) == 0:
raise ValueError(f"标的 {symbol} 未取到任何 K 线数据")
return _normalize_bars_dt(df)
def _run_portfolio_backtest(
stock_data_list: list[Any], req: PortfolioBacktestRequest
) -> dict[str, Any]:
"""执行组合回测并返回清洗后的结果字典(后台线程内调用)。
与单标的 ``_run_backtest`` 对齐:附带组合评级(``grade_portfolio_equity``
净值曲线 5 维度口径)与综合评分(``score_strategy``,无 WF 时权重自动
归一化),供前端/REST 直接消费。
"""
from easy_tdx.backtest.grading import grade_portfolio_equity
from easy_tdx.backtest.portfolio_engine import PortfolioBacktestEngine
from easy_tdx.backtest.scoring import score_strategy
from easy_tdx.backtest.strategies import get_registry
try:
entry = get_registry().get(req.strategy)
except KeyError as exc:
raise ValueError(str(exc)) from exc
strategy = entry.build(req.params)
engine = PortfolioBacktestEngine(
strategy=strategy,
stocks=stock_data_list,
total_cash=req.cash,
commission=req.commission,
min_commission=req.min_commission,
stamp_tax=req.stamp_tax,
slippage=req.slippage,
execution=req.execution,
auto_fees=req.auto_fees,
)
result = engine.run()
out = serialize_result(result)
# 组合评级(净值曲线口径)+ 综合评分——与单标的回测响应同构
if len(result.combined_equity) >= 2:
out["grade"] = grade_portfolio_equity(
result.combined_equity.to_dict(orient="records")
).to_dict()
out["score"] = score_strategy(dict(result.total_performance)).to_dict()
return out
async def _fetch_portfolio_bars(
client: Any,
stocks: list[str],
category: str,
start_date: str | None,
end_date: str | None,
) -> list[Any]:
"""逐个标的取 K 线并组装 StockData 列表(async,必须在 event loop 内调用)。
当 start_date 超出单次 800 根覆盖范围时,自动翻页拉取(与前端 fetchBars
同逻辑)。单个标的取数失败时跳过(不中断整个组合),全部失败返回空列表。
"""
from easy_tdx.backtest.portfolio_engine import StockData
from easy_tdx.web.convert import category_from_str, market_from_str
max_pages = 10 # 翻页上限:10 × 800 = 8000 根
stock_data_list: list[StockData] = []
for symbol in stocks:
market_str, code = symbol.split(":", 1)
frames: list[pd.DataFrame] = []
for page in range(max_pages):
try:
page_df = await client.get_security_bars(
market_from_str(market_str),
code,
category_from_str(category),
page * 800,
800,
)
except Exception:
break # 单页失败则停止该标的的翻页
if len(page_df) == 0:
break
frames.append(page_df)
# 已覆盖到 start_date(本页最早一根 ≤ start_date)则停止
if start_date and len(page_df) > 0:
dt_col = "datetime" if "datetime" in page_df.columns else "date"
oldest = str(page_df[dt_col].iloc[-1])[:10]
if oldest <= start_date:
break
if len(page_df) < 800:
break # 数据起点
if not frames:
continue
df = pd.concat(frames, ignore_index=True)
# 列名归一化:日线返回 date,分钟线返回 datetime
if "datetime" not in df.columns and "date" in df.columns:
df = df.copy()
df["datetime"] = df["date"]
# 翻页拼接后按时间正序排序(页间逆序)
df = df.sort_values("datetime").reset_index(drop=True)
df = _normalize_bars_dt(df)
# 日期范围过滤
if start_date or end_date:
dt_str = df["datetime"].astype(str).str.slice(0, 10)
mask = pd.Series(True, index=df.index)
if start_date:
mask &= dt_str >= start_date
if end_date:
mask &= dt_str <= end_date
df = df[mask]
if len(df) < 2:
continue
stock_data_list.append(
StockData(code=code, market=market_str, df=df.reset_index(drop=True))
)
return stock_data_list
async def _fetch_multi_strategy_bars(
client: Any,
items: list[Any],
) -> list[Any]:
"""逐个策略槽位取行情 + 构造策略实例,组装 StrategySlot 列表(async)。
每条 item 自带 symbol(如 "SH:601088")、category、start/end_date、strategy+params。
单条取数或策略构造失败则跳过(不中断整组)。返回的 StrategySlot 已绑定好策略
实例与 df,可直接交给后台线程跑引擎(避免把 async client 带进线程)。
"""
from easy_tdx.backtest.multi_strategy_engine import StrategySlot
from easy_tdx.backtest.strategies import get_registry
from easy_tdx.web.convert import category_from_str, market_from_str
registry = get_registry()
slots: list[StrategySlot] = []
for item in items:
# 1. 解析策略(未知策略跳过)
try:
entry = registry.get(item.strategy)
except KeyError:
continue
# 2. 逐页取行情(覆盖 start_date,最多 10 页 = 8000 根)
market_str, code = item.symbol.split(":", 1)
frames: list[pd.DataFrame] = []
for page in range(10):
try:
page_df = await client.get_security_bars(
market_from_str(market_str),
code,
category_from_str(item.category),
page * 800,
800,
)
except Exception:
break
if len(page_df) == 0:
break
frames.append(page_df)
if item.start_date and len(page_df) > 0:
dt_col = "datetime" if "datetime" in page_df.columns else "date"
oldest = str(page_df[dt_col].iloc[-1])[:10]
if oldest <= item.start_date:
break
if len(page_df) < 800:
break
if not frames:
continue
df = pd.concat(frames, ignore_index=True)
if "datetime" not in df.columns and "date" in df.columns:
df = df.copy()
df["datetime"] = df["date"]
df = df.sort_values("datetime").reset_index(drop=True)
df = _normalize_bars_dt(df)
# 日期范围过滤
if item.start_date or item.end_date:
df = _filter_df_by_date(df, item.start_date, item.end_date)
if len(df) < 2:
continue
# 3. 构造策略实例(参数非法跳过该条)
try:
strategy = entry.build(item.params)
except ValueError:
continue
label = item.strategy_label or entry.label
slots.append(StrategySlot(label=label, symbol=item.symbol, strategy=strategy, df=df))
return slots
def _run_multi_strategy_backtest(
slots: list[Any], req: MultiStrategyBacktestRequest
) -> dict[str, Any]:
"""执行多策略组合回测并返回清洗后的结果字典(后台线程内调用)。
与组合回测 ``_run_portfolio_backtest`` 同构:附带组合评级(净值口径)
与综合评分,供前端/REST 直接消费。
"""
from easy_tdx.backtest.grading import grade_portfolio_equity
from easy_tdx.backtest.multi_strategy_engine import MultiStrategyEngine
from easy_tdx.backtest.scoring import score_strategy
engine = MultiStrategyEngine(
strategies=slots,
total_cash=req.cash,
commission=req.commission,
min_commission=req.min_commission,
stamp_tax=req.stamp_tax,
slippage=req.slippage,
execution=req.execution,
)
result = engine.run()
out = serialize_result(result)
# 组合评级(净值曲线口径)+ 综合评分——与单标的/多标的组合响应同构
if len(result.combined_equity) >= 2:
out["grade"] = grade_portfolio_equity(
result.combined_equity.to_dict(orient="records")
).to_dict()
out["score"] = score_strategy(dict(result.total_performance)).to_dict()
return out
def _run_multi_strategy_walkforward(
slots: list[Any], req: MultiStrategyBacktestRequest, n_windows: int = 7
) -> dict[str, Any]:
"""执行多策略组合级 Walk-Forward 验证(后台线程内调用)。"""
from easy_tdx.backtest.walkforward import MultiStrategyWalkForwardEngine
wf = MultiStrategyWalkForwardEngine(
strategies=slots,
n_windows=n_windows,
total_cash=req.cash,
commission=req.commission,
min_commission=req.min_commission,
stamp_tax=req.stamp_tax,
slippage=req.slippage,
execution=req.execution,
).run()
return {"walkforward": wf.to_dict()}
def _run_multi_strategy_evaluate(
slots: list[Any], req: MultiStrategyBacktestRequest
) -> dict[str, Any]:
"""执行多策略组合级一条龙评估(后台线程内调用)。"""
from easy_tdx.backtest.benchmark import evaluate_multi
return evaluate_multi(
strategies=slots,
total_cash=req.cash,
commission=req.commission,
min_commission=req.min_commission,
stamp_tax=req.stamp_tax,
slippage=req.slippage,
execution=req.execution,
)
def _resolve_effective_fees(
auto_fees: bool,
symbol: str | None,
commission: float,
min_commission: float,
stamp_tax: float,
) -> tuple[float, float, float]:
"""auto_fees 品种费率解析(与 BacktestEngine 同款口径)。
显式非默认值优先(调用方有意覆盖),默认值按品种费率表替换(如
ETF/可转债免印花税)。ParamGridOptimizer 无 auto_fees 参数,寻优端点
在 web 层预解析成具体费率再传入,保证与单标的回测同口径。
"""
if not auto_fees or not symbol:
return commission, min_commission, stamp_tax
from easy_tdx.backtest.fees import resolve_fee_model
fee = resolve_fee_model(symbol)
if commission == 0.0003:
commission = fee.commission
if min_commission == 5.0:
min_commission = fee.min_commission
if stamp_tax == 0.001:
stamp_tax = fee.stamp_tax
return commission, min_commission, stamp_tax
def _run_optimize(df: pd.DataFrame, req: OptimizeBacktestRequest) -> dict[str, Any]:
"""执行参数网格寻优并返回清洗后的结果字典(后台线程内调用)。"""
from easy_tdx.backtest.benchmark import run_buy_hold_benchmark
from easy_tdx.backtest.optimizer import ParamGridOptimizer
commission, min_commission, stamp_tax = _resolve_effective_fees(
req.auto_fees, req.symbol, req.commission, req.min_commission, req.stamp_tax
)
optimizer = ParamGridOptimizer(
strategy_name=req.strategy,
param_grid=req.param_grid,
df=df,
cash=req.cash,
commission=commission,
min_commission=min_commission,
stamp_tax=stamp_tax,
slippage=req.slippage,
execution=req.execution,
workers=req.workers,
)
result = optimizer.run()
out = result.to_dict()
# 买入持有基准(同区间/同费率/同资金,与一条龙评估同口径),供前端在
# 最优结果旁直观对比「策略 vs 买入不动」。
out["buy_hold"] = run_buy_hold_benchmark(
df,
cash=req.cash,
commission=commission,
min_commission=min_commission,
stamp_tax=stamp_tax,
slippage=req.slippage,
execution=req.execution,
symbol=req.symbol,
auto_fees=req.auto_fees,
)
return out
def _optimize_one_strategy(
strategy_name: str,
grid: dict[str, list[Any]],
df: pd.DataFrame,
cash: float,
commission: float,
slippage: float,
execution: str,
) -> dict[str, Any] | None:
"""跑单个策略的网格寻优,返回其最优点摘要(模块顶层,可被 ProcessPoolExecutor pickle)。
必须是模块级顶层函数:Windows 下 ProcessPoolExecutor 用 spawn 方式启动子进程,
子进程按 ``module.qualname`` 重新 import 本函数。lambda / 闭包 / 嵌套函数不可 pickle。
策略类(``registry.get(name).build()``)在子进程内构造,从不跨进程传递,
因此天然避开了 screen scanner 当年遇到的"策略类不可 pickle"问题。
返回纯 dict(所有值都是 JSON 原生类型),可安全 pickle 回主进程。
"""
from easy_tdx.backtest.optimizer import ParamGridOptimizer
try:
optimizer = ParamGridOptimizer(
strategy_name=strategy_name,
param_grid=grid,
df=df,
cash=cash,
commission=commission,
slippage=slippage,
execution=execution,
)
except ValueError:
# 单策略网格超限(不应发生,预设已控制规模)→ 跳过
return None
result = optimizer.run()
if result.best is None:
return None
return {
"strategy": strategy_name,
"params": result.best.params,
"total_return": result.best.total_return,
"sharpe": result.best.sharpe,
"max_drawdown": result.best.max_drawdown,
"total_trades": result.best.total_trades,
"win_rate": result.best.win_rate,
"profit_factor": result.best.profit_factor,
"grid_points": len(result.results),
}
def _run_optimize_all(df: pd.DataFrame, req: OptimizeAllBacktestRequest) -> dict[str, Any]:
"""对所有策略的预设网格逐策略寻优,汇总成全局排名(后台线程内调用)。
遍历 ``STRATEGY_PRESETS`` 中每个策略,用其预设参数网格跑
:class:`ParamGridOptimizer`,取各策略的最优点(best)组装排名。单个策略
无有效结果(如全网格回测失败)则跳过。
并发:``req.workers >= 2`` 时用 ``ProcessPoolExecutor`` 跨进程并行寻优
(回测是 CPU-boundnumpy/pandas 持 GIL,线程无加速,必须用进程)。
``workers`` 为 0 或 1 时串行。进程池在函数内 ``with`` 创建/销毁,对前端
轮询与 task_runner 透明。
"""
from easy_tdx.backtest.benchmark import run_buy_hold_benchmark
from easy_tdx.backtest.strategies import get_registry
from easy_tdx.backtest.strategies.presets import STRATEGY_PRESETS
registry = get_registry()
# 过滤出已注册的策略 + 解析 label(label 必须在主进程取,避免子进程各自解析不一致)
jobs: list[tuple[str, dict[str, list[Any]]]] = []
labels: dict[str, str] = {}
for strategy_name, grid in STRATEGY_PRESETS.items():
if strategy_name not in registry.names():
continue
labels[strategy_name] = registry.get(strategy_name).label
jobs.append((strategy_name, grid))
# 跑寻优:串行 or 进程池并行
raw_results: list[dict[str, Any]] = []
if req.workers and req.workers >= 2:
import concurrent.futures
with concurrent.futures.ProcessPoolExecutor(max_workers=req.workers) as executor:
futures = {
executor.submit(
_optimize_one_strategy,
name,
grid,
df,
req.cash,
req.commission,
req.slippage,
req.execution,
): name
for name, grid in jobs
}
for future in concurrent.futures.as_completed(futures):
res = future.result()
if res is not None:
raw_results.append(res)
else:
for name, grid in jobs:
res = _optimize_one_strategy(
name, grid, df, req.cash, req.commission, req.slippage, req.execution
)
if res is not None:
raw_results.append(res)
# 组装排名(主进程统一构造 Pydantic 模型,保证类型一致)
ranking: list[OptimizeAllRankEntry] = []
per_strategy: dict[str, OptimizeAllRankEntry] = {}
total_grid = 0
for res in raw_results:
strategy_name = res["strategy"]
entry = OptimizeAllRankEntry(
strategy=strategy_name,
strategy_label=labels[strategy_name],
params=res["params"],
total_return=res["total_return"],
sharpe=res["sharpe"],
max_drawdown=res["max_drawdown"],
total_trades=res["total_trades"],
win_rate=res["win_rate"],
profit_factor=res["profit_factor"],
grid_points=res["grid_points"],
)
ranking.append(entry)
per_strategy[strategy_name] = entry
total_grid += res["grid_points"]
# 按 total_return 降序
ranking.sort(key=lambda r: r.total_return, reverse=True)
best = ranking[0] if ranking else None
result_obj = OptimizeAllResult(
ranking=ranking,
best=best,
per_strategy=per_strategy,
total_grid_points=total_grid,
# 买入持有基准(同区间/同费率/同资金),供前端在全局最佳旁直观对比
buy_hold=run_buy_hold_benchmark(
df,
cash=req.cash,
commission=req.commission,
slippage=req.slippage,
execution=req.execution,
),
)
return result_obj.model_dump()
def _filter_df_by_date(df: pd.DataFrame, start: str | None, end: str | None) -> pd.DataFrame:
"""按日期范围过滤 DataFrame(闭区间,比较 YYYY-MM-DD)。"""
if not start and not end:
return df
dt_col = "datetime" if "datetime" in df.columns else "date"
dt_str = df[dt_col].astype(str).str.slice(0, 10)
mask = pd.Series(True, index=df.index)
if start:
mask &= dt_str >= start
if end:
mask &= dt_str <= end
return df[mask].reset_index(drop=True)
def _now() -> float:
"""获取当前时间戳(隔离 import,便于测试)。"""
import time
return time.time()