"""Tests for sanguo_backtest.cta_engine module.""" import pytest from unittest.mock import Mock, patch, MagicMock from datetime import datetime 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 mock_engine = MagicMock() mock_engine.calculate_result.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() 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_")