"""Tests for sanguo_backtest.cta_optimizer module.""" import pytest from unittest.mock import Mock, patch, MagicMock from sanguo_backtest.cta_optimizer import run_cta_optimization class TestRunCtaOptimization: """Test suite for run_cta_optimization function.""" def test_run_optimization_returns_results(self, temp_db_path): """Test successful optimization execution returns proper results.""" # Mock strategy class mock_strategy_class = Mock() mock_strategy_class.__name__ = "TestStrategy" # Mock engine mock_engine = MagicMock() # Mock run_optimization to return list of results (tuple format) mock_engine.run_optimization.return_value = [ ({"n": 5}, 1.0, {"sharpe_ratio": 1.0, "total_return": 0.10}), ({"n": 10}, 1.5, {"sharpe_ratio": 1.5, "total_return": 0.15}), ] # Mock OptimizationSetting mock_setting = MagicMock() mock_setting.set_target.return_value = None mock_setting.add_parameter.return_value = (True, "") # Mock config mock_cfg = Mock() # Create mock modules mock_ctastrategy = MagicMock() mock_ctastrategy.BacktestingEngine = Mock(return_value=mock_engine) mock_trader = MagicMock() mock_trader.OptimizationSetting = Mock(return_value=mock_setting) with patch.dict("sys.modules", { "vnpy_ctastrategy.backtesting": mock_ctastrategy, "vnpy.trader.optimize": mock_trader }): results = run_cta_optimization( strategy_class=mock_strategy_class, symbol="600000", grid={"n": (5, 20, 5)}, start="2024-01-01", end="2024-03-31", cfg=mock_cfg, db_path=temp_db_path, max_workers=2 ) # Verify results assert len(results) == 2 assert all(result.type == "optimize" for result in results) assert all(result.status == "done" for result in results) assert all(result.strategy == "TestStrategy" for result in results) assert all(result.symbol == "600000" for result in results) assert all(result.start == "2024-01-01" for result in results) assert all(result.end == "2024-03-31" for result in results) # Verify first result assert results[0].params == {"n": 5} assert results[0].statistics == {"sharpe_ratio": 1.0, "total_return": 0.10} # Verify second result assert results[1].params == {"n": 10} assert results[1].statistics == {"sharpe_ratio": 1.5, "total_return": 0.15} # Verify engine methods were called mock_engine.set_parameters.assert_called_once() mock_engine.add_strategy.assert_called_once_with(mock_strategy_class, {}) mock_engine.run_optimization.assert_called_once() # Verify OptimizationSetting configuration mock_setting.set_target.assert_called_once_with("sharpe_ratio") mock_setting.add_parameter.assert_called_once_with("n", 5, 20, 5) def test_run_optimization_handles_dict_format_results(self, temp_db_path): """Test optimization with dict format results (alternative format).""" mock_strategy_class = Mock() mock_strategy_class.__name__ = "DictStrategy" mock_engine = MagicMock() # Mock run_optimization to return list of dict results mock_engine.run_optimization.return_value = [ { "params": {"window": 20}, "statistics": {"sharpe_ratio": 1.2} }, { "params": {"window": 30}, "statistics": {"sharpe_ratio": 1.8} } ] mock_setting = MagicMock() mock_setting.set_target.return_value = None mock_setting.add_parameter.return_value = (True, "") mock_cfg = Mock() mock_ctastrategy = MagicMock() mock_ctastrategy.BacktestingEngine = Mock(return_value=mock_engine) mock_trader = MagicMock() mock_trader.OptimizationSetting = Mock(return_value=mock_setting) with patch.dict("sys.modules", { "vnpy_ctastrategy.backtesting": mock_ctastrategy, "vnpy.trader.optimize": mock_trader }): results = run_cta_optimization( strategy_class=mock_strategy_class, symbol="000001", grid={"window": (20, 30, 10)}, start="2024-01-01", end="2024-03-31", cfg=mock_cfg, db_path=temp_db_path, max_workers=2 ) # Verify dict format results were parsed correctly assert len(results) == 2 assert results[0].params == {"window": 20} assert results[0].statistics == {"sharpe_ratio": 1.2} assert results[1].params == {"window": 30} assert results[1].statistics == {"sharpe_ratio": 1.8} def test_optimization_handles_failure(self, temp_db_path): """Test that exceptions during optimization are handled properly.""" mock_strategy_class = Mock() mock_strategy_class.__name__ = "FailingStrategy" mock_engine = MagicMock() mock_engine.run_optimization.side_effect = RuntimeError("Optimization failed") mock_setting = MagicMock() mock_cfg = Mock() mock_ctastrategy = MagicMock() mock_ctastrategy.BacktestingEngine = Mock(return_value=mock_engine) mock_trader = MagicMock() mock_trader.OptimizationSetting = Mock(return_value=mock_setting) with patch.dict("sys.modules", { "vnpy_ctastrategy.backtesting": mock_ctastrategy, "vnpy.trader.optimize": mock_trader }): results = run_cta_optimization( strategy_class=mock_strategy_class, symbol="600000", grid={"n": (5, 20, 5)}, start="2024-01-01", end="2024-03-31", cfg=mock_cfg, db_path=temp_db_path, max_workers=2 ) # Verify failed result assert len(results) == 1 assert results[0].type == "optimize" assert results[0].status == "failed" assert results[0].strategy == "FailingStrategy" assert results[0].symbol == "600000" assert results[0].error_msg is not None assert "RuntimeError" in results[0].error_msg assert "Optimization failed" in results[0].error_msg def test_optimization_with_multiple_parameters(self, temp_db_path): """Test optimization with multiple parameters in grid.""" mock_strategy_class = Mock() mock_strategy_class.__name__ = "MultiParamStrategy" mock_engine = MagicMock() mock_engine.run_optimization.return_value = [ ({"n": 5, "window": 10}, 1.0, {"sharpe_ratio": 1.0}), ] mock_setting = MagicMock() mock_setting.set_target.return_value = None mock_setting.add_parameter.return_value = (True, "") mock_cfg = Mock() mock_ctastrategy = MagicMock() mock_ctastrategy.BacktestingEngine = Mock(return_value=mock_engine) mock_trader = MagicMock() mock_trader.OptimizationSetting = Mock(return_value=mock_setting) with patch.dict("sys.modules", { "vnpy_ctastrategy.backtesting": mock_ctastrategy, "vnpy.trader.optimize": mock_trader }): results = run_cta_optimization( strategy_class=mock_strategy_class, symbol="600000", grid={ "n": (5, 20, 5), "window": (10, 30, 10) }, start="2024-01-01", end="2024-03-31", cfg=mock_cfg, db_path=temp_db_path, max_workers=4 ) # Verify multiple parameters were added to setting assert mock_setting.add_parameter.call_count == 2 assert len(results) == 1 assert results[0].params == {"n": 5, "window": 10} def test_optimization_generates_unique_task_ids(self, temp_db_path): """Test that each optimization result gets unique task ID.""" mock_strategy_class = Mock() mock_strategy_class.__name__ = "IdStrategy" mock_engine = MagicMock() mock_engine.run_optimization.return_value = [ ({"n": 5}, 1.0, {"sharpe_ratio": 1.0}), ({"n": 10}, 1.5, {"sharpe_ratio": 1.5}), ] mock_setting = MagicMock() mock_setting.set_target.return_value = None mock_setting.add_parameter.return_value = (True, "") mock_cfg = Mock() mock_ctastrategy = MagicMock() mock_ctastrategy.BacktestingEngine = Mock(return_value=mock_engine) mock_trader = MagicMock() mock_trader.OptimizationSetting = Mock(return_value=mock_setting) with patch.dict("sys.modules", { "vnpy_ctastrategy.backtesting": mock_ctastrategy, "vnpy.trader.optimize": mock_trader }): results = run_cta_optimization( strategy_class=mock_strategy_class, symbol="600000", grid={"n": (5, 20, 5)}, start="2024-01-01", end="2024-03-31", cfg=mock_cfg, db_path=temp_db_path, max_workers=2 ) # Verify unique task IDs assert results[0].task_id != results[1].task_id assert results[0].task_id.startswith("opt_") assert results[1].task_id.startswith("opt_")