Files
playwrite/config/loader.py
Misaka 3b7c00377f style: format all Python files with Black
Apply Black formatter to the entire codebase for consistent code style.

Co-Authored-By: Claude (glm-5) <noreply@anthropic.com>
2026-02-26 22:44:03 +08:00

284 lines
10 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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"),
),
)