selector_engine.py

"""
QTrading 选股引擎模块
"""

import pandas as pd
from typing import List, Dict, Optional, Union
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, today_str

logger = get_logger(__name__)

class SelectorEngine:
    """选股引擎"""

    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 select_stocks(
        self,
        strategy_name: str,
        adjust_method: str = "qfq",
        stock_list: List[str] = None,
        start_date: str = None,
        end_date: str = None
    ) -> List[Dict]:
        """
        选股主方法

        Args:
            strategy_name: 策略名称
            adjust_method: 复权方式(qfq: 前复权, hfq: 后复权, none: 不复权)
            stock_list: 股票池(可选,默认从配置文件读取)
            start_date: 开始日期(YYYYMMDD,可选)
            end_date: 结束日期(YYYYMMDD,可选)

        Returns:
            List[Dict]: 选股结果(按分数降序排列)
        """
        if stock_list is None:
            stock_list = config.stock_pool.stock_list

        if not stock_list:
            self.logger.warning("股票池为空,无法选股")
            return []

        end_date = end_date or today_str()

        # 创建策略实例
        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},复权方式: {adjust_method}")
        self.logger.info(f"股票池大小: {len(stock_list)} 只")
        self.logger.info(f"数据日期范围: {start_date} - {end_date}")

        results = []
        success_count = 0
        fail_count = 0

        for i, code in enumerate(stock_list, 1):
            self.logger.debug(f"正在处理 {i}/{len(stock_list)}: {code}")
            try:
                if not validate_code(code):
                    self.logger.warning(f"股票代码格式无效: {code}")
                    continue

                # 获取复权后的数据
                data = self.data_manager.calculate_adjusted_data(
                    code, adjust_method, start_date, end_date
                )
                if data is None or data.empty:
                    self.logger.warning(f"未找到股票 {code} 的数据")
                    continue

                # 计算分数
                score = strategy.score(data)

                # 获取基本信息(最新价格、成交量等)
                latest_data = data.iloc[-1]
                result = {
                    'code': code,
                    'ts_code': latest_data['ts_code'],
                    'trade_date': latest_data['trade_date'],
                    'close': latest_data['close'],
                    'vol': latest_data['vol'],
                    'score': score
                }

                results.append(result)
                success_count += 1

            except Exception as e:
                self.logger.error(f"处理股票 {code} 失败: {e}")
                fail_count += 1

        # 按分数降序排序
        results.sort(key=lambda x: x['score'], reverse=True)

        self.logger.info(f"选股完成: 成功 {success_count} 只,失败 {fail_count} 只")
        self.logger.info(f"选股结果: {len(results)} 只股票(分数 >= 0)")

        return results

    def select_top_n(
        self,
        strategy_name: str,
        n: int = 10,
        adjust_method: str = "qfq",
        stock_list: List[str] = None,
        start_date: str = None,
        end_date: str = None
    ) -> List[Dict]:
        """
        选择前 N 只得分最高的股票

        Args:
            strategy_name: 策略名称
            n: 要选择的股票数量
            adjust_method: 复权方式(qfq: 前复权, hfq: 后复权, none: 不复权)
            stock_list: 股票池(可选,默认从配置文件读取)
            start_date: 开始日期(YYYYMMDD,可选)
            end_date: 结束日期(YYYYMMDD,可选)

        Returns:
            List[Dict]: 前 N 只股票的选股结果(按分数降序排列)
        """
        results = self.select_stocks(
            strategy_name, adjust_method, stock_list, start_date, end_date
        )
        return results[:n]

    def select_by_score_threshold(
        self,
        strategy_name: str,
        threshold: float = 60,
        adjust_method: str = "qfq",
        stock_list: List[str] = None,
        start_date: str = None,
        end_date: str = None
    ) -> List[Dict]:
        """
        选择得分高于阈值的股票

        Args:
            strategy_name: 策略名称
            threshold: 得分阈值
            adjust_method: 复权方式(qfq: 前复权, hfq: 后复权, none: 不复权)
            stock_list: 股票池(可选,默认从配置文件读取)
            start_date: 开始日期(YYYYMMDD,可选)
            end_date: 结束日期(YYYYMMDD,可选)

        Returns:
            List[Dict]: 得分高于阈值的股票(按分数降序排列)
        """
        results = self.select_stocks(
            strategy_name, adjust_method, stock_list, start_date, end_date
        )
        return [result for result in results if result['score'] >= threshold]

    def get_strategy_performance(
        self,
        strategy_name: str,
        adjust_method: str = "qfq",
        stock_list: List[str] = None,
        start_date: str = None,
        end_date: str = None
    ) -> Dict:
        """
        获取策略在股票池上的整体表现

        Args:
            strategy_name: 策略名称
            adjust_method: 复权方式(qfq: 前复权, hfq: 后复权, none: 不复权)
            stock_list: 股票池(可选,默认从配置文件读取)
            start_date: 开始日期(YYYYMMDD,可选)
            end_date: 结束日期(YYYYMMDD,可选)

        Returns:
            Dict: 策略整体表现统计
        """
        results = self.select_stocks(
            strategy_name, adjust_method, stock_list, start_date, end_date
        )

        if not results:
            return {
                'total_stocks': 0,
                'avg_score': 0.0,
                'max_score': 0.0,
                'min_score': 0.0,
                'score_distribution': {},
                'top_10_avg_score': 0.0
            }

        scores = [result['score'] for result in results]
        avg_score = sum(scores) / len(scores)
        max_score = max(scores)
        min_score = min(scores)

        # 分数分布
        score_distribution = {
            '0-20': 0,
            '21-40': 0,
            '41-60': 0,
            '61-80': 0,
            '81-100': 0
        }
        for score in scores:
            if score <= 20:
                score_distribution['0-20'] += 1
            elif score <= 40:
                score_distribution['21-40'] += 1
            elif score <= 60:
                score_distribution['41-60'] += 1
            elif score <= 80:
                score_distribution['61-80'] += 1
            else:
                score_distribution['81-100'] += 1

        # 前10只平均得分
        top_10_avg = sum(result['score'] for result in results[:10]) / min(10, len(results))

        return {
            'total_stocks': len(results),
            'avg_score': round(avg_score, 2),
            'max_score': round(max_score, 2),
            'min_score': round(min_score, 2),
            'score_distribution': score_distribution,
            'top_10_avg_score': round(top_10_avg, 2)
        }

    def compare_strategies(
        self,
        strategy_names: List[str],
        adjust_method: str = "qfq",
        stock_list: List[str] = None,
        start_date: str = None,
        end_date: str = None
    ) -> List[Dict]:
        """
        比较多个策略的选股结果

        Args:
            strategy_names: 策略名称列表
            adjust_method: 复权方式(qfq: 前复权, hfq: 后复权, none: 不复权)
            stock_list: 股票池(可选,默认从配置文件读取)
            start_date: 开始日期(YYYYMMDD,可选)
            end_date: 结束日期(YYYYMMDD,可选)

        Returns:
            List[Dict]: 策略比较结果
        """
        comparison_results = []
        for strategy_name in strategy_names:
            performance = self.get_strategy_performance(
                strategy_name, adjust_method, stock_list, start_date, end_date
            )
            comparison_results.append({
                'strategy_name': strategy_name,
                'description': self.strategy_manager.get_strategy_description(strategy_name),
                'performance': performance
            })

        return comparison_results

    def get_data_coverage(self, stock_list: List[str] = None, end_date: str = None) -> Dict:
        """
        获取数据覆盖情况

        Args:
            stock_list: 股票池(可选,默认从配置文件读取)
            end_date: 结束日期(YYYYMMDD,可选)

        Returns:
            Dict: 数据覆盖情况
        """
        if stock_list is None:
            stock_list = config.stock_pool.stock_list

        end_date = end_date or today_str()

        coverage = {
            'total_stocks': len(stock_list),
            'has_data_count': 0,
            'no_data_count': 0,
            'no_data_codes': []
        }

        for code in stock_list:
            try:
                data = self.data_manager.get_kline_data(code, end_date=end_date)
                if data is not None and not data.empty:
                    coverage['has_data_count'] += 1
                else:
                    coverage['no_data_count'] += 1
                    coverage['no_data_codes'].append(code)
            except Exception as e:
                self.logger.error(f"检查股票 {code} 数据覆盖失败: {e}")
                coverage['no_data_count'] += 1
                coverage['no_data_codes'].append(code)

        coverage['coverage_rate'] = round(
            coverage['has_data_count'] / coverage['total_stocks'] * 100, 2
        )

        return coverage