config.py 6.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198
  1. # -*- coding: utf-8 -*-
  2. """统一配置加载器。
  3. 加载优先级(高 -> 低):
  4. 1. 进程环境变量
  5. 2. config/.env (若存在,仅本地开发用)
  6. 3. config/*.yaml 默认值
  7. YAML 中可使用 ${VAR:default} 占位符引用环境变量,未设置时取 default。
  8. Phase 1 目标:消除硬编码 IP/密码/端口,运行时业务路径全部经此模块读取配置。
  9. """
  10. from __future__ import absolute_import
  11. import os
  12. import re
  13. import logging
  14. try:
  15. import yaml
  16. except ImportError: # pragma: no cover - 依赖缺失时给出明确错误
  17. raise ImportError("PyYAML is required: pip install pyyaml")
  18. __all__ = ["get_config", "get_settings", "get_redis_config", "get_db_config", "get_models_config"]
  19. _CONFIG_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "config")
  20. _CACHE = {}
  21. _PLACEHOLDER_RE = re.compile(r"\$\{([A-Z0-9_]+)(?::([^}]*))?\}")
  22. def _load_env_file():
  23. """读取 config/.env,把键值写入 os.environ(已存在的环境变量优先)。"""
  24. env_path = os.path.join(_CONFIG_DIR, ".env")
  25. if not os.path.exists(env_path):
  26. return
  27. with open(env_path, "r", encoding="utf-8") as f:
  28. for raw in f:
  29. line = raw.strip()
  30. if not line or line.startswith("#") or "=" not in line:
  31. continue
  32. key, _, value = line.partition("=")
  33. key = key.strip()
  34. value = value.strip().strip('"').strip("'")
  35. if key and key not in os.environ:
  36. os.environ[key] = value
  37. def _expand_placeholders(value):
  38. """递归展开 ${VAR:default} 占位符。"""
  39. if isinstance(value, dict):
  40. return {k: _expand_placeholders(v) for k, v in value.items()}
  41. if isinstance(value, list):
  42. return [_expand_placeholders(v) for v in value]
  43. if isinstance(value, str):
  44. def _sub(m):
  45. var, default = m.group(1), m.group(2) or ""
  46. return os.environ.get(var, default)
  47. # 同一字符串内可能有多个占位符,循环到稳定
  48. prev = None
  49. cur = value
  50. while prev != cur:
  51. prev = cur
  52. cur = _PLACEHOLDER_RE.sub(_sub, cur)
  53. return cur
  54. return value
  55. def _read_yaml(name):
  56. path = os.path.join(_CONFIG_DIR, name)
  57. if not os.path.exists(path):
  58. return {}
  59. with open(path, "r", encoding="utf-8") as f:
  60. data = yaml.safe_load(f) or {}
  61. return _expand_placeholders(data)
  62. def get_config(name):
  63. """读取 config/<name>.yaml,带缓存。"""
  64. if name not in _CACHE:
  65. _load_env_file()
  66. _CACHE[name] = _read_yaml("%s.yaml" % name)
  67. return _CACHE[name]
  68. def get_settings():
  69. return get_config("settings")
  70. def get_redis_config():
  71. return get_config("redis")
  72. def get_db_config():
  73. return get_config("db")
  74. def get_models_config():
  75. return get_config("models")
  76. def reload():
  77. """清空缓存,用于测试或热加载。"""
  78. _CACHE.clear()
  79. # ===== 常用便捷读取 =====
  80. def server_port(name="web"):
  81. """获取服务端口。优先环境变量 BIDDINGKG_WEB_PORT / BIDDINGKG_SINGLE_PORT。"""
  82. env_key = "BIDDINGKG_%s_PORT" % name.upper()
  83. if env_key in os.environ:
  84. return int(os.environ[env_key])
  85. settings = get_settings()
  86. port_field = "web_port" if name == "web" else "single_port"
  87. return int(settings.get("server", {}).get(port_field, 15014))
  88. def model_input_shape(name):
  89. """获取模型输入 shape,name ∈ role/money/person。"""
  90. settings = get_settings()
  91. shapes = settings.get("model_input_shape", {})
  92. if name not in shapes:
  93. raise KeyError("model_input_shape.%s not configured" % name)
  94. return tuple(shapes[name])
  95. def pai_eas_settings():
  96. """返回 (use_pai_eas, api_url, use_api)。"""
  97. env_flag = os.environ.get("BIDDINGKG_USE_PAI_EAS", "").lower()
  98. if env_flag in ("1", "true", "yes", "on"):
  99. use_pai = True
  100. elif env_flag in ("0", "false", "no", "off"):
  101. use_pai = False
  102. else:
  103. use_pai = bool(get_settings().get("pai_eas", {}).get("enabled", False))
  104. pai_cfg = get_settings().get("pai_eas", {})
  105. return use_pai, pai_cfg.get("api_url", ""), bool(pai_cfg.get("use_api", False))
  106. def internal_api_url(key="article_extract_url"):
  107. """获取内部 API URL。优先环境变量。"""
  108. settings = get_settings()
  109. return settings.get("internal_api", {}).get(key, "")
  110. def logging_settings():
  111. """返回日志配置 dict:level/format/date_format/file。"""
  112. return get_settings().get("logging", {})
  113. def redis_main_config():
  114. """主 Redis 连接参数(host/port/db/password)。"""
  115. cfg = get_redis_config()
  116. main = dict(cfg.get("main", {}))
  117. env_host = os.environ.get("BIDDINGKG_REDIS_HOST")
  118. env_pass = os.environ.get("BIDDINGKG_REDIS_PASS")
  119. if env_host:
  120. main["host"] = env_host
  121. if env_pass is not None:
  122. main["password"] = env_pass
  123. return main
  124. def redis_queue_names():
  125. """队列名 dict:content/preprocess/preprocess_legacy_alias/error/timeout。"""
  126. return dict(get_redis_config().get("queues", {}))
  127. def redis_consumer_params():
  128. """消费者参数:batch_size/server_sleep/client_sleep/timeout。"""
  129. return dict(get_redis_config().get("consumer", {}))
  130. def pg_db_config(dbname="BiddingKG"):
  131. """返回指定库的 PG 连接参数 dict。"""
  132. cfg = get_db_config()
  133. default = dict(cfg.get("default", {}))
  134. dbs = cfg.get("databases", {})
  135. if dbname not in dbs:
  136. raise KeyError("db.databases.%s not configured" % dbname)
  137. merged = dict(default)
  138. merged.update(dbs[dbname])
  139. env_host = os.environ.get("BIDDINGKG_PG_HOST")
  140. env_user = os.environ.get("BIDDINGKG_PG_USER")
  141. env_pass = os.environ.get("BIDDINGKG_PG_PASS")
  142. if env_host:
  143. merged["host"] = env_host
  144. if env_user:
  145. merged["user"] = env_user
  146. if env_pass is not None:
  147. merged["password"] = env_pass
  148. return merged
  149. def mongo_uri():
  150. """返回 MongoDB URI,优先环境变量 BIDDINGKG_MONGO_URI。"""
  151. return os.environ.get("BIDDINGKG_MONGO_URI", "")