config_manager.py

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
统一配置管理模块
负责程序所有配置的加载、管理和验证
"""

import os
import logging
import json
from pathlib import Path
from dotenv import load_dotenv

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

class ConfigManager:
    """统一配置管理类"""
    
    _instance = None
    
    def __new__(cls, *args, **kwargs):
        """单例模式创建实例"""
        if cls._instance is None:
            cls._instance = super().__new__(cls)
            cls._instance._initialized = False
        return cls._instance
    
    def __init__(self):
        """初始化配置管理器"""
        if self._initialized:
            return
        
        self._initialized = True
        self._config = {}
        self._load_config()
    
    def _load_config(self):
        """加载所有配置"""
        try:
            # 加载环境变量
            self._load_env_variables()
            
            # 加载策略配置
            self._load_strategy_configs()
            
            logger.info("配置加载完成")
            
        except Exception as e:
            logger.error(f"配置加载失败:{e}")
            raise
    
    def _load_env_variables(self):
        """加载环境变量配置"""
        load_dotenv()
        
        # 系统配置
        self._config['system'] = {
            'db_path': os.getenv("DB_PATH", "data/stock_data.db"),
            'log_path': os.getenv("LOG_PATH", "logs/quant_trading.log"),
            'backtest_start_date': os.getenv("BACKTEST_START_DATE", "20260101"),
            'backtest_end_date': os.getenv("BACKTEST_END_DATE", "20260204")
        }
        
        # Tushare API配置
        self._config['tushare'] = {
            'token': os.getenv("TUSHARE_TOKEN"),
            'max_retries': int(os.getenv("TUSHARE_MAX_RETRIES", 3)),
            'retry_delay': int(os.getenv("TUSHARE_RETRY_DELAY", 1))
        }
        
        # 火山引擎API配置
        self._config['volcengine'] = {
            'api_key': os.getenv("VOLCENGINE_API_KEY"),
            'api_base': os.getenv("VOLCENGINE_API_BASE", "https://ark.cn-beijing.volces.com/api/v3/chat/completions"),
            'model': os.getenv("VOLCENGINE_MODEL", "doubao-pro-32k"),
            'temperature': float(os.getenv("TEMPERATURE", 0.1)),
            'top_p': float(os.getenv("TOP_P", 0.9)),
            'max_tokens': int(os.getenv("MAX_TOKENS", 8192))
        }
        
        # 选股参数
        self._config['selection'] = {
            'time': os.getenv("SELECTION_TIME", "14:30"),
            'count': int(os.getenv("STOCK_SELECTION_COUNT", 10)),
            'min_turnover_rate': float(os.getenv("MIN_TURNOVER_RATE", 1.0)),
            'max_stock_count': int(os.getenv("MAX_STOCK_COUNT", 3000))
        }
        
        # 网页服务配置
        self._config['web'] = {
            'host': os.getenv("WEB_HOST", "0.0.0.0"),
            'port': int(os.getenv("WEB_PORT", 5000)),
            'debug': os.getenv("WEB_DEBUG", "false").lower() == "true"
        }
    
    def _load_strategy_configs(self):
        """加载策略配置"""
        # 大模型选股策略配置
        self._config['model_strategy'] = self._load_json_config("config/model_strategy.json", {
            "learning_window": "90日",
            "profit_threshold": "3%",
            "stock_count": "10只",
            "volume_ratio_threshold": "1.5倍"
        })
        
        # VCP极致坍塌模型策略配置
        self._config['vcp_strategy'] = self._load_json_config("config/vcp_strategy.json", {
            "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
        })
        
        # 测试选股策略配置
        self._config['test_strategy'] = self._load_json_config("config/test_strategy.json", {
            "min_price": 10,
            "max_price": 100,
            "min_volume": 1000000,
            "max_volatility": 0.05,
            "min_rsi": 30,
            "max_rsi": 70,
            "min_macd": -0.1,
            "max_macd": 0.1
        })
    
    def _load_json_config(self, file_path, default_config):
        """加载JSON配置文件"""
        try:
            if Path(file_path).exists():
                with open(file_path, 'r', encoding='utf-8') as f:
                    return json.load(f)
            else:
                logger.warning(f"配置文件 {file_path} 不存在,使用默认配置")
                return default_config
        except Exception as e:
            logger.error(f"加载配置文件 {file_path} 失败:{e}")
            return default_config
    
    def get(self, key, default=None):
        """获取配置值"""
        keys = key.split('.')
        value = self._config
        
        for k in keys:
            if k in value:
                value = value[k]
            else:
                return default
        
        return value
    
    def set(self, key, value):
        """设置配置值"""
        keys = key.split('.')
        config = self._config
        
        for k in keys[:-1]:
            if k not in config:
                config[k] = {}
            config = config[k]
        
        config[keys[-1]] = value
    
    def save_strategy_config(self, strategy_name, config):
        """保存策略配置到文件"""
        if strategy_name == 'model':
            file_path = "config/model_strategy.json"
        elif strategy_name == 'vcp':
            file_path = "config/vcp_strategy.json"
        elif strategy_name == 'test':
            file_path = "config/test_strategy.json"
        else:
            logger.error(f"不支持的策略类型:{strategy_name}")
            return False
        
        try:
            with open(file_path, 'w', encoding='utf-8') as f:
                json.dump(config, f, ensure_ascii=False, indent=2)
            
            logger.info(f"策略配置保存成功:{file_path}")
            return True
        except Exception as e:
            logger.error(f"策略配置保存失败:{e}")
            return False
    
    def validate_config(self):
        """验证配置的完整性"""
        required_fields = [
            'tushare.token',
            'volcengine.api_key'
        ]
        
        missing_fields = []
        for field in required_fields:
            if self.get(field) is None:
                missing_fields.append(field)
        
        if missing_fields:
            logger.error(f"配置不完整,缺少以下必填字段:{', '.join(missing_fields)}")
            return False
        
        return True
    
    def print_config(self):
        """打印配置信息"""
        logger.info("当前配置:")
        logger.info(f"数据库路径:{self.get('system.db_path')}")
        logger.info(f"日志路径:{self.get('system.log_path')}")
        logger.info(f"选股时间:{self.get('selection.time')}")
        logger.info(f"选股数量:{self.get('selection.count')}")
        logger.info(f"最小换手率:{self.get('selection.min_turnover_rate')}")
        logger.info(f"最大股票数量:{self.get('selection.max_stock_count')}")
        logger.info(f"网页服务地址:{self.get('web.host')}:{self.get('web.port')}")
        logger.info(f"网页调试模式:{self.get('web.debug')}")

# 创建全局配置实例
config_manager = ConfigManager()