"""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