Files
easy_tdx_max/src/easy_tdx/web/backtest_schemas.py
T

563 lines
22 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.
"""回测 Web API 的请求/响应模型与结果序列化。
与 ``easy_tdx.web.schemas`` 分离,避免主 schemas 文件膨胀。结果序列化复用
主 schemas 的 numpy/datetime 清洗思路,但递归处理嵌套结构(BacktestResult
的 performance 是 dict、equity_curve 是 list[dict]、config 是 dict)。
"""
from __future__ import annotations
from typing import Any, Literal
from pydantic import BaseModel, Field, model_validator
__all__ = [
"BacktestRequest",
"BacktestResultResponse",
"StrategySchemaResponse",
"TaskSubmitResponse",
"TaskStateResponse",
"OptimizeAllBacktestRequest",
"OptimizeAllResult",
"OptimizeAllRankEntry",
"SavedStrategy",
"SavedStrategyCreate",
"SavedStrategyListResponse",
"MultiStrategyItem",
"MultiStrategyBacktestRequest",
"SignalScanRequest",
"SignalScanRecentSignal",
"SignalScanRow",
"SignalScanResult",
"serialize_result",
]
# ── 请求模型 ───────────────────────────────────────────────────────────────────
class BacktestRequest(BaseModel):
"""单标的回测请求。
两种数据来源(二选一):
- ``ohlcv``: 直接内联 OHLCV 记录(前端已有数据时)。
- ``symbol`` + ``category`` + ``count``: 指定标的,由后端取行情。
两者都给则以内联数据为准。
"""
strategy: str = Field(..., description="策略名(见 /backtest/strategies")
params: dict[str, Any] = Field(default_factory=dict, description="策略参数")
cash: float = Field(default=1_000_000.0, gt=0, description="初始资金")
commission: float = Field(default=0.0003, ge=0, le=0.01, description="佣金费率")
min_commission: float = Field(default=5.0, ge=0, description="单笔最低佣金")
stamp_tax: float = Field(default=0.001, ge=0, le=0.01, description="印花税(卖出)")
slippage: float = Field(default=0.0, ge=0, le=0.05, description="滑点费率")
execution: Literal["next_open", "next_close"] = Field(
default="next_open", description="成交模式"
)
auto_fees: bool = Field(
default=False,
description="按标的品种自动解析费率(ETF/可转债免印花税等);显式非默认费率优先",
)
# 数据来源 A:内联 OHLCV(上限与 symbol 路径的 count 上限对齐,防 DoS)
ohlcv: list[dict[str, Any]] | None = Field(
default=None,
max_length=2000,
description="内联 K 线记录(含 datetime/open/high/low/close/vol/amount,最多 2000 条)",
)
# 数据来源 B:按标的取行情
symbol: str | None = Field(
default=None,
pattern=r"^(SZ|SH|BJ):\d{6}$",
description="标的代码,格式 市场:代码,如 SZ:000001",
)
category: Literal["DAY", "WEEK", "MONTH", "MIN_5", "MIN_15", "MIN_30", "MIN_60"] = Field(
default="DAY", description="K 线周期"
)
count: int = Field(default=250, ge=20, le=2000, description="K 线根数")
@model_validator(mode="after")
def _check_data_source(self) -> BacktestRequest:
if self.ohlcv is None and self.symbol is None:
raise ValueError("必须提供 ohlcv(内联数据)或 symbol(标的代码)之一")
return self
class PortfolioBacktestRequest(BaseModel):
"""组合(多标的)回测请求。
与单标的 BacktestRequest 的区别:用 ``stocks`` 列表替代单个 ``symbol``
资金按 equal 模式均分到各标的。日期范围过滤由前端完成(取满 800 根后过滤)。
"""
strategy: str = Field(..., description="策略名")
params: dict[str, Any] = Field(default_factory=dict, description="策略参数")
cash: float = Field(default=1_000_000.0, gt=0, description="组合总资金")
commission: float = Field(default=0.0003, ge=0, le=0.01)
min_commission: float = Field(default=5.0, ge=0)
stamp_tax: float = Field(default=0.001, ge=0, le=0.01)
slippage: float = Field(default=0.0, ge=0, le=0.05)
execution: Literal["next_open", "next_close"] = Field(default="next_open")
auto_fees: bool = Field(
default=False, description="按各标的品种逐只解析费率(ETF/可转债免印花税等)"
)
stocks: list[str] = Field(
...,
min_length=1,
max_length=20,
description='标的列表,格式 "市场:代码",如 ["SZ:000001","SH:600519"]',
)
category: Literal["DAY", "WEEK", "MONTH", "MIN_5", "MIN_15", "MIN_30", "MIN_60"] = Field(
default="DAY"
)
start_date: str | None = Field(default=None, description="开始日期 YYYY-MM-DD(可选过滤)")
end_date: str | None = Field(default=None, description="结束日期 YYYY-MM-DD(可选过滤)")
@model_validator(mode="after")
def _check_stocks_format(self) -> PortfolioBacktestRequest:
for s in self.stocks:
if not s.startswith(("SZ:", "SH:", "BJ:")) or len(s) != 9:
raise ValueError(f"标的格式应为 '市场:6位代码',得到 {s!r}")
return self
class OptimizeBacktestRequest(BaseModel):
"""参数网格寻优请求。
在单个标的上,对策略的 1-2 个参数做网格搜索。数据来源与单标的回测一致
ohlcv 内联或 symbol 取行情)。param_grid 指定寻优参数及其取值列表,
网格大小(各取值数乘积)上限 200。
"""
strategy: str = Field(..., description="策略名")
cash: float = Field(default=1_000_000.0, gt=0)
commission: float = Field(default=0.0003, ge=0, le=0.01)
slippage: float = Field(default=0.0, ge=0, le=0.05)
execution: Literal["next_open", "next_close"] = Field(default="next_open")
workers: int = Field(
default=1,
ge=0,
le=32,
description="并行工作进程数:0/1=串行+指标缓存(默认);2+=进程级并行(实测 36 点网格 "
"800 根 K 线约 2 倍,网格越大收益越高)",
)
param_grid: dict[str, list[int | float | str]] = Field(
...,
min_length=1,
max_length=2,
description='参数取值网格,如 {"fast":[5,10,20], "slow":[15,20,30]}',
)
# 数据来源 A:内联 OHLCV
ohlcv: list[dict[str, Any]] | None = Field(default=None, max_length=2000)
# 数据来源 B:按标的取行情
symbol: str | None = Field(default=None, pattern=r"^(SZ|SH|BJ):\d{6}$")
category: Literal["DAY", "WEEK", "MONTH", "MIN_5", "MIN_15", "MIN_30", "MIN_60"] = Field(
default="DAY"
)
count: int = Field(default=250, ge=20, le=800)
start_date: str | None = Field(default=None)
end_date: str | None = Field(default=None)
@model_validator(mode="after")
def _check_data_source(self) -> OptimizeBacktestRequest:
if self.ohlcv is None and self.symbol is None:
raise ValueError("必须提供 ohlcv 或 symbol 之一")
return self
class OptimizeAllBacktestRequest(BaseModel):
"""一键寻优所有策略请求。
在单个标的上,对所有策略的预设参数网格(见
``easy_tdx.backtest.strategies.presets.STRATEGY_PRESETS``)依次做网格寻优,
取各策略最优点汇总成全局排名,找出最佳策略 + 参数组合。数据来源与单标的
回测一致(ohlcv 内联或 symbol 取行情)。
"""
cash: float = Field(default=1_000_000.0, gt=0)
commission: float = Field(default=0.0003, ge=0, le=0.01)
slippage: float = Field(default=0.0, ge=0, le=0.05)
execution: Literal["next_open", "next_close"] = Field(default="next_open")
workers: int = Field(
default=1,
ge=0,
le=32,
description="一键寻优并发工作进程数:0 或 1=串行(默认);2+=ProcessPoolExecutor "
"并行(CPU-bound,线程无加速)。推荐 min(cpu_count, 8)。",
)
# 数据来源 A:内联 OHLCV
ohlcv: list[dict[str, Any]] | None = Field(default=None, max_length=2000)
# 数据来源 B:按标的取行情
symbol: str | None = Field(default=None, pattern=r"^(SZ|SH|BJ):\d{6}$")
category: Literal["DAY", "WEEK", "MONTH", "MIN_5", "MIN_15", "MIN_30", "MIN_60"] = Field(
default="DAY"
)
count: int = Field(default=250, ge=20, le=800)
start_date: str | None = Field(default=None)
end_date: str | None = Field(default=None)
@model_validator(mode="after")
def _check_data_source(self) -> OptimizeAllBacktestRequest:
if self.ohlcv is None and self.symbol is None:
raise ValueError("必须提供 ohlcv 或 symbol 之一")
return self
class MultiSeedRequest(BaseModel):
"""多 seed 随机抽样验证 + 晋级门槛请求(v1.25)。
在给定股票池上对策略做多次(不同 seed)随机抽样回测,输出跨样本稳定性
(正收益比例 / 平均夏普 / 各 seed 稳定性列)与四项晋级门槛判定。
"""
strategy: str = Field(..., description="策略名(见 /backtest/strategies")
params: dict[str, Any] = Field(default_factory=dict, description="策略参数")
stocks: list[str] = Field(
...,
min_length=2,
max_length=50,
description='股票池,格式 "市场:代码",如 ["SZ:000001","SH:600519"]',
)
category: Literal["DAY", "WEEK", "MONTH", "MIN_5", "MIN_15", "MIN_30", "MIN_60"] = Field(
default="DAY"
)
count: int = Field(default=500, ge=60, le=2000, description="每标的 K 线根数")
n_seeds: int = Field(default=3, ge=1, le=10, description="随机种子数")
sample_size: int | None = Field(
default=None, ge=1, le=50, description="每 seed 抽样标的数(None=全池)"
)
cash: float = Field(default=1_000_000.0, gt=0)
commission: float = Field(default=0.0003, ge=0, le=0.01)
min_commission: float = Field(default=5.0, ge=0)
stamp_tax: float = Field(default=0.001, ge=0, le=0.01)
slippage: float = Field(default=0.0, ge=0, le=0.05)
execution: Literal["next_open", "next_close"] = Field(default="next_open")
auto_fees: bool = Field(default=True, description="按各标的品种计费(默认开)")
gates: dict[str, float] | None = Field(
default=None,
description="晋级门槛覆盖(positive_ratio/mean_sharpe/mean_trades/mean_return",
)
class RotationBacktestRequest(BaseModel):
"""轮动组合回测请求(v1.27)。
按打分排名定期换仓:固定槽位等额、跌出排名自动卖出补位。
打分支持内置动量或通达信公式数值输出。
"""
stocks: list[str] = Field(
...,
min_length=2,
max_length=50,
description='股票池,格式 "市场:代码",如 ["SZ:000001","SH:600519"]',
)
score: Literal["momentum", "formula"] = Field(
default="momentum", description="打分方式:momentum(内置动量)或 formula(公式数值输出)"
)
period: int = Field(default=20, ge=2, le=250, description="动量回看周期(score=momentum 时)")
formula_text: str | None = Field(
default=None, max_length=8000, description="通达信公式(score=formula 时)"
)
score_col: str | None = Field(default=None, description="公式数值输出列名(默认最后一个)")
slots: int = Field(default=5, ge=1, le=20, description="持仓槽位数")
refresh: Literal["daily", "weekly", "monthly"] = Field(default="weekly", description="调仓频率")
category: Literal["DAY", "WEEK", "MONTH", "MIN_5", "MIN_15", "MIN_30", "MIN_60"] = Field(
default="DAY"
)
count: int = Field(default=500, ge=60, le=2000, description="每标的 K 线根数")
cash: float = Field(default=1_000_000.0, gt=0)
commission: float = Field(default=0.0003, ge=0, le=0.01)
min_commission: float = Field(default=5.0, ge=0)
stamp_tax: float = Field(default=0.001, ge=0, le=0.01)
stop_loss: float | None = Field(default=None, gt=0, le=0.9, description="槽内止损比例")
take_profit: float | None = Field(default=None, gt=0, le=10.0, description="槽内止盈比例")
# ── 响应模型 ───────────────────────────────────────────────────────────────────
class StrategySchemaResponse(BaseModel):
"""策略 schema 列表响应(/backtest/strategies)。"""
strategies: list[dict[str, Any]]
count: int
class BacktestResultResponse(BaseModel):
"""回测结果响应(已清洗为 JSON 原生类型)。"""
performance: dict[str, Any]
equity_curve: list[dict[str, Any]]
trades: list[dict[str, Any]]
positions: list[dict[str, Any]]
config: dict[str, Any]
# v1.25:评级(S-D,不看收益)与综合评分(0-100,含收益权重)后端输出
grade: dict[str, Any] | None = None
score: dict[str, Any] | None = None
class TaskSubmitResponse(BaseModel):
"""后台任务提交响应。"""
task_id: str
status: Literal["pending", "running"]
class TaskStateResponse(BaseModel):
"""后台任务状态响应。"""
task_id: str
status: Literal["pending", "running", "done", "failed"]
result: dict[str, Any] | None = None
error: str | None = None
description: str = ""
elapsed: float = 0.0
class TaskSummary(BaseModel):
"""任务摘要(列表用,不含完整 result)。"""
task_id: str
status: Literal["pending", "running", "done", "failed"]
description: str = ""
created_at: float = 0.0
elapsed: float = 0.0
class TaskListResponse(BaseModel):
"""任务摘要列表响应。"""
tasks: list[TaskSummary]
count: int
class OptimizeAllRankEntry(BaseModel):
"""一键寻优全局排名单行:某策略的最优点摘要。"""
strategy: str
strategy_label: str
params: dict[str, Any]
total_return: float = 0.0
sharpe: float = 0.0
max_drawdown: float = 0.0
total_trades: int = 0
win_rate: float = 0.0
profit_factor: float = 0.0
grid_points: int = 0 # 该策略本轮寻优的网格点数
class OptimizeAllResult(BaseModel):
"""一键寻优所有策略的结果:全局排名 + 最佳 + 各策略最优点。"""
ranking: list[OptimizeAllRankEntry] # 按 total_return 降序
best: OptimizeAllRankEntry | None = None
per_strategy: dict[str, OptimizeAllRankEntry] = {} # 策略名 → 最优点
total_grid_points: int = 0 # 所有策略网格点合计
# 买入持有基准(同区间/同费率/同资金;含 total_return 等 6 项指标),
# 供前端在全局最佳旁直观对比「策略 vs 买入不动」。老版本结果可能缺省。
buy_hold: dict[str, Any] | None = None
# ── 已保存策略(策略库 / StrategyLibrary)───────────────────────────────────────
class SavedStrategyCreate(BaseModel):
"""新建一条已保存策略的请求体。
前端在单标的/组合回测结果区点「保存策略」时提交。``strategy`` + ``params``
是回测引擎可直接消费的最小复现形态;``context`` 记录当时测的标的/日期,
``snapshot`` 记录保存时的关键绩效指标("为什么觉得它好")。
"""
name: str = Field(..., min_length=1, max_length=120, description="策略名称(用户自拟)")
kind: Literal["single", "portfolio", "multi"] = Field(
..., description="来源:单标的/组合/多策略组合"
)
strategy: str = Field(..., description="策略名(注册表 key,如 ma_cross")
strategy_label: str = Field(default="", description="策略展示名")
params: dict[str, Any] = Field(default_factory=dict)
context: dict[str, Any] = Field(
default_factory=dict,
description="标的上下文:single 存 symbol/category/start_date/end_date"
"portfolio 存 stocks 列表;multi 存 items(MultiStrategyItem[]) + cash/execution",
)
trade_config: dict[str, Any] = Field(
default_factory=dict, description="资金与成本配置(cash/commission/..."
)
snapshot: dict[str, Any] = Field(
default_factory=dict, description="保存时的成绩快照(total_return/sharpe/..."
)
tags: list[str] = Field(default_factory=list)
notes: str = Field(default="", max_length=2000)
class SavedStrategy(BaseModel):
"""一条已保存策略(响应模型,含 id 与时间戳)。"""
id: str
name: str
kind: Literal["single", "portfolio", "multi"]
strategy: str
strategy_label: str = ""
params: dict[str, Any] = {}
context: dict[str, Any] = {}
trade_config: dict[str, Any] = {}
snapshot: dict[str, Any] = {}
tags: list[str] = []
notes: str = ""
created_at: str = ""
updated_at: str = ""
app_version: str = ""
class SavedStrategyListResponse(BaseModel):
"""策略库列表响应。"""
strategies: list[SavedStrategy]
count: int
# ── 多策略组合回测(资金分仓 / 并行制)──────────────────────────────────────────
class MultiStrategyItem(BaseModel):
"""多策略组合回测的单条策略槽位。
每条 = 一个策略 + 它的参数 + 它要跑的原标的 + 日期范围。资金由请求体的
``cash`` 统一给出,引擎按策略数均分到各条。
"""
strategy: str = Field(..., description="策略名(注册表 key,如 ma_cross")
strategy_label: str = Field(default="", description="策略展示名(用于结果 key")
params: dict[str, Any] = Field(default_factory=dict)
symbol: str = Field(
...,
pattern=r"^(SZ|SH|BJ):\d{6}$",
description='标的完整代码,格式 "市场:6位代码",如 "SH:601088"',
)
category: Literal["DAY", "WEEK", "MONTH", "MIN_5", "MIN_15", "MIN_30", "MIN_60"] = Field(
default="DAY"
)
start_date: str | None = Field(default=None, description="开始日期 YYYY-MM-DD(可选过滤)")
end_date: str | None = Field(default=None, description="结束日期 YYYY-MM-DD(可选过滤)")
class MultiStrategyBacktestRequest(BaseModel):
"""多策略组合回测请求(资金分仓)。
勾选 N 个策略,各跑在各自原标的上,总资金按策略数均分。响应该请求的后台任务
结果是 ``MultiStrategyResult``(结构同 ``PortfolioResult``,前端复用组合页图表)。
"""
items: list[MultiStrategyItem] = Field(
..., min_length=1, max_length=20, description="策略槽位列表(1~20 条)"
)
cash: float = Field(default=1_000_000.0, gt=0, description="组合总资金(均分给各策略)")
commission: float = Field(default=0.0003, ge=0, le=0.01)
min_commission: float = Field(default=5.0, ge=0)
stamp_tax: float = Field(default=0.001, ge=0, le=0.01)
slippage: float = Field(default=0.0, ge=0, le=0.05)
execution: Literal["next_open", "next_close"] = Field(default="next_open")
# ── 信号雷达(一键扫描已保存策略的最近买卖信号)────────────────────────────────
class SignalScanRequest(BaseModel):
"""信号扫描请求:扫描策略库全部已保存策略,只看最近 N 根 K 线内的信号。"""
window_bars: int = Field(
default=5, ge=1, le=30, description="检查最近 N 根 K 线内的信号(日线即 N 个交易日)"
)
class SignalScanRecentSignal(BaseModel):
"""窗口内单根 K 线的信号。"""
date: str
direction: Literal["BUY", "SELL"]
class SignalScanRow(BaseModel):
"""扫描结果单行:一个"策略×标的"子任务的信号摘要。
single 策略 1 行;portfolio 每只标的 1 行;multi 每个子策略 1 行
(行内 ``strategy_name`` 是所属已保存策略的名字)。
"""
strategy_id: str
strategy_name: str
kind: Literal["single", "portfolio", "multi"]
strategy: str
strategy_label: str = ""
params: dict[str, Any] = {}
symbol: str
category: str = "DAY"
latest_signal: Literal["BUY", "SELL"] | None = None # 窗口内最后一根有信号的 K 线
signal_date: str | None = None # 该信号所在 K 线日期
recent_signals: list[SignalScanRecentSignal] = [] # 窗口内全部信号(按时间正序)
position: Literal["holding", "flat"] | None = None # 扫描结束时策略仓位
last_close: float | None = None
last_bar_date: str | None = None
error: str | None = None
class SignalScanResult(BaseModel):
"""信号扫描结果:全部子任务行 + 汇总计数。"""
rows: list[SignalScanRow]
total: int = 0
buy_count: int = 0 # 窗口内有买入信号的行数
sell_count: int = 0 # 窗口内有卖出信号的行数
error_count: int = 0
elapsed: float = 0.0
# ── 结果序列化 ─────────────────────────────────────────────────────────────────
def _clean_value(v: Any) -> Any:
"""把单个值转为 JSON 原生类型(递归处理容器)。"""
# bool 必须先于 int 判断(bool 是 int 子类)
if isinstance(v, bool):
return v
if isinstance(v, int | str):
return v
if isinstance(v, float):
# NaN/Inf 无法 JSON 序列化,转 None
import math
return None if math.isnan(v) or math.isinf(v) else v
if v is None:
return v
if hasattr(v, "isoformat"):
return v.isoformat()
if hasattr(v, "item"):
# numpy scalar → Python native(递归清洗,处理 numpy NaN
return _clean_value(v.item())
if isinstance(v, dict):
return {str(k): _clean_value(val) for k, val in v.items()}
if isinstance(v, list | tuple):
return [_clean_value(item) for item in v]
# 兜底:不可序列化的对象转字符串
return str(v)
def serialize_result(result: Any) -> dict[str, Any]:
"""把 BacktestResult(或其 to_dict())清洗为纯 JSON 兼容字典。
递归处理 performance / equity_curve / trades / positions / config
把 numpy scalar 转 Python 原生、NaN/Inf 转 NoneJSON 不支持)、
datetime 转 ISO 字符串。
"""
if hasattr(result, "to_dict"):
result = result.to_dict()
return {str(k): _clean_value(v) for k, v in result.items()}