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