Files
easy-tdx/tests/unit/test_portfolio_engine.py
T
Justin Gu 38731114b6 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 组合引擎 + 安全回归)
2026-07-03 03:34:34 +08:00

198 lines
6.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""单元测试:多标的组合回测引擎."""
from __future__ import annotations
import numpy as np
import pandas as pd
from easy_tdx.backtest.portfolio_engine import (
PortfolioBacktestEngine,
StockData,
)
from easy_tdx.backtest.strategy import Strategy
class SimpleBuyStrategy(Strategy):
"""简单策略:bar 5 买入,bar 30 卖出."""
def init(self) -> None:
pass
def next(self) -> None:
if self._bar_index == 5 and self.position["size"] == 0:
self.buy(size=0)
elif self._bar_index == 30 and self.position["size"] > 0:
self.sell(size=0)
def _make_df(n: int = 100, seed: int = 42) -> pd.DataFrame:
"""生成随机 OHLCV DataFrame."""
rng = np.random.default_rng(seed)
close = 100.0 + np.cumsum(rng.normal(0, 1, n))
high = close + rng.uniform(0, 1, n)
low = close - rng.uniform(0, 1, n)
open_ = low + rng.uniform(0, high - low, n)
vol = rng.integers(1000000, 10000000, n).astype(float)
return pd.DataFrame(
{
"datetime": pd.date_range("2024-01-01", periods=n, freq="D"),
"open": open_,
"high": high,
"low": low,
"close": close,
"vol": vol,
"amount": vol * close,
}
)
class TestPortfolioBacktest:
"""测试组合回测引擎."""
def test_basic_portfolio_run(self) -> None:
"""基本组合回测应正常完成."""
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 result.total_performance is not None
assert "total_return" in result.total_performance
assert len(result.individual_results) == 2
assert result.total_performance["total_stocks"] == 2
def test_equal_allocation(self) -> None:
"""均等分配:每只标的资金应为总资金/标的数."""
stocks = [
StockData("000001", "SZ", _make_df(100, seed=42)),
StockData("000002", "SZ", _make_df(100, seed=99)),
]
engine = PortfolioBacktestEngine(
strategy=SimpleBuyStrategy,
stocks=stocks,
total_cash=100000,
allocation="equal",
)
result = engine.run()
# 每只标的分配 50000
assert result.equity_allocation["SZ000001"] == 0.5
assert result.equity_allocation["SZ000002"] == 0.5
def test_empty_stocks(self) -> None:
"""空标的列表应返回零绩效."""
engine = PortfolioBacktestEngine(
strategy=SimpleBuyStrategy,
stocks=[],
total_cash=100000,
)
result = engine.run()
assert result.total_performance["total_return"] == 0.0
assert len(result.individual_results) == 0
def test_to_dict_serializable(self) -> None:
"""结果应可序列化为字典."""
stocks = [StockData("000001", "SZ", _make_df(100))]
engine = PortfolioBacktestEngine(
strategy=SimpleBuyStrategy,
stocks=stocks,
total_cash=100000,
)
result = engine.run()
d = result.to_dict()
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",
}