Apply Black formatter to the entire codebase for consistent code style. Co-Authored-By: Claude (glm-5) <noreply@anthropic.com>
284 lines
10 KiB
Python
284 lines
10 KiB
Python
#!/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"),
|
||
),
|
||
)
|