feat(backtest): cta_engine BacktestingEngine wrapper
This commit is contained in:
@@ -0,0 +1,148 @@
|
||||
"""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_")
|
||||
Reference in New Issue
Block a user