feat(backtest): 回测 REST API + 策略注册表 + 组合回测引擎

后端回测系统完整实现:

策略注册表(backtest/strategies/):
- Param schema 声明机制,支持动态表单渲染
- 5 个内置策略:MA交叉/MACD/布林/RSI/KDJ
- 参数校验(含 NaN/Inf 拦截、范围检查、类型强制转换)

REST API(web/routers/backtest.py):
- GET  /backtest/strategies  策略枚举 + 参数 schema
- POST /backtest/run         同步回测(内联 OHLCV)
- POST /backtest/run/async   后台任务回测(含 symbol 取行情)
- POST /backtest/portfolio/run/async  组合回测(多标的)
- GET  /backtest/tasks/{id}  任务轮询

后台任务执行器(web/task_runner.py):
- ThreadPoolExecutor + 进程内 LRU 任务表
- status-aware 淘汰(不淘汰 running 任务)
- 线程安全单例 + lifespan shutdown 接入

组合回测引擎改造(portfolio_engine.py):
- 接受策略实例,参数透传到每个标的
- 新增组合净值曲线(按日期并集 forward-fill 对齐求和)

审计修复(/check 三轮):
- Param.validate 拦截 NaN/Inf/giant-int(防 DoS)
- ohlcv max_length=2000(防内存耗尽)
- LRU 淘汰跳过 running 任务(修复结果丢失竞态)
- get_runner double-checked locking(修复单例竞态)
- shutdown 接入 lifespan(修复资源泄漏)

测试:808 passed(含 39 回测路由 + 8 组合引擎 + 安全回归)
This commit is contained in:
Justin Gu
2026-07-03 03:34:34 +08:00
parent 803f2fc3ad
commit 38731114b6
13 changed files with 2224 additions and 17 deletions
+1 -1
View File
@@ -14,7 +14,7 @@ dependencies = ["pandas>=2.0,<3", "tzdata>=2024.1", "click>=8.0,<9"]
easy-tdx = "easy_tdx.cli:cli" # cli/__init__.py exposes the click group
[project.optional-dependencies]
dev = ["pytest>=8.0", "pytest-asyncio>=0.23", "pytest-cov", "mypy>=1.9", "ruff>=0.4", "scipy>=1.10,<1.16"]
dev = ["pytest>=8.0", "pytest-asyncio>=0.23", "pytest-cov", "mypy>=1.9", "ruff>=0.4", "scipy>=1.10,<1.16", "httpx>=0.27"]
science = ["scipy>=1.10,<1.16"]
web = ["fastapi>=0.110,<1", "uvicorn[standard]>=0.29"]
+1 -1
View File
@@ -389,7 +389,7 @@ def portfolio(
# 4. 创建引擎并运行
engine = PortfolioBacktestEngine(
strategy_cls=strategy_cls,
strategy=strategy_cls,
stocks=stock_data_list,
total_cash=cash,
allocation=allocation,
+75 -8
View File
@@ -39,11 +39,15 @@ class PortfolioResult:
total_performance: 组合整体绩效指标
individual_results: 每只标的的独立回测结果
equity_allocation: 每只标的的资金分配比例
combined_equity: 组合整体净值曲线(按日期对齐各标的求和),
列: datetime/total/drawdown/drawdown_pct。各标的独立回测日期范围
可能不同,此处按日期并集 forward-fill 对齐后求和。
"""
total_performance: dict[str, float]
individual_results: dict[str, BacktestResult]
equity_allocation: dict[str, float]
combined_equity: pd.DataFrame
def to_dict(self) -> dict[str, Any]:
"""转为可序列化字典。"""
@@ -51,6 +55,7 @@ class PortfolioResult:
"total_performance": self.total_performance,
"individual_results": {k: v.to_dict() for k, v in self.individual_results.items()},
"equity_allocation": self.equity_allocation,
"combined_equity": self.combined_equity.to_dict(orient="records"),
}
@@ -63,7 +68,7 @@ class PortfolioBacktestEngine:
用法::
engine = PortfolioBacktestEngine(
strategy_cls=MyStrategy,
strategy=MyStrategy,
stocks=[
StockData("000001", "SZ", df1),
StockData("600000", "SH", df2),
@@ -76,7 +81,7 @@ class PortfolioBacktestEngine:
def __init__(
self,
strategy_cls: type[Strategy],
strategy: Strategy | type[Strategy],
stocks: list[StockData],
total_cash: float = 200_000.0,
allocation: str = "equal",
@@ -90,12 +95,12 @@ class PortfolioBacktestEngine:
"""初始化组合回测引擎。
Args:
strategy_cls: 策略类
strategy: 策略类或已构造的策略实例。传实例时(如带参数的
ParametrizedStrategy),参数会被透传到每个标的的回测。
传类时(CLI 用法)用默认参数。
stocks: 标的列表(StockData
total_cash: 总资金
allocation: 资金分配方式
- "equal": 均等分配
- "capitalization": 按市值加权(需额外数据)
allocation: 资金分配方式(目前仅 "equal" 均等分配)
commission: 佣金率
min_commission: 最低佣金
stamp_tax: 印花税
@@ -103,7 +108,7 @@ class PortfolioBacktestEngine:
execution: 执行模式
chanlun_level: 缠论级别(可选)
"""
self._strategy_cls = strategy_cls
self._strategy = strategy
self._stocks = stocks
self._total_cash = total_cash
self._allocation = allocation
@@ -145,7 +150,7 @@ class PortfolioBacktestEngine:
cash = allocations.get(key, 0)
engine = BacktestEngine(
strategy=self._strategy_cls,
strategy=self._strategy,
cash=cash,
commission=self._commission,
min_commission=self._min_commission,
@@ -164,10 +169,14 @@ class PortfolioBacktestEngine:
total_alloc = sum(allocations.values())
equity_pct = {k: v / total_alloc if total_alloc > 0 else 0 for k, v in allocations.items()}
# 生成组合整体净值曲线(各标的按日期对齐求和)
combined_equity = self._build_combined_equity(individual_results, allocations)
return PortfolioResult(
total_performance=total_perf,
individual_results=individual_results,
equity_allocation=equity_pct,
combined_equity=combined_equity,
)
def _aggregate_performance(
@@ -204,3 +213,61 @@ class PortfolioBacktestEngine:
"total_stocks": len(results),
"total_cash": total_cash,
}
def _build_combined_equity(
self,
results: dict[str, BacktestResult],
allocations: dict[str, float],
) -> pd.DataFrame:
"""把各标的独立净值曲线按日期对齐求和,生成组合整体净值曲线。
各标的独立回测的日期范围可能不同(取数差异、停牌等),这里取所有标的
datetime 的并集,每个标的的 total 列 forward-fill 对齐到并集后求和。
Returns:
DataFrame: datetime / total / drawdown / drawdown_pct。
空结果时返回带表头的空 DataFrame。
"""
empty = pd.DataFrame(columns=["datetime", "total", "drawdown", "drawdown_pct"])
if not results:
return empty
# 收集各标的的 (datetime, total) 系列,以 datetime 为索引
series_list: list[pd.Series] = []
for key, result in results.items():
ec = result.equity_curve
if len(ec) == 0:
continue
# datetime 列可能是 int(YYYYMMDD) 或 datetime;统一转可比字符串/时间戳
dt = ec["datetime"]
if dt.dtype.kind in "iu": # int YYYYMMDD
dt = pd.to_datetime(dt.astype(str), format="%Y%m%d")
elif dt.dtype != "datetime64[ns]":
dt = pd.to_datetime(dt)
s = pd.Series(ec["total"].to_numpy(), index=dt, name=key)
series_list.append(s)
if not series_list:
return empty
# 外连接对齐(并集日期),forward-fill 各标的在缺失日期的净值(持有不动),
# 再求和得组合总净值。缺失值填 0 是为应对某标的完全无该日期数据的情况。
aligned = pd.concat(series_list, axis=1).sort_index()
aligned = aligned.ffill().fillna(0)
total = aligned.sum(axis=1)
# 计算回撤
peak = total.cummax()
drawdown = total - peak
# drawdown_pct:以初始总资金为基准(peak 的首个值),避免除零
initial = peak.iloc[0] if len(peak) > 0 and peak.iloc[0] != 0 else 1.0
drawdown_pct = drawdown / initial
return pd.DataFrame(
{
"datetime": total.index,
"total": total.to_numpy(),
"drawdown": drawdown.to_numpy(),
"drawdown_pct": drawdown_pct.to_numpy(),
}
).reset_index(drop=True)
@@ -0,0 +1,41 @@
"""内置策略包:注册表 + 预置策略。
导入本包即触发所有内置策略的注册。调用方通过 :func:`get_registry` 发现策略::
from easy_tdx.backtest.strategies import get_registry
for entry in get_registry().all():
print(entry.name, entry.label, [p.to_schema() for p in entry.params])
"""
from easy_tdx.backtest.strategies.registry import ( # noqa: F401
Param,
ParametrizedStrategy,
RegisteredStrategy,
StrategyRegistry,
get_registry,
register_strategy,
resolve,
)
__all__ = [
"Param",
"ParametrizedStrategy",
"RegisteredStrategy",
"StrategyRegistry",
"get_registry",
"register_strategy",
"resolve",
]
def _load_builtin() -> None:
"""导入内置策略模块以触发注册(惰性,避免循环导入)。"""
from easy_tdx.backtest.strategies import builtin # noqa: F401
# 触发 REF 符号的导入校验(builtin 顶部已导入,这里仅保留语义占位)
del builtin
# 模块导入时即加载内置策略,确保 get_registry() 调用前策略已注册
_load_builtin()
+177
View File
@@ -0,0 +1,177 @@
"""内置策略集合。
每个策略通过 :func:`~easy_tdx.backtest.strategies.registry.register_strategy`
登记到全局注册表,并声明参数 schema 供 Web API 表单动态渲染。
导入本模块即触发所有策略的注册。Web API / CLI 通过 ``get_registry()```
发现策略,无需手动枚举。
"""
from __future__ import annotations
from easy_tdx.backtest.strategies.registry import (
Param,
ParametrizedStrategy,
register_strategy,
)
from easy_tdx.MyTT import BOLL, CROSS, KDJ, MA, MACD, RSI
__all__: list[str] = [] # 注册副作用即可,无需导出符号
# ── 双均线交叉 ─────────────────────────────────────────────────────────────────
@register_strategy(
name="ma_cross",
label="双均线交叉",
description="快线上穿慢线买入,快线下穿慢线卖出。最经典的趋势跟随策略。",
)
class MaCrossStrategy(ParametrizedStrategy):
"""快慢均线金叉买入、死叉卖出。"""
params = [
Param("fast", int, default=5, min_value=1, max_value=60, label="快线周期"),
Param("slow", int, default=20, min_value=5, max_value=250, label="慢线周期"),
]
def init(self) -> None:
self.ma_fast = self.I(MA, self.data.close, self.p["fast"])
self.ma_slow = self.I(MA, self.data.close, self.p["slow"])
self.gold = self.I(CROSS, self.ma_fast, self.ma_slow)
self.dead = self.I(CROSS, self.ma_slow, self.ma_fast)
def next(self) -> None:
i = self._bar_index
if self.gold[i]:
self.buy()
elif self.dead[i] and self.position["size"] > 0:
self.sell()
# ── MACD 金叉 ──────────────────────────────────────────────────────────────────
@register_strategy(
name="macd",
label="MACD 金叉",
description="DIF 上穿 DEA 买入(金叉),DIF 下穿 DEA 卖出(死叉)。",
)
class MacdStrategy(ParametrizedStrategy):
"""MACD 金叉/死叉。"""
params = [
Param("short", int, default=12, min_value=2, max_value=50, label="短期EMA"),
Param("long", int, default=26, min_value=5, max_value=100, label="长期EMA"),
Param("signal", int, default=9, min_value=2, max_value=50, label="信号周期"),
]
def init(self) -> None:
self.dif, self.dea, self._hist = self.I(
MACD,
self.data.close,
self.p["short"],
self.p["long"],
self.p["signal"],
)
self.gold = self.I(CROSS, self.dif, self.dea)
self.dead = self.I(CROSS, self.dea, self.dif)
def next(self) -> None:
i = self._bar_index
if self.gold[i]:
self.buy()
elif self.dead[i] and self.position["size"] > 0:
self.sell()
# ── 布林带突破 ─────────────────────────────────────────────────────────────────
@register_strategy(
name="boll_breakout",
label="布林带突破",
description="收盘价突破下轨买入,突破上轨卖出(均值回归思路)。",
)
class BollBreakoutStrategy(ParametrizedStrategy):
"""价格触及下轨买入、触及上轨卖出。"""
params = [
Param("n", int, default=20, min_value=5, max_value=100, label="周期"),
Param("p", float, default=2.0, min_value=0.5, max_value=4.0, label="标准差倍数"),
]
def init(self) -> None:
self.upper, self.mid, self.lower = self.I(BOLL, self.data.close, self.p["n"], self.p["p"])
def next(self) -> None:
i = self._bar_index
close = self.data.close[0]
# 触及下轨买入(均值回归);触及上轨获利了结
if close <= self.lower[i] and self.position["size"] == 0:
self.buy()
elif close >= self.upper[i] and self.position["size"] > 0:
self.sell()
# ── RSI 超买超卖 ───────────────────────────────────────────────────────────────
@register_strategy(
name="rsi_reversal",
label="RSI 超卖反弹",
description="RSI 低于超卖线买入,RSI 高于超买线卖出。",
)
class RsiReversalStrategy(ParametrizedStrategy):
"""RSI 超卖买入、超买卖出。"""
params = [
Param("n", int, default=14, min_value=2, max_value=50, label="RSI周期"),
Param("oversold", int, default=30, min_value=5, max_value=45, label="超卖线"),
Param("overbought", int, default=70, min_value=55, max_value=95, label="超买线"),
]
def init(self) -> None:
self.rsi = self.I(RSI, self.data.close, self.p["n"])
def next(self) -> None:
i = self._bar_index
rsi = self.rsi[i]
if rsi <= self.p["oversold"] and self.position["size"] == 0:
self.buy()
elif rsi >= self.p["overbought"] and self.position["size"] > 0:
self.sell()
# ── KDJ 金叉 ───────────────────────────────────────────────────────────────────
@register_strategy(
name="kdj_cross",
label="KDJ 金叉",
description="K 线上穿 D 线买入(金叉),K 线下穿 D 线卖出(死叉)。",
)
class KdjCrossStrategy(ParametrizedStrategy):
"""KDJ K/D 金叉死叉。"""
params = [
Param("n", int, default=9, min_value=2, max_value=30, label="RSV周期"),
]
def init(self) -> None:
self.k, self.d, self._j = self.I(
KDJ,
self.data.close,
self.data.high,
self.data.low,
self.p["n"],
)
self.gold = self.I(CROSS, self.k, self.d)
self.dead = self.I(CROSS, self.d, self.k)
def next(self) -> None:
i = self._bar_index
if self.gold[i]:
self.buy()
elif self.dead[i] and self.position["size"] > 0:
self.sell()
@@ -0,0 +1,298 @@
"""内置策略注册表。
提供策略的发现、参数自描述和实例化机制,供 Web API 表单动态渲染、
CLI 枚举和回测路由使用。
设计要点:
- 每个策略通过类属性 ``params`` 声明参数 schemalist[Param])。
- ``@register_strategy`` 装饰器读取 ``params`` 与元信息登记到全局 registry。
- 策略类继承 :class:`ParametrizedStrategy```__init__`` 接受 kwargs 并
做类型/范围校验,校验通过后存入 ``self.params`` 供 ``init()`` 使用。
- 不修改基类 :class:`~easy_tdx.backtest.strategy.Strategy` 的无参 ``__init__``
契约——``ParametrizedStrategy`` 是单独的 mixin 子类。
示例::
from easy_tdx.MyTT import MA, crossover
from easy_tdx.backtest.strategies import ParametrizedStrategy, Param, register_strategy
@register_strategy(name="ma_cross", label="双均线交叉", description="...")
class MaCrossStrategy(ParametrizedStrategy):
params = [
Param("fast", int, default=5, min_value=1, max_value=60, label="快线周期"),
Param("slow", int, default=20, min_value=5, max_value=250, label="慢线周期"),
]
def init(self) -> None:
self.ma_fast = self.I(MA, self.data.close, self.p["fast"])
self.ma_slow = self.I(MA, self.data.close, self.p["slow"])
self.cross = crossover(self.ma_fast, self.ma_slow)
def next(self) -> None:
if self.cross[self._bar_index]:
self.buy()
elif self.position["size"] > 0:
self.sell()
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, ClassVar, Literal
from easy_tdx.backtest.strategy import Strategy
if TYPE_CHECKING:
from collections.abc import Callable
__all__ = [
"Param",
"ParametrizedStrategy",
"RegisteredStrategy",
"StrategyRegistry",
"get_registry",
"register_strategy",
"resolve",
]
# ── 参数 schema ───────────────────────────────────────────────────────────────
ParamType = Literal["int", "float", "bool", "str"]
def _type_name(tp: type) -> ParamType:
"""把 Python 类型映射到 schema 类型字符串。"""
if tp is int:
return "int"
if tp is float:
return "float"
if tp is bool:
return "bool"
return "str"
@dataclass(frozen=True)
class Param:
"""单个策略参数的 schema 描述。
Attributes:
name: 参数名(与 ``__init__`` 关键字一致,存入 ``self.params`` 的键)。
type: 参数类型(int/float/bool/str)。
default: 默认值。
min_value: 数值型下限(含),None 表示不限。
max_value: 数值型上限(含),None 表示不限。
choices: 字符串型可选值;非空时限制取值集合。
label: 前端展示用的中文标签。
description: 参数说明(可选)。
"""
name: str
type: type
default: Any = None
min_value: float | None = None
max_value: float | None = None
choices: tuple[str, ...] | None = None
label: str = ""
description: str = ""
def to_schema(self) -> dict[str, Any]:
"""序列化为 JSON 兼容的 schema 字典(供前端表单渲染)。"""
schema: dict[str, Any] = {
"name": self.name,
"type": _type_name(self.type),
"default": self.default,
"label": self.label or self.name,
}
if self.min_value is not None:
schema["min_value"] = self.min_value
if self.max_value is not None:
schema["max_value"] = self.max_value
if self.choices:
schema["choices"] = list(self.choices)
if self.description:
schema["description"] = self.description
return schema
def validate(self, value: Any) -> Any:
"""校验并强制转换取值,不合法抛 ValueError。
NaN/Inf 的 float 输入会被拒绝(NaN 绕过所有比较,Inf 仅在有界时被拦,
显式 isfinite 检查堵住两者)。int(inf)/int(nan) 会抛 OverflowError/
ValueError,一并捕获。
Returns:
转换后的值(类型与 ``self.type`` 一致)。
"""
import math
try:
if self.type is bool:
# bool 必须先于 int 判断(bool 是 int 子类)
if isinstance(value, bool):
converted: Any = value
else:
converted = str(value).strip().lower() in {"1", "true", "yes", "on"}
elif self.type is int:
# 先拒绝 float 的 NaN/Infint(inf) 会抛 OverflowError 不在 ValueError 内)
if isinstance(value, float) and not math.isfinite(value):
raise ValueError(f"参数 '{self.name}' 不接受 NaN/Inf")
converted = int(value)
elif self.type is float:
converted = float(value)
# float 的 NaN/Inf 必须显式拦:NaN 比较恒 False 会绕过边界检查
if not math.isfinite(converted):
raise ValueError(f"参数 '{self.name}' 不接受 NaN/Inf")
else:
converted = str(value)
except (TypeError, ValueError, OverflowError) as exc:
raise ValueError(
f"参数 '{self.name}' 期望 {self.type.__name__},得到 {value!r}"
) from exc
if self.choices is not None and self.type is str and converted not in self.choices:
raise ValueError(
f"参数 '{self.name}' 取值 {converted!r} 不在可选范围 {list(self.choices)}"
)
if self.type in (int, float) and self.min_value is not None and converted < self.min_value:
raise ValueError(f"参数 '{self.name}'={converted} 小于下限 {self.min_value}")
if self.type in (int, float) and self.max_value is not None and converted > self.max_value:
raise ValueError(f"参数 '{self.name}'={converted} 大于上限 {self.max_value}")
return converted
# ── 可参数化策略基类 ───────────────────────────────────────────────────────────
class ParametrizedStrategy(Strategy):
"""支持 kwargs 注入与校验的策略基类。
子类声明类属性 ``params: ClassVar[list[Param]]``(参数 schema),
``__init__`` 接受同名关键字参数(缺省取 ``Param.default``),校验后存入
实例属性 ``self.p`` 字典(已解析的参数值)。
策略代码中用 ``self.p["fast"]`` 访问参数值;注册表用 ``cls.params``
读取 schema。两者类型不同(list vs dict),通过 ClassVar 隔离。
"""
# 类属性:参数 schema(子类覆盖)。ClassVar 表明这是类级配置而非实例字段。
params: ClassVar[list[Param]] = []
# 实例属性:已校验的参数值字典(init/next 中通过 self.p[name] 访问)。
p: dict[str, Any]
def __init__(self, **kwargs: Any) -> None:
"""从 kwargs 构造策略参数。
多余的未知参数抛 ValueError;缺失参数取默认值。
"""
super().__init__()
declared = {param.name: param for param in self.params}
unknown = set(kwargs) - set(declared)
if unknown:
raise ValueError(f"策略 {type(self).__name__} 收到未知参数: {sorted(unknown)}")
resolved: dict[str, Any] = {}
for name, param in declared.items():
raw = kwargs.get(name, param.default)
resolved[name] = param.validate(raw)
self.p = resolved
# ── 注册表 ─────────────────────────────────────────────────────────────────────
@dataclass
class RegisteredStrategy:
"""注册表条目。"""
name: str
label: str
description: str
strategy_cls: type[ParametrizedStrategy]
params: list[Param] = field(default_factory=list)
def to_schema(self) -> dict[str, Any]:
"""序列化为 JSON 兼容的策略描述(供前端策略下拉框 + 参数表单)。"""
return {
"name": self.name,
"label": self.label,
"description": self.description,
"params": [p.to_schema() for p in self.params],
}
def build(self, params: dict[str, Any] | None = None) -> ParametrizedStrategy:
"""用给定参数构造策略实例,缺失参数取默认值。"""
return self.strategy_cls(**(params or {}))
class StrategyRegistry:
"""内置策略注册表(全局单例)。"""
def __init__(self) -> None:
self._strategies: dict[str, RegisteredStrategy] = {}
def register(
self,
strategy_cls: type[ParametrizedStrategy],
*,
name: str,
label: str = "",
description: str = "",
) -> type[ParametrizedStrategy]:
"""登记一个策略类。重复 name 抛 ValueError。"""
if name in self._strategies:
raise ValueError(f"策略名 '{name}' 已注册")
params = list(getattr(strategy_cls, "params", []))
self._strategies[name] = RegisteredStrategy(
name=name,
label=label or name,
description=description,
strategy_cls=strategy_cls,
params=params,
)
return strategy_cls
def get(self, name: str) -> RegisteredStrategy:
"""按 name 取策略,不存在抛 KeyError。"""
try:
return self._strategies[name]
except KeyError:
raise KeyError(f"未知策略 '{name}',可选: {sorted(self._strategies)}") from None
def all(self) -> list[RegisteredStrategy]:
"""返回所有已注册策略(按注册顺序)。"""
return list(self._strategies.values())
def names(self) -> list[str]:
"""返回所有策略名。"""
return sorted(self._strategies)
# ── 全局单例与便捷接口 ─────────────────────────────────────────────────────────
_REGISTRY = StrategyRegistry()
def get_registry() -> StrategyRegistry:
"""获取全局策略注册表单例。"""
return _REGISTRY
def register_strategy(
*,
name: str,
label: str = "",
description: str = "",
) -> Callable[[type[ParametrizedStrategy]], type[ParametrizedStrategy]]:
"""类装饰器:把 ParametrizedStrategy 子类登记到全局注册表。"""
def decorator(cls: type[ParametrizedStrategy]) -> type[ParametrizedStrategy]:
_REGISTRY.register(cls, name=name, label=label, description=description)
return cls
return decorator
def resolve(name: str) -> RegisteredStrategy:
"""便捷函数:按 name 解析策略。"""
return _REGISTRY.get(name)
+4 -3
View File
@@ -251,12 +251,13 @@ class Strategy(ABC):
# ── 指标注册 ───────────────────────────────────────────────────────────────
def I( # noqa: E743
self, func: Callable[..., NDArray], *args: Any, **kwargs: Any
) -> NDArray:
self, func: Callable[..., Any], *args: Any, **kwargs: Any
) -> Any:
"""注册指标。
将 _SeriesAccessor 参数自动解包为 numpy 数组,调用 func 计算指标。
返回的数组存入 self._indicators,可在 next() 中访问。
返回结果(通常是 ndarray,但 MACD/KDJ 等返回 tuple)存入
``self._indicators``,可在 ``next()`` 中访问。
Args:
func: 指标函数(如 MyTT.MA
+15
View File
@@ -79,6 +79,18 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]:
except Exception:
pass
# --- 关闭回测任务执行器(取消 pending,等待 running ---
# shutdown 是同步阻塞调用,包在 to_thread 里避免阻塞 event loop
import asyncio
try:
from easy_tdx.web.task_runner import shutdown_runner
await asyncio.to_thread(shutdown_runner)
logger.info("Backtest task runner shutdown")
except Exception:
logger.warning("Backtest task runner shutdown failed", exc_info=True)
def _create_app(
host: str | None = None,
@@ -141,6 +153,7 @@ def _create_app(
# Mount routers
from easy_tdx.web.routers.announcement import router as announcement_router
from easy_tdx.web.routers.backtest import router as backtest_router
from easy_tdx.web.routers.bars import router as bars_router
from easy_tdx.web.routers.block import router as block_router
from easy_tdx.web.routers.board_mac import router as board_mac_router
@@ -172,5 +185,7 @@ def _create_app(
app.include_router(announcement_router, prefix="/api/v1")
# 新浪财报三表路由(独立数据源)
app.include_router(sina_router, prefix="/api/v1")
# 回测路由(纯计算,不依赖行情连接 lifespan)
app.include_router(backtest_router, prefix="/api/v1")
return app
+186
View File
@@ -0,0 +1,186 @@
"""回测 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",
"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=100000.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", "this_close", "worst", "best"] = Field(
default="next_open", description="成交模式"
)
# 数据来源 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=200000.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", "this_close", "worst", "best"] = Field(
default="next_open"
)
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 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]
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
# ── 结果序列化 ─────────────────────────────────────────────────────────────────
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()}
+301
View File
@@ -0,0 +1,301 @@
"""回测路由:策略枚举、同步回测、后台任务回测、任务轮询。
设计要点:
- 回测是纯计算(不依赖行情连接的 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,
PortfolioBacktestRequest,
StrategySchemaResponse,
TaskStateResponse,
TaskSubmitResponse,
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/{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.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)
# ── 内部实现 ───────────────────────────────────────────────────────────────────
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,
)
result = engine.run(df)
return serialize_result(result)
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
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 df
def _run_portfolio_backtest(
stock_data_list: list[Any], req: PortfolioBacktestRequest
) -> dict[str, Any]:
"""执行组合回测并返回清洗后的结果字典(后台线程内调用)。"""
from easy_tdx.backtest.portfolio_engine import PortfolioBacktestEngine
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,
)
result = engine.run()
return serialize_result(result)
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 内调用)。
单个标的取数失败时跳过(不中断整个组合),全部失败返回空列表。
"""
from easy_tdx.backtest.portfolio_engine import StockData
from easy_tdx.web.convert import category_from_str, market_from_str
stock_data_list: list[StockData] = []
for symbol in stocks:
market_str, code = symbol.split(":", 1)
try:
df = await client.get_security_bars(
market_from_str(market_str),
code,
category_from_str(category),
0,
800, # 固定拉满,前端按日期过滤
)
except Exception:
continue # 单个标的失败不中断
if len(df) < 2:
continue
# 列名归一化:日线返回 date,分钟线返回 datetime
if "datetime" not in df.columns and "date" in df.columns:
df = df.copy()
df["datetime"] = df["date"]
# 日期范围过滤
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
def _now() -> float:
"""获取当前时间戳(隔离 import,便于测试)。"""
import time
return time.time()
+231
View File
@@ -0,0 +1,231 @@
"""回测后台任务执行器。
在进程内 ThreadPoolExecutor 上运行回测任务,通过 task_id 轮询结果。
不引入外部依赖(Celery/Redis),适合单机部署。
设计取舍:
- 回测本身是 CPU-boundnumpy/pandas 持 GIL),线程池主要价值是**不阻塞
FastAPI 的 asyncio event loop**——回测在独立线程跑,HTTP handler 立即返回。
- 任务结果保留在进程内存,带 LRU 上限(默认 100),重启即丢。MVP 可接受;
若需持久化历史,未来再加 SQLite。
并发正确性要点(task_runner.py 审计修复):
- ``submit`` 把「注册 task state」与「提交 executor」放在**同一把锁**内,
避免任务在拿到 future 前就被淘汰(Finding: submit-future 竞态窗口)。
- ``_run`` 写状态时**不假设** ``self._tasks[task_id]`` 仍在表中——并发淘汰
可能在任务运行期间移除其条目。``move_to_end`` 用 try/except 容忍,状态
写到本地 ``state`` 引用(即使被淘汰也无害,GC 回收)。
- ``_evict_if_needed_locked`` 跳过 ``running`` 状态的任务——正在执行的任务
恰好是 OrderedDict 头部(完成时才 move_to_end),盲目 FIFO 淘汰会优先
杀掉在途任务。淘汰改用「最旧的 non-running 条目」。
"""
from __future__ import annotations
import logging
import time
import uuid
from collections import OrderedDict
from collections.abc import Callable
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, field
from threading import Lock
from typing import Any, Literal
logger = logging.getLogger(__name__)
TaskStatus = Literal["pending", "running", "done", "failed"]
# 结果表上限:超过后丢弃最旧的已完成/失败任务(LRU)。running 任务不会被淘汰。
# 每个全市场组合回测结果约几百 KB,100 条上限内存占用 < 100 MB。
_MAX_RESULTS = 100
@dataclass
class TaskState:
"""单个回测任务的状态快照。"""
task_id: str
status: TaskStatus = "pending"
result: dict[str, Any] | None = None
error: str | None = None
created_at: float = field(default_factory=time.time)
started_at: float | None = None
finished_at: float | None = None
# 供前端展示的描述(策略名 + 标的等),不参与业务逻辑
description: str = ""
def to_dict(self) -> dict[str, Any]:
"""序列化为 JSON 兼容字典(供轮询接口返回)。"""
return {
"task_id": self.task_id,
"status": self.status,
"result": self.result,
"error": self.error,
"created_at": self.created_at,
"started_at": self.started_at,
"finished_at": self.finished_at,
"description": self.description,
"elapsed": (self.finished_at or time.time()) - (self.started_at or self.created_at),
}
class BacktestTaskRunner:
"""回测任务执行器(全局单例)。
Thread-safety: 所有对 ``self._tasks`` 的读写都在 ``self._lock`` 内。
worker 线程的 ``_run`` 写状态时持锁,但写的是本地 ``state`` 引用——
即使条目已被并发淘汰,写入也无副作用(写入的对象不再被 ``_tasks`` 引用,
GC 回收)。``move_to_end`` 容忍 KeyError。
"""
def __init__(self, max_workers: int = 4, max_results: int = _MAX_RESULTS) -> None:
self._executor = ThreadPoolExecutor(
max_workers=max_workers,
thread_name_prefix="backtest-worker",
)
self._tasks: OrderedDict[str, TaskState] = OrderedDict()
self._lock = Lock()
self._max_results = max_results
self._shutdown = False
def submit(
self,
func: Callable[[], dict[str, Any]],
*,
description: str = "",
) -> str:
"""提交一个回测任务,立即返回 task_id。
注册 task state 与提交 executor 在同一把锁内,避免任务在拿到 future
前就被并发淘汰。
Args:
func: 无参可调用,返回 ``dict[str, Any]`` 形式的回测结果。
内部异常会被捕获并记入 ``TaskState.error``。
description: 任务描述(前端展示用)。
Returns:
task_iduuid4 十六进制)。
"""
task_id = uuid.uuid4().hex
with self._lock:
if self._shutdown:
raise RuntimeError("任务执行器已关闭,拒绝提交")
self._tasks[task_id] = TaskState(task_id=task_id, description=description)
# 提交 executor 也在锁内——避免「注册后被淘汰再提交」的窗口。
# executor.submit 本身很快(入队即返回),不会显著持锁。
self._executor.submit(self._run, task_id, func)
self._evict_if_needed_locked()
return task_id
def get(self, task_id: str) -> TaskState:
"""取任务状态,不存在抛 KeyError。"""
with self._lock:
if task_id not in self._tasks:
raise KeyError(f"未知任务 '{task_id}'")
return self._tasks[task_id]
def peek(self, task_id: str) -> TaskState | None:
"""取任务状态,不存在返回 None(不抛异常,便于轮询)。"""
with self._lock:
return self._tasks.get(task_id)
def status(self, task_id: str) -> TaskStatus | None:
"""取任务状态字符串,不存在返回 None。"""
state = self.peek(task_id)
return state.status if state else None
def shutdown(self, wait: bool = True) -> None:
"""关闭执行器(应用退出时调用)。
取消排队中的 pending 任务,等待 running 任务完成(wait=True)。
之后再 submit 会抛 RuntimeError。
"""
with self._lock:
if self._shutdown:
return
self._shutdown = True
self._executor.shutdown(wait=wait, cancel_futures=True)
# ── 内部实现 ───────────────────────────────────────────────────────────────
def _run(self, task_id: str, func: Callable[[], dict[str, Any]]) -> None:
"""在工作线程内执行:更新状态、跑任务、捕获异常。
状态写入用本地 ``state`` 引用,不假设条目仍在 ``self._tasks`` 中——
并发淘汰可能在任务运行期间移除条目。``move_to_end`` 容忍 KeyError。
"""
# 取本地引用;若已被淘汰则静默退出(无副作用)
with self._lock:
state = self._tasks.get(task_id)
if state is None:
logger.warning("任务 %s 在执行前已被淘汰,跳过", task_id)
return
state.status = "running"
state.started_at = time.time()
try:
result = func()
with self._lock:
# 即使被淘汰也写到本地 state(无害),move_to_end 容忍缺失
state.result = result
state.status = "done"
state.finished_at = time.time()
try:
self._tasks.move_to_end(task_id)
except KeyError:
pass # 已被淘汰,无需移动
except Exception as exc: # noqa: BLE001 — 故意宽口径,任务级兜底
logger.exception("回测任务 %s 失败", task_id)
with self._lock:
state.error = f"{type(exc).__name__}: {exc}"
state.status = "failed"
state.finished_at = time.time()
try:
self._tasks.move_to_end(task_id)
except KeyError:
pass
def _evict_if_needed_locked(self) -> None:
"""超过上限时丢弃最旧的 non-running 任务(调用方需持锁)。
running 任务不会被淘汰(它们恰在 OrderedDict 头部,但盲淘汰会杀在途任务)。
只淘汰 pending/done/failed 中最旧者。
"""
while len(self._tasks) > self._max_results:
# 找第一个 non-running 条目淘汰;若无则停止(全在 running,不强制淘汰)
evict_id: str | None = None
for tid, st in self._tasks.items():
if st.status != "running":
evict_id = tid
break
if evict_id is None:
break # 全部 running,暂时无法淘汰
self._tasks.pop(evict_id, None)
# ── 全局单例 ───────────────────────────────────────────────────────────────────
_RUNNER: BacktestTaskRunner | None = None
_RUNNER_LOCK = Lock()
def get_runner() -> BacktestTaskRunner:
"""获取全局回测任务执行器单例(惰性初始化,线程安全)。"""
global _RUNNER # noqa: PLW0603 — 模块级单例
if _RUNNER is None:
with _RUNNER_LOCK:
# double-checked locking:拿到锁后再确认一次,避免重复创建
if _RUNNER is None:
_RUNNER = BacktestTaskRunner()
return _RUNNER
def shutdown_runner() -> None:
"""关闭全局执行器(应用退出时调用,幂等)。"""
global _RUNNER # noqa: PLW0603
with _RUNNER_LOCK:
if _RUNNER is not None:
_RUNNER.shutdown()
_RUNNER = None
+85 -4
View File
@@ -58,7 +58,7 @@ class TestPortfolioBacktest:
]
engine = PortfolioBacktestEngine(
strategy_cls=SimpleBuyStrategy,
strategy=SimpleBuyStrategy,
stocks=stocks,
total_cash=200000,
)
@@ -77,7 +77,7 @@ class TestPortfolioBacktest:
]
engine = PortfolioBacktestEngine(
strategy_cls=SimpleBuyStrategy,
strategy=SimpleBuyStrategy,
stocks=stocks,
total_cash=100000,
allocation="equal",
@@ -91,7 +91,7 @@ class TestPortfolioBacktest:
def test_empty_stocks(self) -> None:
"""空标的列表应返回零绩效."""
engine = PortfolioBacktestEngine(
strategy_cls=SimpleBuyStrategy,
strategy=SimpleBuyStrategy,
stocks=[],
total_cash=100000,
)
@@ -104,7 +104,7 @@ class TestPortfolioBacktest:
"""结果应可序列化为字典."""
stocks = [StockData("000001", "SZ", _make_df(100))]
engine = PortfolioBacktestEngine(
strategy_cls=SimpleBuyStrategy,
strategy=SimpleBuyStrategy,
stocks=stocks,
total_cash=100000,
)
@@ -114,3 +114,84 @@ class TestPortfolioBacktest:
assert "total_performance" in d
assert "individual_results" in d
assert "equity_allocation" in d
assert "combined_equity" in d
class TestStrategyInstanceParams:
"""测试策略实例(带参数)的透传——Phase 3 引擎改造的核心."""
def test_strategy_instance_params_passed_through(self) -> None:
"""传策略实例时,参数应透传到每个标的(而非用默认值)."""
from easy_tdx.backtest.strategies import get_registry
entry = get_registry().get("ma_cross")
strategy_instance = entry.build({"fast": 10, "slow": 30})
stocks = [
StockData("000001", "SZ", _make_df(120, seed=1)),
StockData("000002", "SZ", _make_df(120, seed=2)),
]
engine = PortfolioBacktestEngine(
strategy=strategy_instance,
stocks=stocks,
total_cash=200000,
)
result = engine.run()
assert len(result.individual_results) == 2
for res in result.individual_results.values():
assert res.performance is not None
class TestCombinedEquity:
"""测试组合净值曲线生成."""
def test_combined_equity_generated(self) -> None:
"""组合净值曲线应生成且含 total/drawdown/drawdown_pct 列."""
stocks = [
StockData("000001", "SZ", _make_df(100, seed=42)),
StockData("600000", "SH", _make_df(100, seed=99)),
]
engine = PortfolioBacktestEngine(
strategy=SimpleBuyStrategy,
stocks=stocks,
total_cash=200000,
)
result = engine.run()
assert len(result.combined_equity) > 0
cols = set(result.combined_equity.columns)
assert {"datetime", "total", "drawdown", "drawdown_pct"} <= cols
assert result.combined_equity["total"].iloc[0] > 0
def test_combined_equity_date_alignment(self) -> None:
"""日期范围不同的标的应正确对齐(forward-fill."""
stocks = [
StockData("000001", "SZ", _make_df(80, seed=1)),
StockData("000002", "SZ", _make_df(100, seed=2)),
]
engine = PortfolioBacktestEngine(
strategy=SimpleBuyStrategy,
stocks=stocks,
total_cash=200000,
)
result = engine.run()
assert len(result.combined_equity) >= 100
def test_combined_equity_empty(self) -> None:
"""空标的列表应返回空净值曲线(带表头)."""
engine = PortfolioBacktestEngine(
strategy=SimpleBuyStrategy,
stocks=[],
total_cash=100000,
)
result = engine.run()
assert len(result.combined_equity) == 0
assert set(result.combined_equity.columns) == {
"datetime",
"total",
"drawdown",
"drawdown_pct",
}
+809
View File
@@ -0,0 +1,809 @@
"""回测 Web API 测试(离线,无网络)。
覆盖:
- 策略注册表(Param schema、校验、序列化)
- 请求模型校验(数据来源二选一、参数范围)
- 结果序列化(numpy/datetime/NaN 清洗)
- 后台任务执行器(提交/轮询/失败/LRU 淘汰)
- router 端到端(策略枚举、同步回测、后台任务、错误路径)
"""
from __future__ import annotations
import time
import numpy as np
import pandas as pd
import pytest
pytest.importorskip("fastapi")
# ── 测试夹具 ───────────────────────────────────────────────────────────────────
@pytest.fixture()
def sample_ohlcv() -> list[dict[str, object]]:
"""带趋势的合成 OHLCV(确保均线策略能产生交易)。"""
np.random.seed(42)
n = 300
close = 10 + np.cumsum(np.random.randn(n) * 0.2 + 0.05)
dates = pd.date_range("2023-01-01", periods=n, freq="B")
return [
{
"datetime": d.strftime("%Y-%m-%d"),
"open": float(c - np.random.rand() * 0.1),
"high": float(c + np.random.rand() * 0.2),
"low": float(c - np.random.rand() * 0.2),
"close": float(c),
"vol": float(np.random.randint(1000, 10000)),
"amount": float(c * 5000),
}
for d, c in zip(dates, close, strict=True)
]
# ---------------------------------------------------------------------------
# 策略注册表
# ---------------------------------------------------------------------------
def test_registry_has_builtin_strategies():
"""导入 strategies 包后应有 5 个内置策略注册。"""
from easy_tdx.backtest.strategies import get_registry
reg = get_registry()
names = reg.names()
assert "ma_cross" in names
assert "macd" in names
assert "boll_breakout" in names
assert "rsi_reversal" in names
assert "kdj_cross" in names
assert len(names) >= 5
def test_strategy_schema_serialization():
"""策略 schema 应含 name/label/description/params 列表。"""
from easy_tdx.backtest.strategies import get_registry
entry = get_registry().get("ma_cross")
schema = entry.to_schema()
assert schema["name"] == "ma_cross"
assert schema["label"] == "双均线交叉"
assert isinstance(schema["description"], str)
param_names = [p["name"] for p in schema["params"]]
assert param_names == ["fast", "slow"]
def test_strategy_build_with_default_params():
"""无参数构造应使用默认值。"""
from easy_tdx.backtest.strategies import get_registry
entry = get_registry().get("ma_cross")
inst = entry.build()
assert inst.p["fast"] == 5
assert inst.p["slow"] == 20
def test_strategy_build_with_custom_params():
"""自定义参数应覆盖默认值。"""
from easy_tdx.backtest.strategies import get_registry
inst = get_registry().get("ma_cross").build({"fast": 10, "slow": 30})
assert inst.p["fast"] == 10
assert inst.p["slow"] == 30
def test_strategy_rejects_unknown_param():
"""未知参数应抛 ValueError。"""
from easy_tdx.backtest.strategies import get_registry
with pytest.raises(ValueError, match="未知参数"):
get_registry().get("ma_cross").build({"foo": 1})
def test_strategy_rejects_out_of_range_param():
"""超出范围的参数应抛 ValueError。"""
from easy_tdx.backtest.strategies import get_registry
with pytest.raises(ValueError, match="上限"):
get_registry().get("ma_cross").build({"fast": 999})
def test_strategy_rejects_unknown_name():
"""未知策略名应抛 KeyError。"""
from easy_tdx.backtest.strategies import get_registry
with pytest.raises(KeyError, match="未知策略"):
get_registry().get("not_a_real_strategy")
def test_param_type_coercion():
"""字符串传入 int 参数应自动转换。"""
from easy_tdx.backtest.strategies import registry
p = registry.Param("n", int, default=10, min_value=1, max_value=100)
assert p.validate("15") == 15
assert p.validate(20) == 20
def test_param_bool_coercion():
"""bool 参数接受多种真值表示。"""
from easy_tdx.backtest.strategies import registry
p = registry.Param("flag", bool, default=False)
assert p.validate("true") is True
assert p.validate("0") is False
assert p.validate(True) is True
def test_param_choices_validation():
"""字符串型 choices 限制取值集合。"""
from easy_tdx.backtest.strategies import registry
p = registry.Param("mode", str, default="a", choices=("a", "b", "c"))
assert p.validate("a") == "a"
with pytest.raises(ValueError, match="可选范围"):
p.validate("z")
# ---------------------------------------------------------------------------
# 请求模型校验
# ---------------------------------------------------------------------------
def test_backtest_request_requires_data_source():
"""请求必须提供 ohlcv 或 symbol 之一。"""
from easy_tdx.web.backtest_schemas import BacktestRequest
with pytest.raises(ValueError, match="ohlcv|symbol"):
BacktestRequest(strategy="ma_cross")
def test_backtest_request_symbol_pattern():
"""symbol 必须符合 市场:代码 格式。"""
from easy_tdx.web.backtest_schemas import BacktestRequest
with pytest.raises(ValueError):
BacktestRequest(strategy="ma_cross", symbol="000001") # 缺市场前缀
with pytest.raises(ValueError):
BacktestRequest(strategy="ma_cross", symbol="XX:000001") # 非法市场
def test_backtest_request_cash_must_be_positive():
"""初始资金必须 > 0。"""
from easy_tdx.web.backtest_schemas import BacktestRequest
with pytest.raises(ValueError):
BacktestRequest(strategy="ma_cross", symbol="SZ:000001", cash=0)
def test_backtest_request_execution_enum():
"""execution 必须是预定义模式之一。"""
from easy_tdx.web.backtest_schemas import BacktestRequest
with pytest.raises(ValueError):
BacktestRequest(strategy="ma_cross", symbol="SZ:000001", execution="invalid_mode")
def test_backtest_request_defaults():
"""默认值应合理。"""
from easy_tdx.web.backtest_schemas import BacktestRequest
req = BacktestRequest(strategy="ma_cross", symbol="SZ:000001")
assert req.cash == 100000.0
assert req.commission == 0.0003
assert req.execution == "next_open"
assert req.category == "DAY"
assert req.count == 250
# ---------------------------------------------------------------------------
# 结果序列化
# ---------------------------------------------------------------------------
def test_serialize_result_cleans_numpy():
"""numpy scalar 应被转为 Python 原生类型。"""
from easy_tdx.web.backtest_schemas import serialize_result
fake = {
"performance": {"total_return": np.float64(0.5), "sharpe": np.int32(3)},
"equity_curve": [{"equity": np.float64(100.0)}],
"trades": [],
"positions": [],
"config": {},
}
out = serialize_result(fake)
assert isinstance(out["performance"]["total_return"], float)
assert isinstance(out["performance"]["sharpe"], int)
assert isinstance(out["equity_curve"][0]["equity"], float)
def test_serialize_result_cleans_nan_inf():
"""NaN/Inf 应被转为 NoneJSON 不支持)。"""
from easy_tdx.web.backtest_schemas import serialize_result
fake = {
"performance": {"sharpe": float("nan"), "sortino": float("inf")},
"equity_curve": [],
"trades": [],
"positions": [],
"config": {},
}
out = serialize_result(fake)
assert out["performance"]["sharpe"] is None
assert out["performance"]["sortino"] is None
def test_serialize_result_cleans_datetime():
"""datetime/timestamp 应被转为 ISO 字符串。"""
from easy_tdx.web.backtest_schemas import serialize_result
ts = pd.Timestamp("2024-01-15")
fake = {
"performance": {},
"equity_curve": [{"date": ts}],
"trades": [],
"positions": [],
"config": {"end_date": ts},
}
out = serialize_result(fake)
assert isinstance(out["equity_curve"][0]["date"], str)
assert "2024-01-15" in out["equity_curve"][0]["date"]
# ---------------------------------------------------------------------------
# 后台任务执行器
# ---------------------------------------------------------------------------
def test_task_runner_submit_and_poll():
"""提交任务后应能轮询到 done 状态与结果。"""
from easy_tdx.web.task_runner import BacktestTaskRunner
runner = BacktestTaskRunner(max_workers=2)
def work() -> dict[str, object]:
return {"ok": True, "value": 42}
task_id = runner.submit(work, description="test")
# 轮询直到完成
for _ in range(100):
state = runner.get(task_id)
if state.status in ("done", "failed"):
break
time.sleep(0.05)
assert state.status == "done"
assert state.result == {"ok": True, "value": 42}
assert state.error is None
assert state.started_at is not None
assert state.finished_at is not None
runner.shutdown()
def test_task_runner_captures_failure():
"""任务内异常应记入 error,状态变 failed。"""
from easy_tdx.web.task_runner import BacktestTaskRunner
runner = BacktestTaskRunner(max_workers=1)
def boom() -> dict[str, object]:
raise RuntimeError("boom")
task_id = runner.submit(boom)
for _ in range(100):
state = runner.get(task_id)
if state.status in ("done", "failed"):
break
time.sleep(0.05)
assert state.status == "failed"
assert "RuntimeError" in (state.error or "")
assert "boom" in (state.error or "")
runner.shutdown()
def test_task_runner_lru_eviction():
"""超过上限应丢弃最旧任务。"""
from easy_tdx.web.task_runner import BacktestTaskRunner
runner = BacktestTaskRunner(max_workers=1, max_results=3)
ids = [runner.submit(lambda: {"i": i}, description=f"t{i}") for i in range(5)]
# 等待全部完成
for _ in range(200):
if all(runner.peek(tid) and runner.peek(tid).status in ("done", "failed") for tid in ids):
break
time.sleep(0.02)
# 前 2 个应被淘汰
assert runner.peek(ids[0]) is None
assert runner.peek(ids[1]) is None
# 后 3 个保留
assert runner.peek(ids[2]) is not None
assert runner.peek(ids[4]) is not None
runner.shutdown()
def test_task_runner_unknown_task_raises():
"""查询不存在的 task_id 应抛 KeyError。"""
from easy_tdx.web.task_runner import BacktestTaskRunner
runner = BacktestTaskRunner(max_workers=1)
with pytest.raises(KeyError, match="未知任务"):
runner.get("nonexistent")
assert runner.peek("nonexistent") is None
assert runner.status("nonexistent") is None
runner.shutdown()
def test_task_runner_does_not_evict_running(monkeypatch):
"""LRU 淘汰应跳过 running 任务——这是审计修复的关键并发正确性回归。
若回退此修复,正在执行的长任务会被新提交的任务淘汰,导致 worker 在
move_to_end 时抛 KeyError、任务结果丢失。
"""
import easy_tdx.web.task_runner as tr_mod
runner = tr_mod.BacktestTaskRunner(max_workers=1, max_results=3)
# 用 Event 钉住第一个任务使其保持 running
release = __import__("threading").Event()
started = __import__("threading").Event()
def slow_task() -> dict[str, object]:
started.set()
release.wait(timeout=5)
return {"slow": True}
running_id = runner.submit(slow_task, description="slow")
# 等它确实进入 running
assert started.wait(timeout=2), "慢任务未启动"
# 提交足够多的新任务触发淘汰(max_results=3
other_ids = [runner.submit(lambda: {"i": i}, description=f"t{i}") for i in range(5)]
for _ in range(200):
if all(runner.peek(t) and runner.peek(t).status in ("done", "failed") for t in other_ids):
break
time.sleep(0.02)
# 关键断言:running 任务绝不能被淘汰
assert runner.peek(running_id) is not None, "running 任务被错误淘汰!"
assert runner.peek(running_id).status == "running"
release.set()
for _ in range(100):
if runner.peek(running_id) and runner.peek(running_id).status in ("done", "failed"):
break
time.sleep(0.02)
assert runner.peek(running_id).status == "done"
runner.shutdown()
def test_task_runner_shutdown_is_idempotent_and_rejects_submit():
"""shutdown 应幂等,关闭后再 submit 应报错。"""
from easy_tdx.web.task_runner import BacktestTaskRunner
runner = BacktestTaskRunner(max_workers=1)
runner.shutdown()
# 幂等:再次 shutdown 不抛
runner.shutdown()
# 关闭后 submit 应被拒
with pytest.raises(RuntimeError, match="已关闭"):
runner.submit(lambda: {})
def test_task_runner_worker_tolerates_eviction():
"""worker 即使在运行期间被(理论上的)淘汰也不应抛未捕获异常。
直接验证 _run 的 move_to_end 容忍 KeyError:手动从 _tasks 移除条目后
触发完成路径。
"""
from easy_tdx.web.task_runner import BacktestTaskRunner
runner = BacktestTaskRunner(max_workers=1, max_results=10)
completed = __import__("threading").Event()
def quick() -> dict[str, object]:
return {"ok": True}
task_id = runner.submit(quick)
for _ in range(100):
if runner.peek(task_id) and runner.peek(task_id).status == "done":
break
time.sleep(0.01)
# 任务完成后手动移除(模拟并发淘汰),再提交新任务不应触发任何 worker 异常
with runner._lock:
runner._tasks.pop(task_id, None)
# 新任务正常工作
tid2 = runner.submit(quick)
for _ in range(100):
if runner.peek(tid2) and runner.peek(tid2).status == "done":
break
time.sleep(0.01)
assert runner.peek(tid2).status == "done"
completed.set()
runner.shutdown()
# ---------------------------------------------------------------------------
# 审计修复回归:参数校验安全(NaN/Inf/OverflowError+ ohlcv 上限
# ---------------------------------------------------------------------------
def test_param_rejects_float_nan():
"""float 参数的 NaN 必须被拦(NaN 绕过比较,原实现漏网)。"""
from easy_tdx.backtest.strategies.registry import Param
p = Param("p", float, default=2.0, min_value=0.5, max_value=4.0)
with pytest.raises(ValueError, match="NaN|Inf|期望"):
p.validate(float("nan"))
def test_param_rejects_float_inf():
"""float 参数的 Inf 必须被拦。"""
from easy_tdx.backtest.strategies.registry import Param
p = Param("p", float, default=2.0, min_value=0.5, max_value=4.0)
with pytest.raises(ValueError):
p.validate(float("inf"))
def test_param_rejects_int_from_inf_without_overflow():
"""int(inf) 应报 ValueError(→400)而非 OverflowError(→500)。"""
from easy_tdx.backtest.strategies.registry import Param
p = Param("n", int, default=10, min_value=1, max_value=100)
with pytest.raises(ValueError): # 不能是 OverflowError
p.validate(float("inf"))
def test_param_rejects_int_from_nan():
"""int(nan) 应报 ValueError(→400)而非 ValueError 逃逸。"""
from easy_tdx.backtest.strategies.registry import Param
p = Param("n", int, default=10, min_value=1, max_value=100)
with pytest.raises(ValueError):
p.validate(float("nan"))
def test_param_normal_values_still_work():
"""安全加固不应误伤正常值。"""
from easy_tdx.backtest.strategies.registry import Param
assert Param("n", int, default=10, min_value=1, max_value=100).validate(50) == 50
assert Param("p", float, default=2.0, min_value=0.5, max_value=4.0).validate(2.0) == 2.0
# 字符串数字仍可转换
assert Param("n", int, default=10, min_value=1, max_value=100).validate("15") == 15
def test_backtest_request_ohlcv_max_length():
"""ohlcv 内联数据必须有上限(防 DoS)。"""
from easy_tdx.web.backtest_schemas import BacktestRequest
bar = {
"datetime": "2024-01-01",
"open": 1,
"high": 1,
"low": 1,
"close": 1,
"vol": 1,
"amount": 1,
}
# 2000 条(上限)应接受
req = BacktestRequest(strategy="ma_cross", ohlcv=[bar] * 2000)
assert len(req.ohlcv) == 2000
# 2001 条应拒绝
with pytest.raises(ValueError):
BacktestRequest(strategy="ma_cross", ohlcv=[bar] * 2001)
# ---------------------------------------------------------------------------
# Router 端到端(TestClient
# ---------------------------------------------------------------------------
@pytest.fixture()
def client():
"""FastAPI TestClient(不触发真实行情连接——lifespan 连接失败不影响回测路由)。"""
from fastapi.testclient import TestClient
from easy_tdx.web import create_app
app = create_app()
with TestClient(app) as c:
yield c
def test_list_strategies_endpoint(client):
"""GET /backtest/strategies 返回策略列表。"""
resp = client.get("/api/v1/backtest/strategies")
assert resp.status_code == 200
body = resp.json()
assert body["count"] >= 5
names = [s["name"] for s in body["strategies"]]
assert "ma_cross" in names
# 每个策略的 schema 结构完整
for s in body["strategies"]:
assert "label" in s
assert "params" in s
assert isinstance(s["params"], list)
def test_sync_backtest_with_inline_data(client, sample_ohlcv):
"""POST /backtest/run 用内联数据同步回测应返回完整结果。"""
resp = client.post(
"/api/v1/backtest/run",
json={
"strategy": "ma_cross",
"params": {"fast": 5, "slow": 20},
"cash": 100000,
"ohlcv": sample_ohlcv,
},
)
assert resp.status_code == 200, resp.text
body = resp.json()
# 结果结构完整
assert "performance" in body
assert "equity_curve" in body
assert "trades" in body
assert "positions" in body
assert "config" in body
# 绩效指标存在且为原生类型
assert "total_return" in body["performance"]
assert isinstance(body["performance"]["total_return"], (int, float))
# 净值曲线有数据
assert len(body["equity_curve"]) > 0
def test_sync_backtest_rejects_missing_ohlcv(client):
"""同步回测不给 ohlcv 应返回 400。"""
resp = client.post(
"/api/v1/backtest/run",
json={"strategy": "ma_cross", "symbol": "SZ:000001"},
)
assert resp.status_code == 400
def test_sync_backtest_rejects_bad_strategy(client, sample_ohlcv):
"""未知策略名应返回 400。"""
resp = client.post(
"/api/v1/backtest/run",
json={"strategy": "nope", "ohlcv": sample_ohlcv},
)
assert resp.status_code == 400
def test_sync_backtest_rejects_bad_params(client, sample_ohlcv):
"""非法参数应返回 400。"""
resp = client.post(
"/api/v1/backtest/run",
json={
"strategy": "ma_cross",
"params": {"fast": 999},
"ohlcv": sample_ohlcv,
},
)
assert resp.status_code == 400
def test_async_backtest_with_inline_data(client, sample_ohlcv):
"""POST /backtest/run/async 提交后台任务,轮询到 done。"""
resp = client.post(
"/api/v1/backtest/run/async",
json={
"strategy": "macd",
"params": {"short": 12, "long": 26, "signal": 9},
"ohlcv": sample_ohlcv,
},
)
assert resp.status_code == 202, resp.text
task_id = resp.json()["task_id"]
# 轮询
final = None
for _ in range(200):
poll = client.get(f"/api/v1/backtest/tasks/{task_id}")
assert poll.status_code == 200
final = poll.json()
if final["status"] in ("done", "failed"):
break
time.sleep(0.05)
assert final is not None
assert final["status"] == "done", final
assert final["result"] is not None
assert "performance" in final["result"]
def test_task_poll_unknown_id(client):
"""轮询不存在的 task_id 应返回 400ValueError → 全局 handler)。"""
resp = client.get("/api/v1/backtest/tasks/nonexistent")
assert resp.status_code == 400
assert "未知任务" in resp.json()["detail"]
def test_async_backtest_failure_recorded(client, monkeypatch):
"""后台任务内的异常应被记录为 failed 状态(而非 500 到客户端)。
通过 monkeypatch 让回测执行函数抛错,验证 task_runner 的异常捕获在
端到端链路里也生效。单元层 test_task_runner_captures_failure 已覆盖纯逻辑。
"""
import easy_tdx.web.routers.backtest as bt_router
def boom(_df, _req): # noqa: ANN001
raise RuntimeError("simulated backtest failure")
monkeypatch.setattr(bt_router, "_run_backtest", boom)
two_bars = [
{
"datetime": "2024-01-01",
"open": 10.0,
"high": 10.5,
"low": 9.5,
"close": 10.2,
"vol": 1000.0,
"amount": 10000.0,
},
{
"datetime": "2024-01-02",
"open": 10.2,
"high": 10.6,
"low": 10.0,
"close": 10.4,
"vol": 1200.0,
"amount": 12000.0,
},
]
resp = client.post(
"/api/v1/backtest/run/async",
json={"strategy": "ma_cross", "ohlcv": two_bars},
)
assert resp.status_code == 202
task_id = resp.json()["task_id"]
final = None
for _ in range(200):
poll = client.get(f"/api/v1/backtest/tasks/{task_id}")
final = poll.json()
if final["status"] in ("done", "failed"):
break
time.sleep(0.05)
assert final["status"] == "failed"
assert "simulated backtest failure" in (final["error"] or "")
# ---------------------------------------------------------------------------
# Phase 3: 组合回测路由
# ---------------------------------------------------------------------------
def test_portfolio_request_validates_stocks_format():
"""组合请求的 stocks 字段格式校验。"""
from easy_tdx.web.backtest_schemas import PortfolioBacktestRequest
req = PortfolioBacktestRequest(strategy="ma_cross", stocks=["SZ:000001", "SH:600519"])
assert len(req.stocks) == 2
with pytest.raises(ValueError):
PortfolioBacktestRequest(strategy="ma_cross", stocks=["000001"]) # 缺市场
with pytest.raises(ValueError):
PortfolioBacktestRequest(strategy="ma_cross", stocks=["SZ:123"]) # 非6位
with pytest.raises(ValueError):
PortfolioBacktestRequest(strategy="ma_cross", stocks=[]) # 空
with pytest.raises(ValueError):
many = [f"SZ:00000{i}" for i in range(21)]
PortfolioBacktestRequest(strategy="ma_cross", stocks=many) # >20
def test_portfolio_request_defaults():
"""组合请求默认值。"""
from easy_tdx.web.backtest_schemas import PortfolioBacktestRequest
req = PortfolioBacktestRequest(strategy="ma_cross", stocks=["SZ:000001"])
assert req.cash == 200000.0
assert req.category == "DAY"
def test_portfolio_backtest_endpoint(client, monkeypatch):
"""POST /backtest/portfolio/run/async 端到端(mock 行情取数)。"""
import pandas as pd
import easy_tdx.web.routers.backtest as bt_router
from easy_tdx.backtest.portfolio_engine import StockData
np.random.seed(3)
async def fake_fetch(client_arg, stocks, category, start, end): # noqa: ANN001
result = []
for sym in stocks:
mkt, code = sym.split(":")
n = 100
close = 10 + np.cumsum(np.random.randn(n) * 0.3 + 0.05)
df = pd.DataFrame(
{
"datetime": pd.date_range("2024-01-01", periods=n, freq="B"),
"open": close - 0.1,
"high": close + 0.2,
"low": close - 0.2,
"close": close,
"vol": np.full(n, 5000.0),
"amount": close * 5000,
}
)
result.append(StockData(code=code, market=mkt, df=df))
return result
monkeypatch.setattr(bt_router, "_fetch_portfolio_bars", fake_fetch)
resp = client.post(
"/api/v1/backtest/portfolio/run/async",
json={
"strategy": "ma_cross",
"params": {"fast": 5, "slow": 20},
"cash": 200000,
"stocks": ["SZ:000001", "SH:600519"],
},
)
assert resp.status_code == 202, resp.text
task_id = resp.json()["task_id"]
final = None
for _ in range(200):
poll = client.get(f"/api/v1/backtest/tasks/{task_id}")
final = poll.json()
if final["status"] in ("done", "failed"):
break
time.sleep(0.05)
assert final["status"] == "done", final
result = final["result"]
assert "total_performance" in result
assert "individual_results" in result
assert "equity_allocation" in result
assert "combined_equity" in result
assert result["total_performance"]["total_stocks"] == 2
assert len(result["individual_results"]) == 2
assert len(result["combined_equity"]) > 0
def test_portfolio_backtest_bad_strategy(client, monkeypatch):
"""未知策略在后台任务内抛错 → failed。"""
import pandas as pd
import easy_tdx.web.routers.backtest as bt_router
from easy_tdx.backtest.portfolio_engine import StockData
async def fake_fetch(client_arg, stocks, category, start, end): # noqa: ANN001
n = 10
df = pd.DataFrame(
{
"datetime": pd.date_range("2024-01-01", periods=n, freq="B"),
"open": np.full(n, 10.0),
"high": np.full(n, 10.5),
"low": np.full(n, 9.5),
"close": np.full(n, 10.2),
"vol": np.full(n, 1000.0),
"amount": np.full(n, 10000.0),
}
)
return [StockData(code="000001", market="SZ", df=df)]
monkeypatch.setattr(bt_router, "_fetch_portfolio_bars", fake_fetch)
resp = client.post(
"/api/v1/backtest/portfolio/run/async",
json={"strategy": "nope", "stocks": ["SZ:000001"]},
)
assert resp.status_code == 202
task_id = resp.json()["task_id"]
for _ in range(200):
poll = client.get(f"/api/v1/backtest/tasks/{task_id}")
final = poll.json()
if final["status"] in ("done", "failed"):
break
time.sleep(0.05)
assert final["status"] == "failed"