"""CTA strategy parameter optimization wrapper using vnpy_ctastrategy.backtesting.""" import sys import os import traceback import uuid from datetime import datetime from pathlib import Path from typing import List, Any # 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 from sanguo_backtest.cta_engine import guess_exchange, Exchange def run_cta_optimization( strategy_class, symbol: str, grid: dict, start: str, end: str, cfg, db_path: str, max_workers: int = 2 ) -> List[BacktestResult]: """ Run CTA strategy parameter optimization using vnpy_ctastrategy BacktestingEngine. Args: strategy_class: CTA strategy class to optimize symbol: Stock symbol (e.g., "600000") grid: Parameter grid dict {name: (start, end, step)} 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 max_workers: Maximum number of parallel optimization workers Returns: List[BacktestResult]: List of result objects with optimization statistics """ # Generate unique task ID for this optimization run task_id = f"opt_{uuid.uuid4().hex[:8]}" try: # Lazy import of BacktestingEngine and OptimizationSetting from vnpy_ctastrategy.backtesting import BacktestingEngine from vnpy.trader.optimize import OptimizationSetting # 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 (same as cta_engine) 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 without parameters (will be set by optimization) engine.add_strategy(strategy_class, {}) # Load historical data engine.load_data() # Create optimization setting setting = OptimizationSetting() setting.set_target("sharpe_ratio") # Optimize for Sharpe ratio # Add parameter ranges to optimization setting for name, (start_val, end_val, step) in grid.items(): setting.add_parameter(name, start_val, end_val, step) # Run optimization (headless, with parallel workers) optimization_results = engine.run_optimization( setting, output=False, # Headless mode max_workers=max_workers ) # Parse optimization results and convert to BacktestResult objects results = [] for item in optimization_results: try: # Handle both tuple format (params, target_value, statistics) # and dict format ({params: ..., statistics: ...}) if isinstance(item, tuple) and len(item) >= 3: params = item[0] statistics = item[2] elif isinstance(item, dict): params = item.get("params", {}) statistics = item.get("statistics", {}) else: # Unknown format, skip this result continue # Create individual result for each optimization run result = BacktestResult( task_id=f"opt_{uuid.uuid4().hex[:8]}", # Unique ID per result type="optimize", status="done", strategy=strategy_class.__name__, symbol=symbol, params=params, start=start, end=end, statistics=statistics, equity_curve=None, # Not available in optimization results trades=None # Not available in optimization results ) results.append(result) # Save each result to database save_result(result, db_path=db_path) except Exception as e: # Handle individual result parsing error error_result = BacktestResult( task_id=f"opt_{uuid.uuid4().hex[:8]}", type="optimize", status="failed", strategy=strategy_class.__name__, symbol=symbol, params={}, start=start, end=end, statistics={}, equity_curve=None, trades=None, error_msg=f"Result parsing error: {type(e).__name__}: {e}" ) results.append(error_result) save_result(error_result, db_path=db_path) return results except Exception as e: # Handle any exceptions during optimization setup/execution error_msg = f"{type(e).__name__}: {e}\n{traceback.format_exc()}" failed_result = BacktestResult( task_id=task_id, type="optimize", status="failed", strategy=strategy_class.__name__, symbol=symbol, params={}, start=start, end=end, statistics={}, equity_curve=None, trades=None, error_msg=error_msg ) # Save failed result to database save_result(failed_result, db_path=db_path) return [failed_result] # Module-level reference for mocking in tests BacktestingEngine = None # Will be set when imported inside run_cta_optimization OptimizationSetting = None # Will be set when imported inside run_cta_optimization