#!/usr/bin/env python # -*- coding: utf-8 -*- """ 配置加载器 负责加载、合并和验证配置,优先从环境变量加载。 """ import json import os from typing import Any, Dict from config.schema import ( AppConfig, ERPConfig, DatabaseConfig, PathConfig, ExtractionConfig, ValidationConfig, DatabaseType, SQLServerConfig, MySQLConfig, ) from config.defaults import DEFAULT_APP_CONFIG, DEFAULT_SETTINGS_DICT from config.env_loader import get_env, get_env_bool, get_env_int class ConfigLoader: """配置加载器""" @staticmethod def load( config_file: str = "config/user_settings.json", use_env: bool = True ) -> AppConfig: """ 加载配置 优先级: 1. 环境变量(如果 use_env=True) 2. JSON 配置文件(如果存在) 3. 默认配置 Args: config_file: 配置文件路径 use_env: 是否使用环境变量,默认为 True Returns: 应用配置对象 """ # 优先从环境变量加载 if use_env: return AppConfig.from_env() # 如果不使用环境变量,则从 JSON 文件加载(向后兼容) if os.path.exists(config_file): try: with open(config_file, "r", encoding="utf-8") as f: loaded_settings = json.load(f) # 合并默认配置和加载的配置 merged_settings = ConfigLoader._merge_settings( DEFAULT_SETTINGS_DICT, loaded_settings ) return ConfigLoader._dict_to_config(merged_settings) except (json.JSONDecodeError, IOError) as e: print(f"加载配置文件失败: {e},使用默认配置") return DEFAULT_APP_CONFIG else: # 首次运行,创建默认配置文件 ConfigLoader.save(DEFAULT_APP_CONFIG, config_file) return DEFAULT_APP_CONFIG @staticmethod def save(config: AppConfig, config_file: str = "config/user_settings.json") -> bool: """ 保存配置到文件 Args: config: 应用配置对象 config_file: 配置文件路径 Returns: 保存是否成功 """ try: # 确保配置目录存在 os.makedirs(os.path.dirname(config_file), exist_ok=True) with open(config_file, "w", encoding="utf-8") as f: json.dump(config.to_dict(), f, ensure_ascii=False, indent=2) return True except IOError as e: print(f"保存配置文件失败: {e}") return False @staticmethod def save_to_env(config: AppConfig, env_file: str = ".env") -> bool: """ 保存配置到 .env 文件 Args: config: 应用配置对象 env_file: .env 文件路径 Returns: 保存是否成功 """ from config.env_loader import save_env_file env_dict = { # ERP 配置 "ERP_URL": config.erp.url, "ERP_USERNAME": config.erp.username, "ERP_PASSWORD": config.erp.password, "ERP_HEADLESS": config.erp.headless, "ERP_IGNORE_HTTPS_ERRORS": config.erp.ignore_https_errors, "ERP_AUTO_CLOSE_BROWSER": config.erp.auto_close_browser, # 数据库配置 "DB_TYPE": ( config.database.db_type.value if isinstance(config.database.db_type, DatabaseType) else config.database.db_type ), "DB_SERVER": config.database.server, "DB_NAME": config.database.database, "DB_USERNAME": config.database.username, "DB_PASSWORD": config.database.password, # SQL Server 特定配置 "DB_SQLSERVER_DRIVER": ( config.database.sqlserver.driver if config.database.sqlserver else "ODBC Driver 18 for SQL Server" ), "DB_TRUST_SERVER_CERTIFICATE": ( config.database.sqlserver.trust_server_certificate if config.database.sqlserver else "yes" ), # MySQL 特定配置 "DB_MYSQL_HOST": ( config.database.mysql.host if config.database.mysql else "" ), "DB_MYSQL_PORT": ( config.database.mysql.port if config.database.mysql else 3306 ), "DB_MYSQL_CHARSET": ( config.database.mysql.charset if config.database.mysql else "utf8mb4" ), # 路径配置 "PATH_DATA_DIR": config.paths.data_dir, "PATH_PRODUCTION_ID_FILE": config.paths.production_id_file, "PATH_DEFAULT_OUTPUT": config.paths.default_output, "PATH_VALIDATION_OUTPUT": config.paths.validation_output, # 数据提取配置 "EXTRACTION_BATCH_SIZE": config.extraction.batch_size, "EXTRACTION_VERBOSE": config.extraction.verbose, "EXTRACTION_AUTO_CONVERT": config.extraction.auto_convert, "EXTRACTION_MERGE_BATCHES": config.extraction.merge_batches, "EXTRACTION_ENABLE_DB_PERSISTENCE": config.extraction.enable_db_persistence, # 校验配置 "VALIDATION_DATA_SOURCE": config.validation.data_source, "VALIDATION_USE_DATABASE": config.validation.use_database, "VALIDATION_BATCH_SIZE": config.validation.batch_size, "VALIDATION_ENABLE_CRUD": config.validation.enable_crud_operations, "VALIDATION_DEFAULT_MANAGER": config.validation.default_manager, "VALIDATION_MATCH_MODE": config.validation.match_mode, } return save_env_file(env_file, env_dict) @staticmethod def _merge_settings(defaults: Dict, loaded: Dict) -> Dict: """ 合并默认配置和加载的配置 Args: defaults: 默认配置 loaded: 加载的配置 Returns: 合并后的配置 """ result = defaults.copy() for key, value in loaded.items(): if ( key in result and isinstance(result[key], dict) and isinstance(value, dict) ): result[key] = ConfigLoader._merge_settings(result[key], value) else: result[key] = value return result @staticmethod def _dict_to_config(settings: Dict) -> AppConfig: """ 将字典转换为配置对象 Args: settings: 配置字典 Returns: 应用配置对象 """ erp_dict = settings.get("erp", {}) database_dict = settings.get("database", {}) paths_dict = settings.get("paths", {}) extraction_dict = settings.get("extraction", {}) validation_dict = settings.get("validation", {}) # 解析数据库类型 db_type_str = database_dict.get("db_type", "sqlserver") try: db_type = DatabaseType(db_type_str) except ValueError: db_type = DatabaseType.SQLSERVER # 解析 SQL Server 配置 sqlserver_dict = database_dict.get("sqlserver", {}) sqlserver_config = SQLServerConfig( driver=sqlserver_dict.get("driver", "ODBC Driver 18 for SQL Server"), trust_server_certificate=sqlserver_dict.get( "trust_server_certificate", "yes" ), ) # 解析 MySQL 配置 mysql_dict = database_dict.get("mysql", {}) mysql_config = MySQLConfig( host=mysql_dict.get("host", database_dict.get("server", "")), port=mysql_dict.get("port", 3306), charset=mysql_dict.get("charset", "utf8mb4"), ) return AppConfig( erp=ERPConfig( url=erp_dict.get("url", ""), username=erp_dict.get("username", ""), password=erp_dict.get("password", ""), headless=erp_dict.get("headless", True), ignore_https_errors=erp_dict.get("ignore_https_errors", True), auto_close_browser=erp_dict.get("auto_close_browser", True), ), database=DatabaseConfig( db_type=db_type, server=database_dict.get("server", ""), database=database_dict.get("database", ""), username=database_dict.get("username", ""), password=database_dict.get("password", ""), sqlserver=sqlserver_config, mysql=mysql_config, ), paths=PathConfig( data_dir=paths_dict.get("data_dir", ""), production_id_file=paths_dict.get("production_id_file", ""), default_output=paths_dict.get( "default_output", "离散备料计划维护_合并.xlsx" ), validation_output=paths_dict.get( "validation_output", "物料状态校验结果.xlsx" ), ), extraction=ExtractionConfig( batch_size=extraction_dict.get("batch_size", 100), verbose=extraction_dict.get("verbose", True), auto_convert=extraction_dict.get("auto_convert", True), merge_batches=extraction_dict.get("merge_batches", True), enable_db_persistence=extraction_dict.get( "enable_db_persistence", False ), ), validation=ValidationConfig( data_source=validation_dict.get("data_source", "database_full"), use_database=validation_dict.get("use_database", True), batch_size=validation_dict.get("batch_size", 2000), enable_crud_operations=validation_dict.get( "enable_crud_operations", False ), default_manager=validation_dict.get("default_manager", ""), match_mode=validation_dict.get("match_mode", "substring"), ), )