"""Tests for sanguo_backtest.cta_engine module.""" # Mock vnpy and tzlocal modules before importing anything that depends on them import sys import importlib.util from unittest.mock import MagicMock # 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 at collection time. # (Previous unconditional sys.modules[...]=MagicMock() here broke datareader/spike/factor/metrics # tests because it ran at collection time and never restored.) 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, "empyrical": MagicMock, } _SAVED_MODULES = {} for _name, _factory in _MOCK_FACTORIES.items(): try: _found = importlib.util.find_spec(_name) except ModuleNotFoundError: _found = None if _found is None: _SAVED_MODULES[_name] = sys.modules.get(_name) sys.modules[_name] = _factory() 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 @pytest.fixture(scope="module", autouse=True) def _restore_modules_after(): """Restore sys.modules after this module's tests, so downstream tests import the REAL vnpy/empyrical instead of the mocks we installed above. Also drop cta_engine/ metrics from the module cache (they were imported while mocks were active) so they re-import fresh with real dependencies.""" yield for cached in ("sanguo_backtest.cta_engine", "sanguo_backtest.metrics", "sanguo_backtest.ashare_engine"): sys.modules.pop(cached, None) for k, orig in _SAVED_MODULES.items(): if orig is None: sys.modules.pop(k, None) else: sys.modules[k] = orig 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