# -*- coding: utf-8 -*- """统一配置加载器。 加载优先级(高 -> 低): 1. 进程环境变量 2. config/.env (若存在,仅本地开发用) 3. config/*.yaml 默认值 YAML 中可使用 ${VAR:default} 占位符引用环境变量,未设置时取 default。 Phase 1 目标:消除硬编码 IP/密码/端口,运行时业务路径全部经此模块读取配置。 """ from __future__ import absolute_import import os import re import logging try: import yaml except ImportError: # pragma: no cover - 依赖缺失时给出明确错误 raise ImportError("PyYAML is required: pip install pyyaml") __all__ = ["get_config", "get_settings", "get_redis_config", "get_db_config", "get_models_config"] _CONFIG_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "config") _CACHE = {} _PLACEHOLDER_RE = re.compile(r"\$\{([A-Z0-9_]+)(?::([^}]*))?\}") def _load_env_file(): """读取 config/.env,把键值写入 os.environ(已存在的环境变量优先)。""" env_path = os.path.join(_CONFIG_DIR, ".env") if not os.path.exists(env_path): return with open(env_path, "r", encoding="utf-8") as f: for raw in f: line = raw.strip() if not line or line.startswith("#") or "=" not in line: continue key, _, value = line.partition("=") key = key.strip() value = value.strip().strip('"').strip("'") if key and key not in os.environ: os.environ[key] = value def _expand_placeholders(value): """递归展开 ${VAR:default} 占位符。""" if isinstance(value, dict): return {k: _expand_placeholders(v) for k, v in value.items()} if isinstance(value, list): return [_expand_placeholders(v) for v in value] if isinstance(value, str): def _sub(m): var, default = m.group(1), m.group(2) or "" return os.environ.get(var, default) # 同一字符串内可能有多个占位符,循环到稳定 prev = None cur = value while prev != cur: prev = cur cur = _PLACEHOLDER_RE.sub(_sub, cur) return cur return value def _read_yaml(name): path = os.path.join(_CONFIG_DIR, name) if not os.path.exists(path): return {} with open(path, "r", encoding="utf-8") as f: data = yaml.safe_load(f) or {} return _expand_placeholders(data) def get_config(name): """读取 config/.yaml,带缓存。""" if name not in _CACHE: _load_env_file() _CACHE[name] = _read_yaml("%s.yaml" % name) return _CACHE[name] def get_settings(): return get_config("settings") def get_redis_config(): return get_config("redis") def get_db_config(): return get_config("db") def get_models_config(): return get_config("models") def reload(): """清空缓存,用于测试或热加载。""" _CACHE.clear() # ===== 常用便捷读取 ===== def server_port(name="web"): """获取服务端口。优先环境变量 BIDDINGKG_WEB_PORT / BIDDINGKG_SINGLE_PORT。""" env_key = "BIDDINGKG_%s_PORT" % name.upper() if env_key in os.environ: return int(os.environ[env_key]) settings = get_settings() port_field = "web_port" if name == "web" else "single_port" return int(settings.get("server", {}).get(port_field, 15014)) def model_input_shape(name): """获取模型输入 shape,name ∈ role/money/person。""" settings = get_settings() shapes = settings.get("model_input_shape", {}) if name not in shapes: raise KeyError("model_input_shape.%s not configured" % name) return tuple(shapes[name]) def pai_eas_settings(): """返回 (use_pai_eas, api_url, use_api)。""" env_flag = os.environ.get("BIDDINGKG_USE_PAI_EAS", "").lower() if env_flag in ("1", "true", "yes", "on"): use_pai = True elif env_flag in ("0", "false", "no", "off"): use_pai = False else: use_pai = bool(get_settings().get("pai_eas", {}).get("enabled", False)) pai_cfg = get_settings().get("pai_eas", {}) return use_pai, pai_cfg.get("api_url", ""), bool(pai_cfg.get("use_api", False)) def internal_api_url(key="article_extract_url"): """获取内部 API URL。优先环境变量。""" settings = get_settings() return settings.get("internal_api", {}).get(key, "") def logging_settings(): """返回日志配置 dict:level/format/date_format/file。""" return get_settings().get("logging", {}) def redis_main_config(): """主 Redis 连接参数(host/port/db/password)。""" cfg = get_redis_config() main = dict(cfg.get("main", {})) env_host = os.environ.get("BIDDINGKG_REDIS_HOST") env_pass = os.environ.get("BIDDINGKG_REDIS_PASS") if env_host: main["host"] = env_host if env_pass is not None: main["password"] = env_pass return main def redis_queue_names(): """队列名 dict:content/preprocess/preprocess_legacy_alias/error/timeout。""" return dict(get_redis_config().get("queues", {})) def redis_consumer_params(): """消费者参数:batch_size/server_sleep/client_sleep/timeout。""" return dict(get_redis_config().get("consumer", {})) def pg_db_config(dbname="BiddingKG"): """返回指定库的 PG 连接参数 dict。""" cfg = get_db_config() default = dict(cfg.get("default", {})) dbs = cfg.get("databases", {}) if dbname not in dbs: raise KeyError("db.databases.%s not configured" % dbname) merged = dict(default) merged.update(dbs[dbname]) env_host = os.environ.get("BIDDINGKG_PG_HOST") env_user = os.environ.get("BIDDINGKG_PG_USER") env_pass = os.environ.get("BIDDINGKG_PG_PASS") if env_host: merged["host"] = env_host if env_user: merged["user"] = env_user if env_pass is not None: merged["password"] = env_pass return merged def mongo_uri(): """返回 MongoDB URI,优先环境变量 BIDDINGKG_MONGO_URI。""" return os.environ.get("BIDDINGKG_MONGO_URI", "")