strategy_manager.py

"""
QTrading 策略管理模块
"""

from typing import Type, Dict, List, Optional
from strategy.base_strategy import BaseStrategy
from strategy.simple_strategies import (
    MA20Strategy,
    RSIStrategy,
    MACDStrategy,
    BollingerBandStrategy,
    MomentumStrategy
)
from config.config import config
from utils.logger import get_logger

logger = get_logger(__name__)

class StrategyManager:
    """策略管理器"""

    def __init__(self):
        """初始化策略管理器"""
        self.strategies: Dict[str, Type[BaseStrategy]] = {}
        self._load_strategies()

    def _load_strategies(self) -> None:
        """加载所有策略类"""
        # 内置策略列表
        builtin_strategies = [
            MA20Strategy,
            RSIStrategy,
            MACDStrategy,
            BollingerBandStrategy,
            MomentumStrategy
        ]

        for strategy_class in builtin_strategies:
            self._register_strategy(strategy_class)

        logger.info(f"成功加载 {len(self.strategies)} 个策略")
        config.strategy.available_strategies = list(self.strategies.keys())

    def _register_strategy(self, strategy_class: Type[BaseStrategy]) -> None:
        """
        注册策略类

        Args:
            strategy_class: 策略类(继承自 BaseStrategy)
        """
        strategy_name = strategy_class.name
        if strategy_name in self.strategies:
            logger.warning(f"策略 {strategy_name} 已存在,跳过注册")
            return
        self.strategies[strategy_name] = strategy_class
        logger.debug(f"策略 {strategy_name} 注册成功")

    def get_strategy_class(self, strategy_name: str) -> Optional[Type[BaseStrategy]]:
        """
        获取策略类

        Args:
            strategy_name: 策略名称

        Returns:
            Type[BaseStrategy]: 策略类,找不到返回 None
        """
        return self.strategies.get(strategy_name)

    def create_strategy_instance(
        self, strategy_name: str, adjust_method: str = "qfq"
    ) -> Optional[BaseStrategy]:
        """
        创建策略实例

        Args:
            strategy_name: 策略名称
            adjust_method: 复权方式(qfq: 前复权, hfq: 后复权, none: 不复权)

        Returns:
            BaseStrategy: 策略实例,找不到返回 None
        """
        strategy_class = self.get_strategy_class(strategy_name)
        if strategy_class:
            try:
                return strategy_class(adjust_method)
            except Exception as e:
                logger.error(f"创建策略 {strategy_name} 实例失败: {e}")
                return None
        else:
            logger.warning(f"策略 {strategy_name} 未找到")
            return None

    def list_strategies(self) -> List[Dict[str, str]]:
        """
        获取策略列表(包含名称和描述)

        Returns:
            List[Dict[str, str]]: 策略列表
        """
        strategy_list = []
        for name, strategy_class in self.strategies.items():
            strategy_list.append({
                "name": name,
                "description": strategy_class.description
            })
        return strategy_list

    def get_strategy_count(self) -> int:
        """
        获取策略数量

        Returns:
            int: 策略数量
        """
        return len(self.strategies)

    def get_default_strategy(self, adjust_method: str = "qfq") -> Optional[BaseStrategy]:
        """
        获取默认策略实例

        Args:
            adjust_method: 复权方式

        Returns:
            BaseStrategy: 默认策略实例
        """
        return self.create_strategy_instance(config.strategy.default_strategy, adjust_method)

    def get_strategy_description(self, strategy_name: str) -> Optional[str]:
        """
        获取策略描述

        Args:
            strategy_name: 策略名称

        Returns:
            str: 策略描述,找不到返回 None
        """
        strategy_class = self.get_strategy_class(strategy_name)
        if strategy_class:
            return strategy_class.description
        else:
            logger.warning(f"策略 {strategy_name} 未找到")
            return None

    def load_custom_strategies(self, module_path: str = None) -> None:
        """
        加载自定义策略(从指定模块路径)

        Args:
            module_path: 模块路径(可选)
        """
        if module_path is None:
            logger.info("未指定自定义策略模块路径")
            return

        try:
            import importlib.util
            spec = importlib.util.spec_from_file_location("custom_strategies", module_path)
            module = importlib.util.module_from_spec(spec)
            spec.loader.exec_module(module)

            # 查找继承自 BaseStrategy 的类
            for attr in dir(module):
                obj = getattr(module, attr)
                if isinstance(obj, type) and issubclass(obj, BaseStrategy) and obj != BaseStrategy:
                    self._register_strategy(obj)

            logger.info(f"成功加载自定义策略模块: {module_path}")
        except Exception as e:
            logger.error(f"加载自定义策略模块失败: {e}")