"""核心配置层:加载根目录 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 config_path: Path = CONFIG_PATH env_path: Path = ENV_PATH project_root: Path = PROJECT_ROOT 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()}" @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" database_url = os.environ.get(url_env) or default_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 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, )