base_strategy.py

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
策略基类 - 定义策略接口
所有选股策略必须继承此基类,实现统一接口
"""

import logging
from abc import ABC, abstractmethod
from typing import Dict, Any, Optional

# 配置日志
logger = logging.getLogger(__name__)

class BaseStrategy(ABC):
    """选股策略基类,所有策略必须继承此类"""
    
    def __init__(self, params: Dict[str, Any] = None):
        """
        初始化策略
        
        参数:
            params: 策略参数配置
        """
        self.name = self.__class__.__name__
        self.description = self._get_description()
        self.params = params or self._get_default_params()
        logger.info(f"策略初始化: {self.name} - {self.description}")
    
    @abstractmethod
    def _get_description(self) -> str:
        """获取策略描述(子类必须实现)"""
        pass
    
    @abstractmethod
    def _get_default_params(self) -> Dict[str, Any]:
        """获取策略默认参数(子类必须实现)"""
        pass
    
    @abstractmethod
    def check_stock(self, ts_code: str, df_stock: Any) -> bool:
        """
        检查单只股票是否符合策略条件(核心接口)
        
        参数:
            ts_code: 股票代码
            df_stock: 股票历史数据(包含OHLCV等)
            
        返回:
            bool: 是否符合条件
        """
        pass
    
    def set_params(self, params: Dict[str, Any]) -> None:
        """
        设置策略参数
        
        参数:
            params: 策略参数配置
        """
        if params:
            self.params.update(params)
            logger.info(f"策略参数已更新: {self.name} - {params}")
    
    def get_params(self) -> Dict[str, Any]:
        """获取策略参数"""
        return self.params.copy()
    
    def get_name(self) -> str:
        """获取策略名称"""
        return self.name
    
    def get_description(self) -> str:
        """获取策略描述"""
        return self.description