backtest_engine.py

"""
QTrading 回测引擎模块
"""

import pandas as pd
import numpy as np
from typing import List, Dict, Optional
from data.data_manager import DataManager
from strategy.strategy_manager import StrategyManager
from strategy.base_strategy import BaseStrategy
from config.config import config
from utils.logger import get_logger
from utils.common import validate_code, format_date, calculate_drawdown, today_str

logger = get_logger(__name__)

class BacktestEngine:
    """回测引擎"""

    def __init__(self, tushare_token: str = None):
        """
        初始化回测引擎

        Args:
            tushare_token: Tushare 接口令牌(可选)
        """
        self.data_manager = DataManager(tushare_token)
        self.strategy_manager = StrategyManager()
        self.logger = logger

    def backtest_strategy(
        self,
        strategy_name: str,
        adjust_method: str = "qfq",
        stock_list: List[str] = None,
        start_date: str = None,
        end_date: str = None,
        holding_period: int = 20,
        rebalance_frequency: int = 20,
        initial_capital: float = None,
        commission_rate: float = None,
        slippage: float = None
    ) -> Dict:
        """
        回测策略

        Args:
            strategy_name: 策略名称
            adjust_method: 复权方式(qfq: 前复权, hfq: 后复权, none: 不复权)
            stock_list: 股票池(可选,默认从配置文件读取)
            start_date: 开始日期(YYYYMMDD,可选)
            end_date: 结束日期(YYYYMMDD,可选)
            holding_period: 持有周期(天数,默认20)
            rebalance_frequency: 调仓频率(天数,默认20)
            initial_capital: 初始资金(可选,默认从配置读取)
            commission_rate: 佣金率(可选,默认从配置读取)
            slippage: 滑点(可选,默认从配置读取)

        Returns:
            Dict: 回测结果
        """
        if stock_list is None:
            stock_list = config.stock_pool.stock_list

        if start_date is None:
            start_date = config.stock_pool.start_date
        if end_date is None:
            end_date = today_str()

        if initial_capital is None:
            initial_capital = config.backtest.initial_capital
        if commission_rate is None:
            commission_rate = config.backtest.commission_rate
        if slippage is None:
            slippage = config.backtest.slippage

        # 创建策略实例
        strategy = self.strategy_manager.create_strategy_instance(strategy_name, adjust_method)
        if strategy is None:
            self.logger.error(f"策略 {strategy_name} 不存在或创建失败")
            return {}

        self.logger.info(f"开始回测策略: {strategy_name}")
        self.logger.info(f"回测参数:")
        self.logger.info(f"  - 复权方式: {adjust_method}")
        self.logger.info(f"  - 股票池: {len(stock_list)} 只")
        self.logger.info(f"  - 日期范围: {start_date} - {end_date}")
        self.logger.info(f"  - 持有周期: {holding_period} 天")
        self.logger.info(f"  - 调仓频率: {rebalance_frequency} 天")
        self.logger.info(f"  - 初始资金: {initial_capital}")
        self.logger.info(f"  - 佣金率: {commission_rate:.2%}")
        self.logger.info(f"  - 滑点: {slippage:.2%}")

        # 获取交易日历
        all_trade_dates = self._get_trade_dates(start_date, end_date)
        if len(all_trade_dates) < holding_period + rebalance_frequency:
            self.logger.warning("回测时间窗口过短")
            return {}

        # 回测主逻辑
        portfolio = {
            'dates': [],
            'nav': [initial_capital],
            'positions': [],
            'cash': [initial_capital],
            'total_value': [initial_capital]
        }

        # 股票历史数据缓存
        stock_data_cache = {}
        for code in stock_list:
            try:
                data = self.data_manager.calculate_adjusted_data(
                    code, adjust_method, start_date, end_date
                )
                if data is not None and not data.empty:
                    stock_data_cache[code] = data.sort_values('trade_date')
            except Exception as e:
                self.logger.warning(f"获取股票 {code} 数据失败: {e}")

        # 开始回测
        position_count = 0
        current_date = start_date
        position_count = len([code for code in stock_list if code in stock_data_cache])
        self.logger.info(f"有效股票数据: {position_count} 只")

        for i, date in enumerate(all_trade_dates):
            if i == 0:
                continue

            # 调仓逻辑
            if i % rebalance_frequency == 0:
                self._rebalance_portfolio(
                    strategy, stock_data_cache, date, portfolio, initial_capital,
                    commission_rate, slippage
                )

            # 更新持仓价值
            self._update_portfolio_value(
                stock_data_cache, date, portfolio
            )

            # 更新日期记录
            portfolio['dates'].append(date)

        # 计算回测指标
        backtest_result = self._calculate_backtest_metrics(portfolio)

        self.logger.info("回测完成")
        return backtest_result

    def _get_trade_dates(self, start_date: str, end_date: str) -> List[str]:
        """获取交易日历"""
        trade_cal = self.data_manager.tushare.get_trade_cal(start_date, end_date)
        if trade_cal is None or trade_cal.empty:
            self.logger.warning("未找到交易日历,使用连续日期")
            return []
        trade_cal = trade_cal[trade_cal['is_open'] == 1]
        return trade_cal['cal_date'].tolist()

    def _rebalance_portfolio(
        self, strategy: BaseStrategy, stock_data_cache: Dict, date: str,
        portfolio: Dict, initial_capital: float, commission_rate: float, slippage: float
    ) -> None:
        """调仓策略"""
        # 计算选股分数
        stock_scores = []
        for code in stock_data_cache:
            try:
                data = stock_data_cache[code]
                data = data[data['trade_date'] <= date]
                if len(data) < 20:  # 需要至少20天数据
                    continue
                score = strategy.score(data)
                stock_scores.append({
                    'code': code,
                    'score': score
                })
            except Exception as e:
                self.logger.warning(f"计算股票 {code} 分数失败: {e}")
                continue

        # 按分数降序排序,选择前 N 只
        stock_scores.sort(key=lambda x: x['score'], reverse=True)
        selected_codes = [item['code'] for item in stock_scores[:10]]

        # 计算每只股票的持仓比例
        position_size = 1.0 / len(selected_codes) if selected_codes else 0

        # 记录当前持仓
        portfolio['positions'].append(selected_codes)

        # 计算需要购买的股票
        current_positions = portfolio['positions'][-1] if portfolio['positions'] else []
        cash = portfolio['cash'][-1]

        # 简单模拟交易逻辑(不考虑持仓比例)
        for code in selected_codes:
            if code not in current_positions:
                # 模拟买入
                if code in stock_data_cache:
                    data = stock_data_cache[code]
                    data = data[data['trade_date'] <= date]
                    if not data.empty:
                        price = data['close'].iloc[-1]
                        shares = (cash * position_size) / (price * (1 + commission_rate + slippage))
                        # 简化处理,直接减去成本
                        cost = shares * price * (1 + commission_rate + slippage)
                        cash -= cost

        # 记录现金
        portfolio['cash'].append(cash)

    def _update_portfolio_value(
        self, stock_data_cache: Dict, date: str, portfolio: Dict
    ) -> None:
        """更新持仓价值"""
        current_positions = portfolio['positions'][-1] if portfolio['positions'] else []
        cash = portfolio['cash'][-1]
        position_value = 0.0

        for code in current_positions:
            if code in stock_data_cache:
                data = stock_data_cache[code]
                data = data[data['trade_date'] <= date]
                if not data.empty:
                    # 简化处理,直接取最新价格计算市值
                    price = data['close'].iloc[-1]
                    # 简单模拟持仓成本和数量(实际需要详细记录)
                    shares = 1000  # 简化假设
                    position_value += shares * price

        total_value = cash + position_value
        portfolio['total_value'].append(total_value)

        # 计算净值(NAV)
        nav = total_value / portfolio['total_value'][0]
        portfolio['nav'].append(nav)

    def _calculate_backtest_metrics(self, portfolio: Dict) -> Dict:
        """计算回测指标"""
        if not portfolio['total_value'] or len(portfolio['total_value']) < 2:
            return {}

        # 计算收益相关指标
        nav_series = pd.Series(portfolio['nav'])
        returns = nav_series.pct_change().dropna()

        total_return = (nav_series.iloc[-1] - 1) * 100

        # 年化收益率
        num_years = len(portfolio['dates']) / 252
        annual_return = (nav_series.iloc[-1] ** (1 / num_years) - 1) * 100

        # 最大回撤
        max_drawdown = calculate_drawdown(nav_series.values)

        # 波动率
        annual_volatility = returns.std() * np.sqrt(252) * 100

        # 夏普比率(假设无风险利率为3%)
        risk_free_rate = 0.03
        excess_returns = returns - (risk_free_rate / 252)
        sharpe_ratio = (excess_returns.mean() / excess_returns.std()) * np.sqrt(252)

        # 胜率
        positive_returns = returns[returns > 0]
        win_rate = len(positive_returns) / len(returns) * 100

        # 盈亏比
        average_win = positive_returns.mean()
        negative_returns = returns[returns < 0]
        average_loss = abs(negative_returns.mean())
        profit_loss_ratio = average_win / average_loss if average_loss != 0 else 0

        return {
            'total_return': round(total_return, 2),
            'annual_return': round(annual_return, 2),
            'max_drawdown': round(max_drawdown, 2),
            'annual_volatility': round(annual_volatility, 2),
            'sharpe_ratio': round(sharpe_ratio, 2),
            'win_rate': round(win_rate, 2),
            'profit_loss_ratio': round(profit_loss_ratio, 2),
            'total_trading_days': len(portfolio['dates']),
            'nav_series': nav_series.tolist(),
            'date_series': portfolio['dates']
        }

    def compare_strategies(
        self, strategy_names: List[str], adjust_method: str = "qfq",
        **kwargs
    ) -> List[Dict]:
        """比较多个策略的回测结果"""
        comparison_results = []
        for strategy_name in strategy_names:
            result = self.backtest_strategy(strategy_name, adjust_method, **kwargs)
            if result:
                comparison_results.append({
                    'strategy_name': strategy_name,
                    'description': self.strategy_manager.get_strategy_description(strategy_name),
                    'backtest_result': result
                })

        return comparison_results

    def run_parameter_sweep(
        self, strategy_name: str, parameter_name: str, parameter_values: List,
        adjust_method: str = "qfq", **kwargs
    ) -> Dict:
        """参数扫描回测"""
        sweep_results = []
        for value in parameter_values:
            self.logger.info(f"正在测试参数 {parameter_name} = {value}")
            result = self.backtest_strategy(strategy_name, adjust_method, **kwargs)
            if result:
                sweep_results.append({
                    parameter_name: value,
                    'backtest_result': result
                })

        # 找到最优结果
        if sweep_results:
            best_result = max(
                sweep_results,
                key=lambda x: x['backtest_result']['sharpe_ratio']
            )
            return {
                'parameter_name': parameter_name,
                'parameter_values': parameter_values,
                'sweep_results': sweep_results,
                'best_result': best_result
            }

        return {}