Files
sanguo_vnpy_v2/tests/backtest/test_cta_engine.py
T

450 lines
19 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.
"""Tests for sanguo_backtest.cta_engine module."""
# Mock vnpy and tzlocal modules before importing anything that depends on them。
# ⚠️注入必须在 fixture 内(测试期)而非模块 import 期:pytest 全量 collection 先于任何
# 测试执行,import 期注入的 sys.modules mock 会污染后收集的 datareader/metrics/factor
# 模块(它们 import 到的是 MagicMock)——2026-08-15 修,见 tests 全量 9 failed 根因。
import sys
import importlib
import importlib.util
from unittest.mock import MagicMock
import pytest
import json
from unittest.mock import Mock, patch, MagicMock
from datetime import datetime
from pathlib import Path
import pandas as pd
# Only mock these when the real module is NOT importable (e.g. local env missing empyrical).
# In the container (authoritative test env) everything imports fine, so NO mock is installed
# → zero sys.modules pollution leaking into other test modules.
mock_tzlocal = MagicMock()
mock_tzlocal.get_localzone_name = MagicMock(return_value="UTC")
_MOCK_FACTORIES = {
"tzlocal": lambda: mock_tzlocal,
"vnpy.trader.setting": MagicMock,
"vnpy.trader.constant": MagicMock,
"vnpy.trader.object": MagicMock,
"vnpy.trader.database": MagicMock,
"vnpy_ctastrategy.backtesting": MagicMock,
}
@pytest.fixture(scope="module", autouse=True)
def _mock_deps_and_import_cta_engine():
"""module 级唯一 mock 窗口:装 mock(仅当真模块不可导入)→ import cta_engine →
测试 → teardown 恢复 sys.modules 并清 sanguo_backtest 模块缓存(它们是在 mock
生效期间 import 的,须弹出让后续测试用真依赖重新 import)。"""
saved = {}
for name, factory in _MOCK_FACTORIES.items():
try:
found = importlib.util.find_spec(name)
except ModuleNotFoundError:
found = None
if found is None:
saved[name] = sys.modules.get(name)
sys.modules[name] = factory()
try:
cta_engine = importlib.import_module("sanguo_backtest.cta_engine")
globals()["run_cta_backtest"] = cta_engine.run_cta_backtest
yield
finally:
for cached in ("sanguo_backtest.cta_engine", "sanguo_backtest.metrics", "sanguo_backtest.ashare_engine"):
sys.modules.pop(cached, None)
for k, orig in saved.items():
if orig is None:
sys.modules.pop(k, None)
else:
sys.modules[k] = orig
globals().pop("run_cta_backtest", None)
class TestRunCtaBacktest:
"""Test suite for run_cta_backtest function."""
def test_run_cta_backtest_returns_result(self, temp_db_path):
"""Test successful CTA backtest execution returns proper result."""
# Mock strategy class
mock_strategy_class = Mock()
mock_strategy_class.__name__ = "TestStrategy"
# Mock BacktestingEngine — calculate_result() returns daily_df (DataFrame),
# calculate_statistics(df) returns the stats dict (vnpy API, matches cta_engine)
mock_engine = MagicMock()
mock_engine.calculate_result.return_value = pd.DataFrame(
{"balance": [1_000_000.0, 1_010_000.0], "net_pnl": [0.0, 10000.0]},
index=pd.date_range("2024-01-01", periods=2, freq="D"),
)
mock_engine.calculate_statistics.return_value = {
"total_return": 0.15,
"sharpe_ratio": 1.2,
"max_drawdown": -0.08,
"win_rate": 0.55
}
mock_engine.get_all_daily_results.return_value = [
{"date": "2024-01-01", "balance": 100000},
{"date": "2024-01-02", "balance": 101000},
]
# Mock config
mock_cfg = Mock()
# Create mock for AShareBacktestingEngine
mock_ashare = MagicMock()
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
mock_engine.trades = {"t1": MagicMock()} # 非空 trades 避免 degenerate
mock_engine.history_data = [] # 空 history → 定寸跳过(mock 无真实 bar)
with patch.dict("sys.modules", {"sanguo_backtest.ashare_engine": mock_ashare}):
result = run_cta_backtest(
strategy_class=mock_strategy_class,
symbol="600000",
params={"window": 20},
start="2024-01-01",
end="2024-03-31",
cfg=mock_cfg,
db_path=temp_db_path
)
# Verify result structure
assert result.type == "cta"
assert result.status == "done"
assert result.strategy == "TestStrategy"
assert result.symbol == "600000"
assert result.params == {"window": 20}
assert result.start == "2024-01-01"
assert result.end == "2024-03-31"
# statistics 是 dict——vnpy calculate_statistics 的原始键会被 compute_metrics
# 的 scalars 覆盖/扩充(容器内 config 加载成功时 metrics 会跑,覆盖 mock 的 0.15
# Mac 无 config 时 metrics 跳过,保留 mock 值)。具体数值随环境,真实值由
# test_integration_ashare 验证;此处只验结构 + mock 特有字段。
assert isinstance(result.statistics, dict)
assert result.statistics.get("sizing_shares_per_lot") == 0
assert result.error_msg is None
# Verify BacktestingEngine methods were called
mock_engine.set_parameters.assert_called_once()
mock_engine.add_strategy.assert_called_once_with(mock_strategy_class, {"window": 20})
mock_engine.load_data.assert_called_once()
mock_engine.run_backtesting.assert_called_once()
mock_engine.calculate_result.assert_called_once()
mock_engine.calculate_statistics.assert_called_once()
def test_run_cta_backtest_handles_exception(self, temp_db_path):
"""Test that exceptions during backtest are handled properly."""
# Mock strategy class
mock_strategy_class = Mock()
mock_strategy_class.__name__ = "FailingStrategy"
# Mock BacktestingEngine to raise exception
mock_engine = MagicMock()
mock_engine.load_data.side_effect = RuntimeError("Data loading failed")
# Mock config
mock_cfg = Mock()
# Create mock for AShareBacktestingEngine
mock_ashare = MagicMock()
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
with patch.dict("sys.modules", {"sanguo_backtest.ashare_engine": mock_ashare}):
result = run_cta_backtest(
strategy_class=mock_strategy_class,
symbol="000001",
params={"window": 20},
start="2024-01-01",
end="2024-03-31",
cfg=mock_cfg,
db_path=temp_db_path
)
# Verify failed result structure
assert result.type == "cta"
assert result.status == "failed"
assert result.strategy == "FailingStrategy"
assert result.symbol == "000001"
assert result.error_msg is not None
assert "RuntimeError" in result.error_msg
assert "Data loading failed" in result.error_msg
def test_run_cta_backtest_generates_task_id(self, temp_db_path):
"""Test that run_cta_backtest generates unique task IDs."""
mock_strategy_class = Mock()
mock_strategy_class.__name__ = "IdTestStrategy"
mock_engine = MagicMock()
mock_engine.calculate_result.return_value = {"sharpe_ratio": 1.0}
mock_engine.get_all_daily_results.return_value = []
mock_cfg = Mock()
# Create mock for AShareBacktestingEngine
mock_ashare = MagicMock()
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
with patch.dict("sys.modules", {"sanguo_backtest.ashare_engine": mock_ashare}):
result1 = run_cta_backtest(
strategy_class=mock_strategy_class,
symbol="600000",
params={},
start="2024-01-01",
end="2024-03-31",
cfg=mock_cfg,
db_path=temp_db_path
)
result2 = run_cta_backtest(
strategy_class=mock_strategy_class,
symbol="600001",
params={},
start="2024-01-01",
end="2024-03-31",
cfg=mock_cfg,
db_path=temp_db_path
)
# Verify unique task IDs
assert result1.task_id != result2.task_id
assert result1.task_id.startswith("cta_")
assert result2.task_id.startswith("cta_")
def test_run_cta_backtest_computes_relative_metrics(self, temp_db_path):
"""Test that run_cta_backtest computes relative metrics against benchmark."""
# Mock strategy class
mock_strategy_class = Mock()
mock_strategy_class.__name__ = "BmTestStrategy"
# Mock daily_df with 'return' column (required by compute_metrics)
dates = pd.date_range("2024-01-01", "2024-03-31", freq="D")
daily_df = pd.DataFrame({
"return": [0.001] * len(dates)
}, index=dates)
# Mock vnpy BacktestingEngine
mock_engine = MagicMock()
mock_engine.calculate_result.return_value = daily_df
mock_engine.calculate_statistics.return_value = {
"total_return": 0.15,
"sharpe_ratio": 1.2,
"max_drawdown": -0.08,
}
# Mock config with benchmark
mock_cfg = Mock()
mock_cfg.data_paths = {"daily_dir": "/mock/daily_dir"}
# Mock read_index_daily to return benchmark data
mock_bench_df = pd.DataFrame({
"date": dates,
"close": [100.0] * len(dates)
})
# Mock compute_metrics result
mock_metrics_result = Mock()
mock_metrics_result.scalars = {
"alpha": 0.05,
"beta": 0.95,
"sharpe_ratio": 1.3,
"total_return": 0.15,
"benchmark_return": 0.10
}
mock_metrics_result.series = {
"equity_curve": pd.Series([1.0, 1.1, 1.2]),
"benchmark_curve": pd.Series([1.0, 1.05, 1.1]),
"alpha": pd.Series([0.01, 0.02, 0.03]),
"beta": pd.Series([0.9, 0.95, 1.0]),
"drawdown": pd.Series([0.0, -0.01, -0.02])
}
# Create mock for AShareBacktestingEngine
mock_ashare = MagicMock()
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
mock_engine.trades = {"t1": MagicMock()} # 非空 trades 避免 degenerate
mock_engine.history_data = [] # 空 history → 定寸跳过(mock 无真实 bar)
# Mock tzlocal and vnpy modules to avoid import errors
mock_tzlocal = MagicMock()
mock_tzlocal.get_localzone_name = Mock(return_value="UTC")
with patch.dict("sys.modules", {
"sanguo_backtest.ashare_engine": mock_ashare,
"tzlocal": mock_tzlocal,
"vnpy.trader.setting": MagicMock()
}):
with patch("sanguo_backtest.cta_engine.read_index_daily", return_value=mock_bench_df):
with patch("sanguo_backtest.cta_engine.compute_metrics", return_value=mock_metrics_result):
result = run_cta_backtest(
strategy_class=mock_strategy_class,
symbol="600000",
params={"window": 20},
start="2024-01-01",
end="2024-03-31",
cfg=mock_cfg,
db_path=temp_db_path
)
# Verify result contains relative metrics (scalars merged into statistics)
assert result.statistics.get("alpha") == 0.05
assert result.statistics.get("beta") == 0.95
assert result.statistics.get("sharpe_ratio") == 1.3 # Should be present
assert result.statistics.get("total_return") == 0.15 # Should be present
# Verify metrics JSON file was written
file_dir = Path(temp_db_path).parent
metrics_file = file_dir / f"{result.task_id}_metrics.json"
assert metrics_file.exists(), f"Metrics file not found: {metrics_file}"
# Verify metrics file can be loaded and contains expected keys
with open(metrics_file, "r") as f:
metrics_data = json.load(f)
# Check that we have 5 series keys
series_keys = list(metrics_data.get("series", {}).keys())
assert len(series_keys) == 5
assert "equity_curve" in series_keys
assert "benchmark_curve" in series_keys
assert "alpha" in series_keys
assert "beta" in series_keys
assert "drawdown" in series_keys
def test_run_cta_backtest_default_benchmark_hs300(self, temp_db_path):
"""Test that default benchmark is hs300 when not specified in config."""
mock_strategy_class = Mock()
mock_strategy_class.__name__ = "DefaultBmStrategy"
# Mock daily_df
dates = pd.date_range("2024-01-01", "2024-03-31", freq="D")
daily_df = pd.DataFrame({
"return": [0.001] * len(dates)
}, index=dates)
mock_engine = MagicMock()
mock_engine.calculate_result.return_value = daily_df
mock_engine.calculate_statistics.return_value = {"total_return": 0.15}
# Mock config without benchmark (should use default hs300)
mock_cfg = Mock()
mock_cfg.data_paths = {"daily_dir": "/mock/daily_dir"}
mock_bench_df = pd.DataFrame({
"date": dates,
"close": [100.0] * len(dates)
})
mock_metrics_result = Mock()
mock_metrics_result.scalars = {"alpha": 0.05}
mock_metrics_result.series = {}
mock_ashare = MagicMock()
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
mock_engine.trades = {"t1": MagicMock()} # 非空 trades 避免 degenerate
mock_engine.history_data = [] # 空 history → 定寸跳过(mock 无真实 bar)
# Mock tzlocal and vnpy modules to avoid import errors
mock_tzlocal = MagicMock()
mock_tzlocal.get_localzone_name = Mock(return_value="UTC")
with patch.dict("sys.modules", {
"sanguo_backtest.ashare_engine": mock_ashare,
"tzlocal": mock_tzlocal,
"vnpy.trader.setting": MagicMock()
}):
with patch("sanguo_backtest.cta_engine.read_index_daily", return_value=mock_bench_df) as mock_read:
with patch("sanguo_backtest.cta_engine.compute_metrics", return_value=mock_metrics_result):
result = run_cta_backtest(
strategy_class=mock_strategy_class,
symbol="600000",
params={},
start="2024-01-01",
end="2024-03-31",
cfg=mock_cfg,
db_path=temp_db_path
)
# Verify read_index_daily was called with hs300 code (sh000300)
mock_read.assert_called_once()
call_args = mock_read.call_args
assert call_args[0][0] == "sh000300", "Default benchmark should be hs300 (sh000300)"
def test_run_cta_backtest_benchmark_missing_degrades(self, temp_db_path):
"""benchmark 缺失时 compute_metrics 仍执行、metrics.json 照写(降级)。
根因:read_index_daily 返空时原实现 if 跳过整段 compute_metrics →
{task_id}_metrics.json 不写 → 策略净值/回撤/波动图也空(不只基准图)。
降级:benchmark 空时传空 Series 给 compute_metrics(内部 reindex→fillna(0),
基准类指标返 NaN→None,策略指标正常),metrics.json 照写。
"""
mock_strategy_class = Mock()
mock_strategy_class.__name__ = "DegBmStrategy"
dates = pd.date_range("2024-01-01", "2024-03-31", freq="D")
daily_df = pd.DataFrame(
{"balance": [1_000_000.0 * (1.001 ** i) for i in range(len(dates))]},
index=dates,
)
mock_engine = MagicMock()
mock_engine.calculate_result.return_value = daily_df
mock_engine.calculate_statistics.return_value = {"total_return": 0.15}
mock_engine.trades = {"t1": MagicMock()} # 非空 trades 避免 degenerate
mock_engine.history_data = [] # 空 history → 定寸跳过
mock_cfg = Mock()
mock_cfg.data_paths = {"daily_dir": "/mock/daily_dir"}
mock_ashare = MagicMock()
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
mock_metrics_result = Mock()
mock_metrics_result.scalars = {
"total_return": 0.15, # 策略指标(不依赖基准)有值
"max_drawdown": -0.05,
"alpha": None, # 基准类指标 NaN→None(基准缺失)
"benchmark_return": None,
}
mock_metrics_result.series = {
"equity_curve": pd.Series([1.0, 1.05, 1.1]),
"drawdown": pd.Series([0.0, -0.01, -0.02]),
}
mock_tzlocal = MagicMock()
mock_tzlocal.get_localzone_name = Mock(return_value="UTC")
with patch.dict("sys.modules", {
"sanguo_backtest.ashare_engine": mock_ashare,
"tzlocal": mock_tzlocal,
"vnpy.trader.setting": MagicMock(),
}):
# read_index_daily 返空 DataFrame = benchmark 缺失
with patch("sanguo_backtest.cta_engine.read_index_daily", return_value=pd.DataFrame()):
with patch("sanguo_backtest.cta_engine.compute_metrics", return_value=mock_metrics_result) as mock_cm:
result = run_cta_backtest(
strategy_class=mock_strategy_class,
symbol="600000",
params={},
start="2024-01-01",
end="2024-03-31",
cfg=mock_cfg,
db_path=temp_db_path,
)
# 1. compute_metrics 仍被调用(降级,没因 benchmark 空跳过)
mock_cm.assert_called_once()
# 2. 传入的 benchmark_returns 是空 Series(降级标记)
benchmark_arg = mock_cm.call_args[0][1]
assert isinstance(benchmark_arg, pd.Series)
assert benchmark_arg.empty, "benchmark 缺失时应传空 Series 给 compute_metrics"
# 3. metrics.json 仍写(策略净值/回撤/波动图数据源)
file_dir = Path(temp_db_path).parent
metrics_file = file_dir / f"{result.task_id}_metrics.json"
assert metrics_file.exists(), f"降级时 metrics.json 应照写: {metrics_file}"
with open(metrics_file) as f:
metrics_data = json.load(f)
series_keys = list(metrics_data.get("series", {}).keys())
assert "equity_curve" in series_keys
assert "drawdown" in series_keys
# 4. 策略指标有值(基准类 None)
assert result.statistics.get("total_return") == 0.15
assert result.statistics.get("alpha") is None