vcp_selector.py

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
VCP选股器类
负责VCP极致坍塌模型选股策略的选股操作
与股票遍历逻辑解耦,使用策略分析器分析单只股票
"""

import os
import logging
import pandas as pd
import numpy as np
import json
import time
from pathlib import Path
from datetime import datetime, timedelta
from dotenv import load_dotenv

from src.config_manager import config_manager
from src.utils import ensure_dir_exists, load_json_file, save_json_file, get_trade_date

# 加载环境变量
load_dotenv()

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

class VCPSelector:
    """VCP选股器类"""
    
    def __init__(self):
        """初始化VCP选股器"""
        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)
        
        # 初始化数据管理模块
        from src.data_manager import DataManager
        self.data_manager = DataManager()
        
        # 加载策略参数
        self.params = self._load_strategy_parameters()
    
    def _load_strategy_parameters(self):
        """加载策略参数"""
        params_file = "config/vcp_strategy.json"
        
        if Path(params_file).exists():
            try:
                with open(params_file, 'r', encoding='utf-8') as f:
                    return json.load(f)
            except Exception as e:
                logger.error(f"加载策略参数失败:{e}")
        
        # 默认参数
        return {
            "min_contraction_period": 20,
            "max_contraction_period": 60,
            "volatility_threshold": 0.15,
            "breakout_threshold": 0.05,
            "min_price": 5,
            "max_price": 200,
            "min_volume_ratio": 1.5,
            "min_macd_signal": 0.01,
            "rsi_upper_bound": 70,
            "rsi_lower_bound": 30
        }
    
    def run(self, trade_date=None, trade_time=None):
        """执行选股流程"""
        try:
            logger.info("开始VCP选股流程...")
            
            # 确定日期和时间
            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)} 只股票有数据")
            
            # 外层主程序遍历股票池(策略逻辑与遍历解耦)
            selected_stocks = []
            
            # 动态导入策略管理器(避免循环导入)
            from src.strategy_manager import strategy_manager
            
            for ts_code, df_stock in all_stock_data.items():
                try:
                    # 使用策略分析器分析单只股票(策略逻辑与主程序完全解耦)
                    analysis_result = strategy_manager.analyze_single_stock('vcp', ts_code, df_stock)
                    
                    if analysis_result["score"] > 0:
                        # 获取股票基本信息
                        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'],
                                "score": analysis_result["score"],
                                "level": analysis_result["level"],
                                "reason": analysis_result["reason"],
                                "signals": analysis_result["signals"],
                                "factors": self._get_selection_factors(df_stock)
                            })
                    
                except Exception as e:
                    logger.error(f"处理股票 {ts_code} 失败:{e}")
                    continue
            
            # 转换为DataFrame
            df_selected = pd.DataFrame(selected_stocks)
            
            # 筛选评分最高的股票
            if not df_selected.empty:
                df_selected = df_selected.sort_values(by='score', ascending=False).head(self.stock_count)
                logger.info(f"VCP选股完成,共选出 {len(df_selected)} 只股票")
                return df_selected
            else:
                logger.warning("未选出符合条件的股票")
                return pd.DataFrame()
                
        except Exception as e:
            logger.error(f"VCP选股失败:{e}")
            return pd.DataFrame()
    
    def _is_trading_day(self, date_str):
        """检查是否为交易日"""
        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 _get_selection_factors(self, df_stock):
        """获取选股因素(可自定义)"""
        try:
            factors = []
            
            # 计算价格波动率
            df_stock['volatility'] = df_stock['close'].pct_change().rolling(20).std()
            
            # 寻找波动率收缩的时期
            contraction_periods = self._find_volatility_contraction(df_stock['volatility'])
            
            if contraction_periods:
                latest_contraction = contraction_periods[-1]
                factors.append({
                    "name": "收缩期长度",
                    "value": f"{latest_contraction['length']}个交易日"
                })
                factors.append({
                    "name": "波动率收缩幅度",
                    "value": f"{latest_contraction['contraction_ratio']:.1%}"
                })
                
                # 计算突破幅度
                breakout_range = self._calculate_breakout_range(df_stock, latest_contraction['end_index'])
                factors.append({
                    "name": "突破幅度",
                    "value": f"{breakout_range:.1%}"
                })
            
            # 计算RSI指标
            rsi = self._calculate_rsi(df_stock['close'], 14).iloc[-1]
            factors.append({
                "name": "RSI",
                "value": f"{rsi:.1f}"
            })
            
            # 计算MACD指标
            macd, _, _ = self._calculate_macd(df_stock['close'])
            factors.append({
                "name": "MACD",
                "value": f"{macd[-1]:.3f}"
            })
            
            return factors
            
        except Exception as e:
            logger.error(f"选股因素获取失败:{e}")
            return []
    
    def _find_volatility_contraction(self, volatility_series):
        """寻找波动率收缩的时期"""
        contraction_periods = []
        start_idx = None
        
        for i in range(1, len(volatility_series)):
            if start_idx is None and volatility_series[i] < volatility_series[i-1]:
                start_idx = i-1
            
            if start_idx is not None and volatility_series[i] > volatility_series[i-1]:
                contraction_length = i - start_idx
                contraction_ratio = (volatility_series[start_idx] - volatility_series[i-1]) / volatility_series[start_idx]
                
                contraction_periods.append({
                    "start_index": start_idx,
                    "end_index": i-1,
                    "length": contraction_length,
                    "contraction_ratio": contraction_ratio,
                    "start_vol": volatility_series[start_idx],
                    "end_vol": volatility_series[i-1]
                })
                
                start_idx = None
        
        return contraction_periods
    
    def _calculate_breakout_range(self, df_stock, end_index):
        """计算突破幅度"""
        try:
            breakout_data = df_stock.iloc[end_index:end_index+5]
            max_price = breakout_data['high'].max()
            min_price = breakout_data['low'].min()
            breakout_range = (max_price - min_price) / min_price
            
            return breakout_range
        except Exception as e:
            logger.error(f"计算突破幅度失败:{e}")
            return 0
    
    def _calculate_rsi(self, prices, period=14):
        """计算RSI指标"""
        delta = prices.diff()
        gain = (delta.where(delta > 0, 0)).rolling(window=period).mean()
        loss = (-delta.where(delta < 0, 0)).rolling(window=period).mean()
        
        rs = gain / loss
        rsi = 100 - (100 / (1 + rs))
        
        return rsi
    
    def _calculate_macd(self, prices, fast=12, slow=26, signal_period=9):
        """计算MACD指标"""
        ema_fast = prices.ewm(span=fast, adjust=False).mean()
        ema_slow = prices.ewm(span=slow, adjust=False).mean()
        
        macd = ema_fast - ema_slow
        signal = macd.ewm(span=signal_period, adjust=False).mean()
        hist = macd - signal
        
        return macd.values, signal.values, hist.values
    
    def save_results(self, df_selected, trade_date=None):
        """保存选股结果"""
        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}_vcp_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}_vcp_selected_stocks.json"
            df_selected.to_json(json_path, orient='records', force_ascii=False, indent=2)
            
            logger.info(f"VCP选股结果已保存:")
            logger.info(f"CSV文件:{csv_path}")
            logger.info(f"JSON文件:{json_path}")
            
            return True
            
        except Exception as e:
            logger.error(f"VCP选股结果保存失败:{e}")
            return False
    
    def get_results(self, trade_date=None):
        """获取选股结果"""
        try:
            if not trade_date:
                trade_date = get_trade_date()
            
            # 检查CSV文件是否存在
            csv_path = Path(self.result_dir) / f"{trade_date}_vcp_selected_stocks.csv"
            if not csv_path.exists():
                logger.warning(f"{trade_date} VCP选股结果文件不存在")
                return {}
            
            # 读取CSV文件
            df_selected = pd.read_csv(csv_path)
            
            # 转换为字典
            results = {}
            for index, row in df_selected.iterrows():
                results[row['ts_code']] = row['name']
            
            logger.info(f"成功获取 {trade_date} VCP选股结果")
            return results
            
        except Exception as e:
            logger.error(f"获取VCP选股结果失败:{e}")
            return {}
    
    def print_results(self, df_selected):
        """打印选股结果"""
        if df_selected.empty:
            logger.warning("无选股结果需要打印")
            return
        
        logger.info("=" * 80)
        logger.info("VCP选股策略选股结果")
        logger.info("=" * 80)
        
        print(f"{'代码':<10} {'名称':<10} {'行业':<10} {'评分':<8} {'等级':<4} {'选股理由'}")
        print("-" * 80)
        
        for _, row in df_selected.iterrows():
            print(f"{row['ts_code']:<10} {row['name']:<10} {row['industry']:<10} {row['score']:<8.1f} {row['level']:<4} {row['reason']}")
        
        logger.info("=" * 80)
    
    def test_selection(self):
        """测试选股功能"""
        logger.info("开始测试VCP选股功能...")
        
        # 测试获取有效股票列表
        df_valid = self.data_manager.get_valid_stocks()
        logger.info(f"有效股票数量:{len(df_valid)}")
        
        # 测试获取历史数据
        test_date = '20241231'
        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)} 只股票有数据")
        
        logger.info("VCP选股功能测试完成")

if __name__ == "__main__":
    # 配置日志
    logging.basicConfig(
        level=logging.INFO,
        format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
    )
    
    # 测试选股模块
    selector = VCPSelector()
    selector.test_selection()
    # result = selector.run('20240102', '14:30:00')
    # selector.save_results(result, '20240102')
    # selector.print_results(result)