264 lines
9.5 KiB
Python
264 lines
9.5 KiB
Python
"""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_")
|