stock_selector.py

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
选股器核心调度器 - 通用唯一入口
负责从股票数据库中循环读取每一只股票的数据、加载管理策略、汇总结果
"""

import logging
import pandas as pd
from pathlib import Path
from datetime import datetime, timedelta
from typing import Dict, List, Any, Optional

from src.config_manager import config_manager
from src.utils import ensure_dir_exists, get_trade_date
from src.strategy_manager import strategy_manager
from src.data_manager import DataManager

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

class StockSelector:
    """选股器核心调度器"""
    
    def __init__(self):
        """初始化选股器"""
        self.stock_count = config_manager.get('selection.count')
        self.min_turnover_rate = config_manager.get('selection.min_turnover_rate')
        
        # 结果保存目录
        self.result_dir = "results/selections"
        ensure_dir_exists(self.result_dir)
        
        # 初始化数据管理模块
        self.data_manager = DataManager()
        
        logger.info("选股器核心调度器初始化完成")
    
    def run(self, strategy_names: List[str], trade_date: Optional[str] = None, trade_time: Optional[str] = None) -> pd.DataFrame:
        """
        执行选股流程
        
        参数:
            strategy_names: 要使用的策略名称列表
            trade_date: 选股日期(格式:YYYYMMDD)
            trade_time: 选股时间(格式:HH:MM)
            
        返回:
            pd.DataFrame: 选股结果
        """
        try:
            logger.info("开始选股流程...")
            
            # 确定日期和时间
            if not trade_date:
                trade_date = get_trade_date()
            if not trade_time:
                trade_time = config_manager.get('selection.time')
            
            logger.info(f"选股日期: {trade_date},选股时间: {trade_time}")
            
            # 检查是否为交易日
            if not self._is_trading_day(trade_date):
                logger.warning(f"{trade_date} 不是交易日,选股终止")
                return pd.DataFrame()
            
            # 获取有效股票列表
            df_valid = self.data_manager.get_valid_stocks()
            logger.info(f"有效股票数量: {len(df_valid)}")
            if df_valid.empty:
                logger.warning("无有效股票数据,选股终止")
                return pd.DataFrame()
            
            # 获取所有股票的历史数据(用于选股)
            logger.info("获取股票历史数据...")
            all_stock_data = self.data_manager.get_all_stock_data_for_backtest(
                start_date=(datetime.strptime(trade_date, '%Y%m%d') - timedelta(days=120)).strftime('%Y%m%d'),
                end_date=trade_date
            )
            
            logger.info(f"历史数据获取完成,共 {len(all_stock_data)} 只股票有数据")
            
            # 加载指定策略
            strategies = []
            for strategy_name in strategy_names:
                strategy = strategy_manager.get_strategy(strategy_name)
                if strategy:
                    strategies.append(strategy)
                    logger.info(f"策略加载成功: {strategy.get_name()} - {strategy.get_description()}")
                else:
                    logger.error(f"策略加载失败: {strategy_name}")
            
            if not strategies:
                logger.error("未加载到任何有效策略,选股终止")
                return pd.DataFrame()
            
            # 遍历股票池,使用策略进行选股
            selected_stocks = []
            
            for ts_code, df_stock in all_stock_data.items():
                try:
                    # 检查股票是否符合所有策略条件(逻辑与)
                    all_conditions_met = True
                    for strategy in strategies:
                        if not strategy.check_stock(ts_code, df_stock):
                            all_conditions_met = False
                            break
                    
                    if all_conditions_met:
                        # 获取股票基本信息
                        stock_info = df_valid[df_valid['ts_code'] == ts_code]
                        if not stock_info.empty:
                            selected_stocks.append({
                                "ts_code": ts_code,
                                "name": stock_info.iloc[0]['name'],
                                "industry": stock_info.iloc[0]['industry']
                            })
                    
                except Exception as e:
                    logger.error(f"处理股票 {ts_code} 失败: {e}")
                    continue
            
            # 转换为DataFrame
            df_selected = pd.DataFrame(selected_stocks)
            
            # 筛选股票数量
            if not df_selected.empty:
                if len(df_selected) > self.stock_count:
                    df_selected = df_selected.head(self.stock_count)
                
                logger.info(f"选股完成,共选出 {len(df_selected)} 只股票")
                return df_selected
            else:
                logger.warning("未选出符合条件的股票")
                return pd.DataFrame()
                
        except Exception as e:
            logger.error(f"选股失败: {e}")
            return pd.DataFrame()
    
    def run_strategy_combination(self, strategy_combinations: List[Dict[str, Any]], 
                               trade_date: Optional[str] = None, 
                               trade_time: Optional[str] = None) -> Dict[str, pd.DataFrame]:
        """
        执行策略组合选股
        
        参数:
            strategy_combinations: 策略组合配置
            trade_date: 选股日期
            trade_time: 选股时间
            
        返回:
            Dict[str, pd.DataFrame]: 不同策略组合的选股结果
        """
        results = {}
        
        for combination in strategy_combinations:
            combination_name = combination.get('name', 'default')
            strategy_names = combination.get('strategies', [])
            
            logger.info(f"执行策略组合: {combination_name}")
            result = self.run(strategy_names, trade_date, trade_time)
            results[combination_name] = result
        
        return results
    
    def save_results(self, df_selected: pd.DataFrame, trade_date: Optional[str] = None, 
                   strategy_name: str = "combined") -> bool:
        """
        保存选股结果
        
        参数:
            df_selected: 选股结果
            trade_date: 选股日期
            strategy_name: 策略名称
            
        返回:
            bool: 保存是否成功
        """
        if df_selected.empty:
            logger.warning("无选股结果需要保存")
            return False
        
        try:
            if not trade_date:
                trade_date = get_trade_date()
            
            # 保存CSV文件
            csv_path = Path(self.result_dir) / f"{trade_date}_{strategy_name}_selected_stocks.csv"
            df_selected.to_csv(csv_path, index=False, encoding='utf-8-sig')
            
            # 保存JSON文件
            json_path = Path(self.result_dir) / f"{trade_date}_{strategy_name}_selected_stocks.json"
            df_selected.to_json(json_path, orient='records', force_ascii=False, indent=2)
            
            logger.info(f"选股结果已保存:")
            logger.info(f"CSV文件: {csv_path}")
            logger.info(f"JSON文件: {json_path}")
            
            return True
            
        except Exception as e:
            logger.error(f"选股结果保存失败: {e}")
            return False
    
    def get_results(self, trade_date: Optional[str] = None, 
                  strategy_name: str = "combined") -> pd.DataFrame:
        """
        获取选股结果
        
        参数:
            trade_date: 选股日期
            strategy_name: 策略名称
            
        返回:
            pd.DataFrame: 选股结果
        """
        try:
            if not trade_date:
                trade_date = get_trade_date()
            
            # 检查CSV文件是否存在
            csv_path = Path(self.result_dir) / f"{trade_date}_{strategy_name}_selected_stocks.csv"
            if not csv_path.exists():
                logger.warning(f"{trade_date} {strategy_name}选股结果文件不存在")
                return pd.DataFrame()
            
            # 读取CSV文件
            df_selected = pd.read_csv(csv_path)
            
            logger.info(f"成功获取 {trade_date} {strategy_name}选股结果")
            return df_selected
            
        except Exception as e:
            logger.error(f"获取选股结果失败: {e}")
            return pd.DataFrame()
    
    def print_results(self, df_selected: pd.DataFrame, strategy_name: str = "combined"):
        """
        打印选股结果
        
        参数:
            df_selected: 选股结果
            strategy_name: 策略名称
        """
        if df_selected.empty:
            logger.warning("无选股结果需要打印")
            return
        
        logger.info("=" * 80)
        logger.info(f"{strategy_name}选股策略选股结果")
        logger.info("=" * 80)
        
        print(f"{'代码':<10} {'名称':<10} {'行业':<10}")
        print("-" * 80)
        
        for _, row in df_selected.iterrows():
            print(f"{row['ts_code']:<10} {row['name']:<10} {row['industry']:<10}")
        
        logger.info("=" * 80)
    
    def _is_trading_day(self, date_str: str) -> bool:
        """
        检查是否为交易日
        
        参数:
            date_str: 日期字符串(格式:YYYYMMDD)
            
        返回:
            bool: 是否为交易日
        """
        try:
            from chinese_calendar import is_workday
            date = datetime.strptime(date_str, '%Y%m%d').date()
            return is_workday(date)
        except Exception as e:
            logger.error(f"交易日检查失败: {e}")
            return False
    
    def test_selector(self, strategy_names: List[str], test_date: str = '20241231') -> None:
        """
        测试选股器功能
        
        参数:
            strategy_names: 要测试的策略名称列表
            test_date: 测试日期
        """
        logger.info("开始测试选股器功能...")
        
        try:
            # 测试获取有效股票列表
            df_valid = self.data_manager.get_valid_stocks()
            logger.info(f"有效股票数量: {len(df_valid)}")
            
            # 测试获取历史数据
            all_stock_data = self.data_manager.get_all_stock_data_for_backtest(
                start_date=(datetime.strptime(test_date, '%Y%m%d') - timedelta(days=120)).strftime('%Y%m%d'),
                end_date=test_date
            )
            logger.info(f"历史数据获取完成,共 {len(all_stock_data)} 只股票有数据")
            
            # 测试策略加载
            strategies = []
            for strategy_name in strategy_names:
                strategy = strategy_manager.get_strategy(strategy_name)
                if strategy:
                    strategies.append(strategy)
            
            logger.info(f"策略加载完成,共 {len(strategies)} 个有效策略")
            
            logger.info("选股器功能测试完成")
            
        except Exception as e:
            logger.error(f"选股器测试失败: {e}")

if __name__ == "__main__":
    # 配置日志
    logging.basicConfig(
        level=logging.INFO,
        format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
    )
    
    # 测试选股器
    selector = StockSelector()
    selector.test_selector(['vcp', 'model'])