"""核心配置层:加载根目录 config.yaml + .env,向全应用提供 Settings。 加载优先级:环境变量 > 根目录 .env > 根目录 config.yaml > 代码默认值。 密钥只存在于 .env(config.yaml 经 *_env 字段引用变量名)。 """ from __future__ import annotations import os from dataclasses import dataclass from functools import lru_cache from pathlib import Path import yaml # 项目根:/backend/app/core/config.py → parents[3] PROJECT_ROOT = Path(__file__).resolve().parents[3] BACKEND_ROOT = PROJECT_ROOT / "backend" CONFIG_PATH = PROJECT_ROOT / "config.yaml" ENV_PATH = PROJECT_ROOT / ".env" _SQLITE_PREFIX = "sqlite:///" def _load_env_file(path: Path) -> None: """把 .env 读入 os.environ(不覆盖已存在的环境变量)。""" if not path.exists(): return for raw in path.read_text(encoding="utf-8").splitlines(): line = raw.strip() if not line or line.startswith("#") or "=" not in line: continue key, _, value = line.partition("=") key = key.strip() if key and key not in os.environ: os.environ[key] = value.strip().strip("\"'") def _load_config_yaml(path: Path) -> dict: if not path.exists(): return {} with path.open(encoding="utf-8") as fh: data = yaml.safe_load(fh) return data if isinstance(data, dict) else {} def _deep(cfg: dict, dotted: str) -> object: node: object = cfg for part in dotted.split("."): if not isinstance(node, dict) or part not in node: return None node = node[part] return node @dataclass(frozen=True) class Settings: """应用全局配置(已解析、不可变)。""" app_name: str app_version: str debug: bool secret_key: str api_prefix: str database_url: str sqlalchemy_echo: bool migrations_dir: Path data_source_primary: str data_source_fallback: str tushare_token: str llm_api_key: str llm_base_url: str | None llm_model: str storage: dict[str, Path] job_mode: str job_memory_limit_gb: int job_max_concurrent: int # research.archive_curve_limit:归档个股收益曲线数量上限;None = 完整存档(默认) research_archive_curve_limit: int | None # research.archive_max_chars:归档结果 JSON 的字节预算(MEDIUMTEXT 上限 16MB 的余量) research_archive_max_chars: int | None config_path: Path = CONFIG_PATH env_path: Path = ENV_PATH project_root: Path = PROJECT_ROOT def _env_or(value_env: str, configured: object) -> object: """环境变量优先于 config.yaml(空串视为未设置)。""" env = os.environ.get(value_env) if env not in (None, ""): return env return configured def _optional_int(value: object) -> int | None: """可选整数配置:None / 空串 / 非法值都视为「未配置」(返回 None)。 用于 `research.archive_curve_limit`:None 表示不截断(完整存档)。 """ if value is None or value == "": return None try: return int(value) # type: ignore[arg-type] except (TypeError, ValueError): return None def _resolve_storage_dirs(cfg: dict) -> dict[str, Path]: storage_cfg = _deep(cfg, "storage") or {} return { key: PROJECT_ROOT / str(value) for key, value in storage_cfg.items() if isinstance(value, str) } def _normalize_sqlite_url(url: str) -> str: """把相对路径的 sqlite URL 解析为项目根下的绝对路径。""" if not url.startswith(_SQLITE_PREFIX): return url rest = url[len(_SQLITE_PREFIX) :] if rest.startswith("/"): # 已是绝对路径(sqlite:////abs/path) return url return f"{_SQLITE_PREFIX}{(PROJECT_ROOT / rest).resolve()}" _URL_RESERVED = set("@:/?#%") def _quote_userinfo(value: str) -> str: """仅编码会破坏 SQLAlchemy URL 解析的字符(@ : / ? # % 与空白)。 其余字符(含 ! - _ . 等非保留符)原样保留:避免 URL 出现 %XX 干扰 Alembic configparser 的 interpolation。 """ out: list[str] = [] for ch in value: if ch in _URL_RESERVED or ch.isspace(): out.append(f"%{ord(ch):02X}") elif ord(ch) > 127: # 非 ASCII:按 UTF-8 逐字节 percent 编码 out.append("".join(f"%{b:02X}" for b in ch.encode("utf-8"))) else: out.append(ch) return "".join(out) def _build_mysql_url(mysql: dict | None) -> str | None: """由 config.yaml database.mysql 段组装 mysql+pymysql URL。 规则:仅当 enabled 且 host/user/db 齐全时返回 URL;密码经 password_env 指定的环境变量读取(AGENT.md §33 密钥只放 .env),未设置则按空密码处理。 """ if not mysql or not mysql.get("enabled"): return None host = mysql.get("host") user = mysql.get("user") db = mysql.get("db") if not (host and user and db): return None port = mysql.get("port") or 3306 charset = mysql.get("charset") or "utf8mb4" password = os.environ.get(mysql.get("password_env") or "MYSQL_PASSWORD", "") auth = ( f"{_quote_userinfo(user)}:{_quote_userinfo(password)}" if password else _quote_userinfo(user) ) return f"mysql+pymysql://{auth}@{host}:{port}/{db}?charset={charset}" # 被禁止的数据库目标主机(用户明确要求:禁止使用 192.168.1.10 作为 DB 目标)。 # 可用环境变量 QLIB_FORBIDDEN_DB_HOSTS 覆盖(逗号分隔;设为空串表示不做限制)。 _FORBIDDEN_DB_HOSTS_ENV = "QLIB_FORBIDDEN_DB_HOSTS" _FORBIDDEN_DB_HOSTS_DEFAULT = "192.168.1.10" def _forbidden_db_hosts() -> set[str]: raw = os.environ.get(_FORBIDDEN_DB_HOSTS_ENV) if raw is None: raw = _FORBIDDEN_DB_HOSTS_DEFAULT return {h.strip().lower() for h in raw.split(",") if h.strip()} def assert_db_target_allowed(database_url: str) -> None: """拒绝把数据库指向被禁主机;命中即**抛错**,不允许「起得来但连错库」。 为什么必须硬失败(AGENT.md §7 不静默):库目标错了不会有任何报错或界面异常—— 回测结果、实验归档、策略、信号会**安静地读写另一台机器上的数据**, 而用户以为看的是本机库;这类错误事后极难发现(数据看似正常,只是「不对」)。 因此在这里直接失败并说明原因与修改方法。 只拦主机名/IP 精确匹配(不做网段推断):`192.168.1.10` 与写进 userinfo 的 同名字符串、SQLite 路径都不会误判;IPv6/带方括号的地址会去掉括号后比较。 """ forbidden = _forbidden_db_hosts() if not forbidden or not database_url: return try: from sqlalchemy.engine import make_url url = make_url(database_url) except Exception: # noqa: BLE001 —— 解析失败交给 SQLAlchemy 自己报错,这里不抢 return host = (url.host or "").lower().strip("[]") if host and host in forbidden: raise RuntimeError( f"数据库目标 {host} 已被禁止({_FORBIDDEN_DB_HOSTS_ENV}=" f"{','.join(sorted(forbidden))})。" "本项目只允许使用**本机 MariaDB**:把 config.yaml 的 database.mysql.host " "设为 127.0.0.1,或清空 .env 的 DATABASE_URL 让它走本机配置。" "确需临时放开请显式设置 QLIB_FORBIDDEN_DB_HOSTS(设为空串表示不限制)。" ) @lru_cache def get_settings() -> Settings: """加载并缓存 Settings。config / env 路径可通过参数覆盖以便测试。""" _load_env_file(ENV_PATH) cfg = _load_config_yaml(CONFIG_PATH) secret_env = _deep(cfg, "app.secret_key_env") or "APP_SECRET_KEY" url_env = _deep(cfg, "database.url_env") or "DATABASE_URL" token_env = _deep(cfg, "data_source.tushare_token_env") or "TUSHARE_TOKEN" # Agent LLM:URL/模型名在 config.yaml 明文,api_key 只从 .env 读(env 可覆盖 url/model) agent_llm = _deep(cfg, "agent.llm") or {} llm_key_env = agent_llm.get("api_key_env") or "LLM_API_KEY" llm_url_env = agent_llm.get("base_url_env") or "LLM_BASE_URL" llm_model_env = agent_llm.get("model_env") or "LLM_MODEL" default_url = "sqlite:///./data/quant.db" # URL 优先级:环境变量(DATABASE_URL) > config.yaml database.mysql 段 > SQLite 兜底 database_url = os.environ.get(url_env) or _build_mysql_url( _deep(cfg, "database.mysql") or {} ) or default_url # 用户约束:禁止把数据写向 192.168.1.10(远端库)——命中直接抛错,不静默降级 assert_db_target_allowed(database_url) migrations_rel = _deep(cfg, "database.migrations_dir") or ( "app/infrastructure/persistence/migrations" ) # Job 执行:环境变量 JOB_MODE / QLIB_JOB_MEM_LIMIT_GB 可覆盖 config.yaml(测试用) job_cfg = _deep(cfg, "job") or {} try: job_memory_limit_gb = int( os.environ.get("QLIB_JOB_MEM_LIMIT_GB") or job_cfg.get("max_memory_gb") or 6 ) except (TypeError, ValueError): job_memory_limit_gb = 6 try: job_max_concurrent = int( os.environ.get("QLIB_JOB_MAX_CONCURRENT") or job_cfg.get("max_concurrent_jobs") or 2 ) except (TypeError, ValueError): job_max_concurrent = 2 # 归档(research 段):曲线数量上限(None = 完整存档)+ 结果 JSON 字节预算 research_cfg = _deep(cfg, "research") research_cfg = research_cfg if isinstance(research_cfg, dict) else {} archive_curve_limit = _optional_int( _env_or("QLIB_ARCHIVE_CURVE_LIMIT", research_cfg.get("archive_curve_limit")) ) archive_max_chars = _optional_int( _env_or("QLIB_ARCHIVE_MAX_CHARS", research_cfg.get("archive_max_chars")) ) return Settings( app_name=str(_deep(cfg, "app.name") or "qlib-platform"), app_version=str(_deep(cfg, "app.version") or "0.1.0"), debug=bool(_deep(cfg, "app.debug") or False), secret_key=os.environ.get(secret_env, "dev-insecure-key"), api_prefix=str(_deep(cfg, "api.prefix") or "/api"), database_url=_normalize_sqlite_url(database_url), sqlalchemy_echo=bool(_deep(cfg, "database.echo") or False), migrations_dir=BACKEND_ROOT / str(migrations_rel), data_source_primary=str(_deep(cfg, "data_source.primary") or "tushare"), data_source_fallback=str(_deep(cfg, "data_source.fallback") or "sina"), tushare_token=os.environ.get(token_env, ""), llm_api_key=os.environ.get(llm_key_env, ""), llm_base_url=os.environ.get(llm_url_env) or (agent_llm.get("base_url") or None), llm_model=os.environ.get(llm_model_env) or agent_llm.get("model") or "qwen-plus", storage=_resolve_storage_dirs(cfg), job_mode=os.environ.get("JOB_MODE") or job_cfg.get("mode") or "subprocess", job_memory_limit_gb=job_memory_limit_gb, job_max_concurrent=job_max_concurrent, research_archive_curve_limit=archive_curve_limit, research_archive_max_chars=archive_max_chars, )