Files
qlib/backend/app/core/config.py
T
Simon 6c2f198261 feat(db): SQLite 全量迁移至 MySQL(config.yaml 配置化 + 迁移脚本 + 一致性校验)
- config.yaml database.mysql:host/port/db/user/charset 明文可提交;密码经 password_env
  引用 .env 的 MYSQL_PASSWORD(AGENT.md §33 密钥不进 git)
- config.py _build_mysql_url:URL 优先级 DATABASE_URL env > database.mysql 段 > sqlite 兜底
- pyproject 引入 pymysql>=1.1
- tests/conftest.py 强制每进程 /tmp SQLite(测试绝不触 MySQL 开发库);test_config 覆盖
  mysql 组装/密码可选/sqlite 兜底分支
- scripts/migrate_sqlite_to_mysql.py:sqlite→mysql 一次性迁移工具(keyset 分页 + chunk
  多值 INSERT + 攒批 commit + 幂等续传 + 低配 MySQL 节流 --throttle-sec + --verify-only
  行数与抽样一致性校验);已用于 data/quant.db 约 1605 万行迁移并经校验一致
2026-09-08 23:58:48 +08:00

183 lines
6.6 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
from urllib.parse import quote_plus
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()}"
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_plus(user)}:{quote_plus(password)}" if password else quote_plus(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,
)