base_strategy.py

"""
QTrading 策略基类模块
"""

import pandas as pd
import numpy as np
from typing import Any, Dict, Optional, List
from abc import ABC, abstractmethod
from utils.logger import get_logger
from utils.common import calculate_change_percent, calculate_volatility, calculate_drawdown

logger = get_logger(__name__)

class BaseStrategy(ABC):
    """策略基类"""

    name: str = "BaseStrategy"
    description: str = "策略基类,提供通用指标计算方法"

    def __init__(self, adjust_method: str = "qfq"):
        """
        初始化策略基类

        Args:
            adjust_method: 复权方式(qfq: 前复权, hfq: 后复权, none: 不复权)
        """
        self.adjust_method = adjust_method
        self.logger = logger
        self.params: Dict[str, Any] = {}

    @abstractmethod
    def score(self, data: pd.DataFrame) -> float:
        """
        计算股票得分(0~100)

        Args:
            data: 股票 K 线数据

        Returns:
            float: 得分(0~100)
        """
        pass

    # 通用指标计算方法

    def calculate_ma(self, data: pd.DataFrame, period: int = 20) -> pd.Series:
        """计算移动平均线(MA)"""
        if len(data) < period:
            return pd.Series([np.nan] * len(data))
        return data['close'].rolling(window=period).mean()

    def calculate_ema(self, data: pd.DataFrame, period: int = 20) -> pd.Series:
        """计算指数移动平均线(EMA)"""
        if len(data) < period:
            return pd.Series([np.nan] * len(data))
        return data['close'].ewm(span=period, adjust=False).mean()

    def calculate_vwap(self, data: pd.DataFrame) -> pd.Series:
        """计算成交量加权平均价格(VWAP)"""
        if 'amount' not in data or 'vol' not in data:
            return pd.Series([np.nan] * len(data))
        return data['amount'] / data['vol']

    def calculate_macd(
        self, data: pd.DataFrame, fast_period: int = 12, slow_period: int = 26, signal_period: int = 9
    ) -> Dict[str, pd.Series]:
        """计算 MACD 指标"""
        ema12 = self.calculate_ema(data, fast_period)
        ema26 = self.calculate_ema(data, slow_period)
        diff = ema12 - ema26
        dea = self.calculate_ema(pd.DataFrame({'close': diff}), signal_period)
        macd = (diff - dea['close']) * 2
        return {
            'diff': diff,
            'dea': dea['close'],
            'macd': macd
        }

    def calculate_rsi(self, data: pd.DataFrame, period: int = 14) -> pd.Series:
        """计算相对强弱指数(RSI)"""
        if len(data) < period + 1:
            return pd.Series([np.nan] * len(data))
        delta = data['close'].diff()
        gain = np.where(delta > 0, delta, 0)
        loss = np.where(delta < 0, -delta, 0)
        avg_gain = pd.Series(gain).rolling(window=period).mean()
        avg_loss = pd.Series(loss).rolling(window=period).mean()
        rs = avg_gain / avg_loss
        rsi = 100 - (100 / (1 + rs))
        return rsi

    def calculate_bollinger_bands(
        self, data: pd.DataFrame, period: int = 20, num_std: float = 2
    ) -> Dict[str, pd.Series]:
        """计算布林带(BOLL)"""
        ma = self.calculate_ma(data, period)
        std = data['close'].rolling(window=period).std()
        upper = ma + num_std * std
        lower = ma - num_std * std
        return {
            'upper': upper,
            'middle': ma,
            'lower': lower
        }

    def calculate_atr(self, data: pd.DataFrame, period: int = 14) -> pd.Series:
        """计算平均真实波幅(ATR)"""
        if len(data) < period:
            return pd.Series([np.nan] * len(data))
        tr1 = data['high'] - data['low']
        tr2 = abs(data['high'] - data['close'].shift(1))
        tr3 = abs(data['low'] - data['close'].shift(1))
        tr = pd.concat([tr1, tr2, tr3], axis=1).max(axis=1)
        return tr.rolling(window=period).mean()

    def calculate_money_flow_index(
        self, data: pd.DataFrame, period: int = 14
    ) -> pd.Series:
        """计算资金流向指标(MFI)"""
        if len(data) < period + 1:
            return pd.Series([np.nan] * len(data))
        typical_price = (data['high'] + data['low'] + data['close']) / 3
        money_flow = typical_price * data['vol']
        positive_flow = np.where(data['close'] > data['close'].shift(1), money_flow, 0)
        negative_flow = np.where(data['close'] < data['close'].shift(1), money_flow, 0)
        positive_mf = pd.Series(positive_flow).rolling(window=period).sum()
        negative_mf = pd.Series(negative_flow).rolling(window=period).sum()
        money_ratio = positive_mf / negative_mf
        mfi = 100 - (100 / (1 + money_ratio))
        return mfi

    # 高级指标计算方法

    def calculate_price_momentum(self, data: pd.DataFrame, period: int = 20) -> float:
        """计算价格动量"""
        if len(data) < period:
            return 0.0
        current_close = data['close'].iloc[-1]
        past_close = data['close'].iloc[-period]
        return calculate_change_percent(current_close, past_close)

    def calculate_price_volatility(self, data: pd.DataFrame, period: int = 20) -> float:
        """计算价格波动率"""
        if len(data) < period:
            return 0.0
        returns = data['close'].pct_change().dropna()
        if len(returns) < period:
            return 0.0
        recent_returns = returns[-period:]
        return np.std(recent_returns) * np.sqrt(252) * 100  # 年化波动率

    def calculate_volume_momentum(self, data: pd.DataFrame, period: int = 20) -> float:
        """计算成交量动量"""
        if len(data) < period:
            return 0.0
        current_vol = data['vol'].iloc[-1]
        avg_vol = data['vol'].rolling(window=period).mean().iloc[-1]
        return calculate_change_percent(current_vol, avg_vol)

    def calculate_price_volume_correlation(
        self, data: pd.DataFrame, period: int = 20
    ) -> float:
        """计算价格与成交量相关性"""
        if len(data) < period + 1:
            return 0.0
        returns = data['close'].pct_change().dropna()
        volume_change = data['vol'].pct_change().dropna()
        if len(returns) < period or len(volume_change) < period:
            return 0.0
        recent_returns = returns[-period:]
        recent_volume = volume_change[-period:]
        return np.corrcoef(recent_returns, recent_volume)[0, 1]

    def calculate_bullish_engulfing_pattern(self, data: pd.DataFrame) -> int:
        """检查看涨吞没形态"""
        if len(data) < 2:
            return 0
        prev = data.iloc[-2]
        curr = data.iloc[-1]
        if (prev['open'] > prev['close'] and curr['open'] < curr['close'] and
                curr['open'] < prev['close'] and curr['close'] > prev['open']):
            return 1
        return 0

    def calculate_bearish_engulfing_pattern(self, data: pd.DataFrame) -> int:
        """检查看跌吞没形态"""
        if len(data) < 2:
            return 0
        prev = data.iloc[-2]
        curr = data.iloc[-1]
        if (prev['open'] < prev['close'] and curr['open'] > curr['close'] and
                curr['open'] > prev['close'] and curr['close'] < prev['open']):
            return 1
        return 0

    def calculate_hammer_pattern(self, data: pd.DataFrame) -> int:
        """检查锤子线形态"""
        if len(data) < 1:
            return 0
        curr = data.iloc[-1]
        body = abs(curr['open'] - curr['close'])
        lower_shadow = min(curr['open'], curr['close']) - curr['low']
        upper_shadow = curr['high'] - max(curr['open'], curr['close'])
        if (body > 0 and lower_shadow > 2 * body and upper_shadow < 0.5 * body):
            return 1
        return 0

    def calculate_inverted_hammer_pattern(self, data: pd.DataFrame) -> int:
        """检查倒锤子线形态"""
        if len(data) < 1:
            return 0
        curr = data.iloc[-1]
        body = abs(curr['open'] - curr['close'])
        upper_shadow = curr['high'] - max(curr['open'], curr['close'])
        lower_shadow = min(curr['open'], curr['close']) - curr['low']
        if (body > 0 and upper_shadow > 2 * body and lower_shadow < 0.5 * body):
            return 1
        return 0

    # 风险指标计算方法

    def calculate_drawdown(self, data: pd.DataFrame) -> float:
        """计算最大回撤"""
        return calculate_drawdown(data['close'].values)

    def calculate_max_drawdown_duration(self, data: pd.DataFrame) -> int:
        """计算最大回撤持续时间(天数)"""
        if len(data) < 2:
            return 0
        prices = data['close'].values
        peak_idx = 0
        max_duration = 0
        current_duration = 0

        for i in range(1, len(prices)):
            if prices[i] > prices[peak_idx]:
                peak_idx = i
                current_duration = 0
            else:
                current_duration += 1
                if current_duration > max_duration:
                    max_duration = current_duration

        return max_duration

    def calculate_sharpe_ratio(self, data: pd.DataFrame, risk_free_rate: float = 0.03) -> float:
        """计算夏普比率"""
        if len(data) < 2:
            return 0.0
        returns = data['close'].pct_change().dropna()
        if len(returns) < 2:
            return 0.0
        excess_returns = returns - (risk_free_rate / 252)
        return np.mean(excess_returns) / np.std(excess_returns) * np.sqrt(252)

    def calculate_sortino_ratio(self, data: pd.DataFrame, risk_free_rate: float = 0.03) -> float:
        """计算索提诺比率"""
        if len(data) < 2:
            return 0.0
        returns = data['close'].pct_change().dropna()
        if len(returns) < 2:
            return 0.0
        excess_returns = returns - (risk_free_rate / 252)
        downside_returns = np.where(excess_returns < 0, excess_returns, 0)
        downside_deviation = np.sqrt(np.mean(downside_returns ** 2))
        return np.mean(excess_returns) / downside_deviation * np.sqrt(252)

    def calculate_calmar_ratio(self, data: pd.DataFrame) -> float:
        """计算卡玛比率"""
        if len(data) < 2:
            return 0.0
        returns = data['close'].pct_change().dropna()
        if len(returns) < 2:
            return 0.0
        annual_return = np.mean(returns) * 252
        max_drawdown = self.calculate_drawdown(data) / 100
        if max_drawdown == 0:
            return 0.0
        return annual_return / max_drawdown