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()