Files
qlib/backend/app/core/config.py
T
Simon c60dc78c88 feat(selection): M6.3 选股结果落库 + /api/selections(快照 + 逐候选行)
- domain/repositories/selection.py + SelectionMeta:选股 Repository Protocol
- models/selection.py + migration a6c91d4e7f20:selection_snapshot(查询/统计快照)+
  selection_result(逐候选:rank/score/factor_values/filter_status/reason JSON)
- repositories/selection_impl.py:save/get/list_recent(读回重建 SelectionResult)
- api/selections.py:POST /api/selections(同步执行+落库,返回 selection_id+result)、
  GET /{id}、GET 列表(as_of/method 过滤);deps 装配 SelectionService/Repo
- config._build_mysql_url:仅编码破坏 URL 结构的字符(修 Alembic configparser 遇 %21 崩)
- 迁移已在 MySQL 应用(alembic head a6c91d4e7f20);tests/test_selections_api.py 4 例
  提交→读回一致/404/列表过滤/condition;全量 pytest 通过
2026-09-09 00:20:42 +08:00

206 lines
7.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""核心配置层:加载根目录 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
# 项目根:<root>/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()}"
_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}"
@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
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,
)