feat(backtest): cta_engine BacktestingEngine wrapper

This commit is contained in:
2026-07-06 10:57:29 +08:00
parent e95ab91526
commit a4ce3aed85
3 changed files with 297 additions and 1 deletions
+148
View File
@@ -0,0 +1,148 @@
"""CTA strategy backtesting engine wrapper using vnpy_ctastrategy.backtesting."""
import sys
import os
import traceback
import uuid
from datetime import datetime
from pathlib import Path
# Add vnpy source to path for local development
_VNPY_SRC = os.path.join(os.path.dirname(__file__), "..", "vnpy_v4.4.0")
_VNPY_SRC = os.path.abspath(_VNPY_SRC)
if _VNPY_SRC not in sys.path:
sys.path.insert(0, _VNPY_SRC)
from sanguo_backtest.result_store import BacktestResult, save_result
# Mock Exchange enum for local use (replaces vnpy.trader.constant.Exchange)
class MockExchange:
SSE = "SSE" # Shanghai Stock Exchange
SZSE = "SZSE" # Shenzhen Stock Exchange
class Exchange:
SSE = "SSE"
SZSE = "SZSE"
def __init__(self, value):
self.value = value
def __repr__(self):
return f"Exchange.{self.value}"
Exchange = MockExchange.Exchange
def guess_exchange(symbol: str) -> Exchange:
"""按代码前缀判断交易所:6/68/5x→SSE0/3/15x→SZSE"""
if symbol.startswith(("60", "68", "51", "56", "58")):
return Exchange("SSE")
if symbol.startswith(("00", "30", "15")):
return Exchange("SZSE")
return Exchange("SSE")
def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end: str, cfg, db_path: str) -> BacktestResult:
"""
Run CTA strategy backtest using vnpy_ctastrategy BacktestingEngine.
Args:
strategy_class: CTA strategy class to backtest
symbol: Stock symbol (e.g., "600000")
params: Strategy parameters dict
start: Backtest start date (YYYY-MM-DD format)
end: Backtest end date (YYYY-MM-DD format)
cfg: Configuration object (may contain data paths)
db_path: SQLite database path for saving results
Returns:
BacktestResult: Result object with backtest statistics and status
"""
# Generate unique task ID
task_id = f"cta_{uuid.uuid4().hex[:8]}"
try:
# Lazy import of BacktestingEngine (local env may not have vnpy_ctastrategy)
from vnpy_ctastrategy.backtesting import BacktestingEngine
# Build vt_symbol for A-shares
vt_symbol = f"{symbol}.{guess_exchange(symbol).value}"
# Convert date strings to datetime objects
start_dt = datetime.strptime(start, "%Y-%m-%d")
end_dt = datetime.strptime(end, "%Y-%m-%d") if end else None
# Create and configure backtesting engine
engine = BacktestingEngine()
# Set parameters with A-share specific values
engine.set_parameters(
vt_symbol=vt_symbol,
interval="1d", # Daily interval for A-shares
start=start_dt,
end=end_dt,
rate=0.001, # Commission rate (0.1% for A-shares)
slippage=0, # No slippage for simplicity
size=1, # Contract size (1 for stocks)
pricetick=0.01, # Minimum price tick (0.01 yuan for A-shares)
capital=0 # No initial capital limit
)
# Add strategy
engine.add_strategy(strategy_class, params)
# Load historical data
engine.load_data()
# Run backtesting
engine.run_backtesting()
# Calculate statistics
statistics = engine.calculate_result()
# Get daily results for equity curve
daily_results = engine.get_all_daily_results()
# Build result object
result = BacktestResult(
task_id=task_id,
type="cta",
status="done",
strategy=strategy_class.__name__,
symbol=symbol,
params=params,
start=start,
end=end,
statistics=statistics,
equity_curve=daily_results, # Simplified: store raw daily results
trades=None # Not implemented in this MVP
)
except Exception as e:
# Handle any exceptions and return failed result
error_msg = f"{type(e).__name__}: {e}\n{traceback.format_exc()}"
result = BacktestResult(
task_id=task_id,
type="cta",
status="failed",
strategy=strategy_class.__name__,
symbol=symbol,
params=params,
start=start,
end=end,
statistics={},
equity_curve=None,
trades=None,
error_msg=error_msg
)
# Save result to database
save_result(result, db_path=db_path)
return result
# Module-level reference for mocking in tests
BacktestingEngine = None # Will be set when imported inside run_cta_backtest
+1 -1
View File
@@ -3,7 +3,7 @@ import os
from pathlib import Path from pathlib import Path
# Add real vnpy source code to sys.path # Add real vnpy source code to sys.path
_VNPY_SRC = os.path.join(os.path.dirname(__file__), "..", "..", "vnpy_v4.4.0") _VNPY_SRC = os.path.join(os.path.dirname(__file__), "..", "vnpy_v4.4.0")
_VNPY_SRC = os.path.abspath(_VNPY_SRC) _VNPY_SRC = os.path.abspath(_VNPY_SRC)
if _VNPY_SRC not in sys.path: if _VNPY_SRC not in sys.path:
sys.path.insert(0, _VNPY_SRC) sys.path.insert(0, _VNPY_SRC)
+148
View File
@@ -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_")