"""网关配置加载与校验。 从 gateway/config.json 读取配置,缺字段给默认值,启动时校验必填项。 """ import json import logging from pathlib import Path from typing import List logger = logging.getLogger("gateway.config") CONFIG_PATH = Path(__file__).resolve().parent / "config.json" DEFAULTS = { "public_base_url": "https://hair.xiangsilian.com", "static_dir": "static/annotations", "health_check": { "path": "/health", "interval_seconds": 8, "timeout_seconds": 3, "unhealthy_threshold": 2, "healthy_threshold": 1, }, "dispatch": { "per_worker_concurrency": 1, "queue_wait_seconds": 30, "request_timeout_seconds": 60, "retry_on_failure": True, "max_retries": 1, }, "cleanup": { "interval_minutes": 60, "max_age_hours": 24, }, } _config_cache = None def _deep_merge(defaults: dict, overrides: dict) -> dict: """递归合并:overrides 中的值覆盖 defaults,嵌套 dict 递归处理。""" result = defaults.copy() for key, value in overrides.items(): if key in result and isinstance(result[key], dict) and isinstance(value, dict): result[key] = _deep_merge(result[key], value) else: result[key] = value return result def load_config() -> dict: """加载并校验配置文件,结果缓存仅加载一次。""" global _config_cache if _config_cache is not None: return _config_cache if not CONFIG_PATH.exists(): raise FileNotFoundError( f"配置文件不存在: {CONFIG_PATH}\n" f"请从 gateway/config.example.json 复制并修改: " f"cp gateway/config.example.json gateway/config.json" ) with open(CONFIG_PATH, "r", encoding="utf-8") as f: raw = json.load(f) # 合并默认值 cfg = _deep_merge(DEFAULTS, raw) # --- 校验 --- workers: List[str] = cfg.get("workers", []) if not workers: raise ValueError("配置错误: workers 列表不能为空,至少需要一个 worker 地址") for i, w in enumerate(workers): if not isinstance(w, str) or not w: raise ValueError(f"配置错误: workers[{i}] 必须是非空字符串") password: str = cfg.get("shared_password", "") if not password or password == "REPLACE_ME_ROTATE_PERIODICALLY": logger.warning( "⚠ 安全警告: shared_password 未设置或仍为占位值 " "'REPLACE_ME_ROTATE_PERIODICALLY',请立即更换!" ) if len(password) < 8: logger.warning("⚠ 安全警告: shared_password 长度不足 8 位,建议使用更长的密码") public_base_url: str = cfg.get("public_base_url", "") if public_base_url.endswith("/"): cfg["public_base_url"] = public_base_url.rstrip("/") logger.warning("public_base_url 末尾含 '/',已自动去除") # 确保 static_dir 是绝对路径 static_dir = Path(cfg["static_dir"]) if not static_dir.is_absolute(): cfg["static_dir"] = str(Path(__file__).resolve().parents[1] / static_dir) logger.info( "配置加载完成 | workers=%s | public_base_url=%s | " "hc_interval=%ds | dispatch_timeout=%ds | queue_wait=%ds", cfg["workers"], cfg["public_base_url"], cfg["health_check"]["interval_seconds"], cfg["dispatch"]["request_timeout_seconds"], cfg["dispatch"]["queue_wait_seconds"], ) _config_cache = cfg return cfg def get_config() -> dict: """获取已加载的配置(必须先调用 load_config)。""" if _config_cache is None: raise RuntimeError("配置尚未加载,请先调用 load_config()") return _config_cache def reload_config() -> dict: """强制重新加载配置(用于热更新)。""" global _config_cache _config_cache = None return load_config()