Files
qlib/backend/app/core/config.py
T
Simon d9be75a98f feat: Phase 5 — AI Research Agent(受控工具白名单 + LLM 编排 + API)
- agent/tools.py:Tool 元数据(JSON Schema)+ 白名单调用(异常转可读反馈,不中断对话)
- agent/tools_impl.py:6 个受控工具 search_stocks / get_market_data / test_factor / run_backtest / get_experiment / compare_experiments —— 全部只读经 Job/Experiment 链路,研究自动归档;无 shell/任意执行/写删数据能力
- agent/llm.py:LLMClient 抽象 + OpenAI 兼容客户端(LLM_API_KEY/LLM_BASE_URL/LLM_MODEL 走 .env,未配置给出引导提示)+ 研究纪律 system prompt(反过拟合/样本外/成本)
- agent/service.py:编排循环(tool/final JSON 决策 → 执行 → 回喂 → 结论),轮次上限兜底,未知工具拒绝
- /api/agent/chat;httpx 移至主依赖;Job 默认工厂抽取(api/agent/executor 复用)
- 测试 6 项(白名单无 shell、完整研究循环产出、未知工具拒绝、轮次兜底),全量 79 passed / ruff clean
2026-09-06 17:22:05 +08:00

136 lines
4.5 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]
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"
llm_key_env = _deep(cfg, "agent.llm_key_env") or "LLM_API_KEY"
llm_url_env = _deep(cfg, "agent.llm_base_url_env") or "LLM_BASE_URL"
llm_model_env = _deep(cfg, "agent.llm_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"
)
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 None,
llm_model=os.environ.get(llm_model_env) or "qwen-plus",
storage=_resolve_storage_dirs(cfg),
)