From a4ce3aed853526c6088aa0c866b73bb66d3f5315 Mon Sep 17 00:00:00 2001 From: claude_dev Date: Mon, 6 Jul 2026 10:57:29 +0800 Subject: [PATCH] feat(backtest): cta_engine BacktestingEngine wrapper --- sanguo_backtest/cta_engine.py | 148 ++++++++++++++++++++++++++++++ sanguo_data/datareader.py | 2 +- tests/backtest/test_cta_engine.py | 148 ++++++++++++++++++++++++++++++ 3 files changed, 297 insertions(+), 1 deletion(-) create mode 100644 sanguo_backtest/cta_engine.py create mode 100644 tests/backtest/test_cta_engine.py diff --git a/sanguo_backtest/cta_engine.py b/sanguo_backtest/cta_engine.py new file mode 100644 index 0000000..70991b2 --- /dev/null +++ b/sanguo_backtest/cta_engine.py @@ -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→SSE,0/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 \ No newline at end of file diff --git a/sanguo_data/datareader.py b/sanguo_data/datareader.py index c05585e..b7ab15b 100644 --- a/sanguo_data/datareader.py +++ b/sanguo_data/datareader.py @@ -3,7 +3,7 @@ import os from pathlib import 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) if _VNPY_SRC not in sys.path: sys.path.insert(0, _VNPY_SRC) diff --git a/tests/backtest/test_cta_engine.py b/tests/backtest/test_cta_engine.py new file mode 100644 index 0000000..a10274a --- /dev/null +++ b/tests/backtest/test_cta_engine.py @@ -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_") \ No newline at end of file