diff --git a/pyproject.toml b/pyproject.toml index b3fe8f4..7d4c6e8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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"] diff --git a/src/easy_tdx/backtest/cli.py b/src/easy_tdx/backtest/cli.py index 6a31821..1d59d76 100644 --- a/src/easy_tdx/backtest/cli.py +++ b/src/easy_tdx/backtest/cli.py @@ -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, diff --git a/src/easy_tdx/backtest/portfolio_engine.py b/src/easy_tdx/backtest/portfolio_engine.py index d684945..f2b6c36 100644 --- a/src/easy_tdx/backtest/portfolio_engine.py +++ b/src/easy_tdx/backtest/portfolio_engine.py @@ -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) diff --git a/src/easy_tdx/backtest/strategies/__init__.py b/src/easy_tdx/backtest/strategies/__init__.py new file mode 100644 index 0000000..89d2288 --- /dev/null +++ b/src/easy_tdx/backtest/strategies/__init__.py @@ -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() diff --git a/src/easy_tdx/backtest/strategies/builtin.py b/src/easy_tdx/backtest/strategies/builtin.py new file mode 100644 index 0000000..98cc8d8 --- /dev/null +++ b/src/easy_tdx/backtest/strategies/builtin.py @@ -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() diff --git a/src/easy_tdx/backtest/strategies/registry.py b/src/easy_tdx/backtest/strategies/registry.py new file mode 100644 index 0000000..4bd1d4b --- /dev/null +++ b/src/easy_tdx/backtest/strategies/registry.py @@ -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) diff --git a/src/easy_tdx/backtest/strategy.py b/src/easy_tdx/backtest/strategy.py index 0ad45e6..818de25 100644 --- a/src/easy_tdx/backtest/strategy.py +++ b/src/easy_tdx/backtest/strategy.py @@ -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) diff --git a/src/easy_tdx/web/app.py b/src/easy_tdx/web/app.py index 3245efb..a17c499 100644 --- a/src/easy_tdx/web/app.py +++ b/src/easy_tdx/web/app.py @@ -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 diff --git a/src/easy_tdx/web/backtest_schemas.py b/src/easy_tdx/web/backtest_schemas.py new file mode 100644 index 0000000..3b43819 --- /dev/null +++ b/src/easy_tdx/web/backtest_schemas.py @@ -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()} diff --git a/src/easy_tdx/web/routers/backtest.py b/src/easy_tdx/web/routers/backtest.py new file mode 100644 index 0000000..55139b7 --- /dev/null +++ b/src/easy_tdx/web/routers/backtest.py @@ -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() diff --git a/src/easy_tdx/web/task_runner.py b/src/easy_tdx/web/task_runner.py new file mode 100644 index 0000000..c752825 --- /dev/null +++ b/src/easy_tdx/web/task_runner.py @@ -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 diff --git a/tests/unit/test_portfolio_engine.py b/tests/unit/test_portfolio_engine.py index ad5e9cc..372f398 100644 --- a/tests/unit/test_portfolio_engine.py +++ b/tests/unit/test_portfolio_engine.py @@ -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", + } diff --git a/tests/unit/test_web_backtest.py b/tests/unit/test_web_backtest.py new file mode 100644 index 0000000..9cb0e27 --- /dev/null +++ b/tests/unit/test_web_backtest.py @@ -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"