mirror of
https://ghfast.top/https://github.com/aeroxw/easy_tdx_max.git
synced 2026-09-12 16:54:20 +08:00
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:
+1
-1
@@ -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"]
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
@@ -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`` 声明参数 schema(list[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/Inf(int(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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 转 None(JSON 不支持)、
|
||||
datetime 转 ISO 字符串。
|
||||
"""
|
||||
if hasattr(result, "to_dict"):
|
||||
result = result.to_dict()
|
||||
return {str(k): _clean_value(v) for k, v in result.items()}
|
||||
@@ -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()
|
||||
@@ -0,0 +1,231 @@
|
||||
"""回测后台任务执行器。
|
||||
|
||||
在进程内 ThreadPoolExecutor 上运行回测任务,通过 task_id 轮询结果。
|
||||
不引入外部依赖(Celery/Redis),适合单机部署。
|
||||
|
||||
设计取舍:
|
||||
- 回测本身是 CPU-bound(numpy/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_id(uuid4 十六进制)。
|
||||
"""
|
||||
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
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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 应被转为 None(JSON 不支持)。"""
|
||||
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 应返回 400(ValueError → 全局 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"
|
||||
Reference in New Issue
Block a user