| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198 |
- # -*- 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/<name>.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", "")
|