mirror of
https://ghfast.top/https://github.com/aeroxw/easy_tdx_max.git
synced 2026-09-12 21:34:21 +08:00
1384 lines
54 KiB
Python
1384 lines
54 KiB
Python
"""回测路由:策略枚举、同步回测、后台任务回测、任务轮询。
|
||
|
||
设计要点:
|
||
- 回测是纯计算(不依赖行情连接的 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"])
|
||
|
||
|
||
# ── 策略枚举 ───────────────────────────────────────────────────────────────────
|
||
|
||
|
||
@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)
|
||
state = runner.get(task_id)
|
||
# 提交瞬间任务应是 pending/running;极端情况下线程已跑完则报实际状态
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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)
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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)
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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)
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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)
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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)
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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)
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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)
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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)
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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)
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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)
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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)
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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} 轮询。
|
||
"""
|
||
from easy_tdx.web.convert import category_from_str, market_from_str
|
||
|
||
stock_dfs: dict[str, pd.DataFrame] = {}
|
||
for symbol in req.stocks:
|
||
market_str, code = symbol.split(":", 1)
|
||
try:
|
||
page = await client.get_security_bars(
|
||
market_from_str(market_str),
|
||
code,
|
||
category_from_str(req.category),
|
||
0,
|
||
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)
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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``)。
|
||
"""
|
||
from easy_tdx.web.convert import category_from_str, market_from_str
|
||
|
||
stock_dfs: dict[str, pd.DataFrame] = {}
|
||
for symbol in req.stocks:
|
||
market_str, code = symbol.split(":", 1)
|
||
try:
|
||
page = await client.get_security_bars(
|
||
market_from_str(market_str), code, category_from_str(req.category), 0, 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)
|
||
status: Any = state.status if state.status in ("pending", "running") else "running"
|
||
return TaskSubmitResponse(task_id=task_id, status=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(client: Any, symbol: str, category: str, count: int) -> pd.DataFrame:
|
||
"""按标的取 K 线(async,必须在 event loop 内调用)。"""
|
||
from easy_tdx.web.convert import category_from_str, market_from_str
|
||
|
||
market_str, code = symbol.split(":", 1)
|
||
df = await client.get_security_bars(
|
||
market_from_str(market_str),
|
||
code,
|
||
category_from_str(category),
|
||
0,
|
||
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 _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
|
||
|
||
optimizer = ParamGridOptimizer(
|
||
strategy_name=req.strategy,
|
||
param_grid=req.param_grid,
|
||
df=df,
|
||
cash=req.cash,
|
||
commission=req.commission,
|
||
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=req.commission,
|
||
slippage=req.slippage,
|
||
execution=req.execution,
|
||
)
|
||
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-bound,numpy/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()
|