Files
sanguo_vnpy_v2/tests/backtest/test_cta_engine.py
T

320 lines
12 KiB
Python

"""Tests for sanguo_backtest.cta_engine module."""
# Mock vnpy and tzlocal modules before importing anything that depends on them
import sys
from unittest.mock import MagicMock
mock_tzlocal = MagicMock()
mock_tzlocal.get_localzone_name = MagicMock(return_value="UTC")
sys.modules["tzlocal"] = mock_tzlocal
sys.modules["vnpy.trader.setting"] = MagicMock()
sys.modules["vnpy.trader.constant"] = MagicMock()
sys.modules["vnpy.trader.object"] = MagicMock()
sys.modules["vnpy.trader.database"] = MagicMock()
sys.modules["vnpy_ctastrategy.backtesting"] = MagicMock()
sys.modules["empyrical"] = MagicMock()
import pytest
import json
from unittest.mock import Mock, patch, MagicMock
from datetime import datetime
from pathlib import Path
import pandas as pd
from sanguo_backtest.cta_engine import run_cta_backtest
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 = MagicMock(name="daily_df")
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 module with BacktestingEngine
mock_module = MagicMock()
mock_module.BacktestingEngine = Mock(return_value=mock_engine)
with patch.dict("sys.modules", {"vnpy_ctastrategy.backtesting": mock_module}):
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"
assert result.statistics == {
"total_return": 0.15,
"sharpe_ratio": 1.2,
"max_drawdown": -0.08,
"win_rate": 0.55
}
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 module with BacktestingEngine
mock_module = MagicMock()
mock_module.BacktestingEngine = Mock(return_value=mock_engine)
with patch.dict("sys.modules", {"vnpy_ctastrategy.backtesting": mock_module}):
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 module with BacktestingEngine
mock_module = MagicMock()
mock_module.BacktestingEngine = Mock(return_value=mock_engine)
with patch.dict("sys.modules", {"vnpy_ctastrategy.backtesting": mock_module}):
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 module with BacktestingEngine
mock_module = MagicMock()
mock_module.BacktestingEngine = Mock(return_value=mock_engine)
# 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", {
"vnpy_ctastrategy.backtesting": mock_module,
"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_module = MagicMock()
mock_module.BacktestingEngine = Mock(return_value=mock_engine)
# 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", {
"vnpy_ctastrategy.backtesting": mock_module,
"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)"