feat(llm): 配置三层解析+config_store 双表——resolve_config 增 overrides 层(DB>env>默认,overrides=None 零行为变化),pipeline.db 新 llm_config/audit 表;api_key 打码出库/掩码串不回写/审计只存打码形 (spec §13 P5.1) [vps] [no-doc]

This commit is contained in:
2026-10-05 09:18:02 +08:00
parent 14f30a00f8
commit 63fdf27d78
4 changed files with 359 additions and 7 deletions
+177
View File
@@ -0,0 +1,177 @@
# sanguo_api/llm/config_store.py
"""LLM 连接 DB 配置(P5.1,spec §13):pipeline.db 两表,仿 gate_config+audit.
三层解析=DB 行 > env > 代码默认;env 保持首次部署兜底通道(治 VPS 手配
env 痛点)。key 集=provider/base_url/model/api_key/timeout/enabled,value
全 TEXT(enabled 存 "1"/"0")。api_key 绝不明文出 GET:读侧打码、写侧含
'*' 串跳过(防掩码回显覆盖真值)、审计行只存打码形。
"""
from __future__ import annotations
import os
import sqlite3
from typing import Any, Mapping
from .provider import (
LLMConfig, LLMConfigError, PROVIDERS, resolve_config)
LLM_KEYS = ("provider", "base_url", "model", "api_key", "timeout", "enabled")
_SCHEMA = """
CREATE TABLE IF NOT EXISTS llm_config (
key TEXT PRIMARY KEY,
value TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS llm_config_audit (
id INTEGER PRIMARY KEY AUTOINCREMENT,
key TEXT NOT NULL,
old_value TEXT,
new_value TEXT NOT NULL,
changed_by TEXT NOT NULL,
changed_at TEXT NOT NULL DEFAULT (datetime('now', 'localtime'))
);
"""
_MASK_HEAD, _MASK_TAIL = 4, 4
def mask_secret(value: str) -> str:
"""打码形:头 4+****+尾 4;短值全掩."""
if len(value) <= _MASK_HEAD + _MASK_TAIL:
return "****"
return f"{value[:_MASK_HEAD]}****{value[-_MASK_TAIL:]}"
def _connect(path: str) -> sqlite3.Connection:
"""同 pipeline_store._connect 纪律:幂等建表+busy_timeout(并发写防锁)."""
d = os.path.dirname(os.path.abspath(path))
os.makedirs(d, exist_ok=True)
conn = sqlite3.connect(path, timeout=30)
conn.execute("PRAGMA busy_timeout=30000")
conn.row_factory = sqlite3.Row
conn.executescript(_SCHEMA)
return conn
def get_llm_overrides(path: str) -> dict[str, str]:
"""DB 全行原文(api_key 明文,仅进程内解析用,绝不出 API)."""
with _connect(path) as conn:
rows = conn.execute("SELECT key, value FROM llm_config").fetchall()
return {r["key"]: r["value"] for r in rows}
def resolve_effective_config(db_path: str,
env: Mapping[str, str] = os.environ) -> LLMConfig:
"""draft/decompose 唯一解析入口(P5.1 三层):DB>env>默认."""
return resolve_config(env, get_llm_overrides(db_path))
def get_masked_config(path: str,
env: Mapping[str, str] = os.environ) -> dict:
"""设置页视图:生效值+逐字段来源(db/env/default)+resolved/missing.
来源判定与 resolve_config 同序;api_key 打码,未配置任一层时空串
(前端 placeholder 提示去 env/设置页配)。
"""
ov = get_llm_overrides(path)
def pick(key: str, env_name: str, default: str = "") -> tuple[str, str]:
if ov.get(key):
return ov[key], "db"
if (env.get(env_name) or "").strip():
return env[env_name].strip(), "env"
return default, "default"
provider, p_src = pick("provider", "SANGUO_LLM_PROVIDER", "glm_cn")
spec = PROVIDERS.get(provider) or PROVIDERS["glm_cn"]
base_url, b_src = pick("base_url", "SANGUO_LLM_BASE_URL", spec.base_url)
model, m_src = pick("model", "SANGUO_LLM_MODEL", spec.default_model)
api_key, k_src = pick("api_key", spec.api_key_env)
timeout, t_src = pick("timeout", "SANGUO_LLM_TIMEOUT", "60.0")
try:
resolve_effective_config(path, env)
resolved, missing = True, []
except LLMConfigError as e:
resolved, missing = False, [str(e)]
return {
"config": {"provider": provider, "base_url": base_url,
"model": model,
"api_key": mask_secret(api_key) if api_key else "",
"timeout": float(timeout),
"enabled": ov.get("enabled") not in ("0", "false", "False")},
"sources": {"provider": p_src, "base_url": b_src, "model": m_src,
"api_key": k_src, "timeout": t_src,
"enabled": "db" if "enabled" in ov else "default"},
"resolved": resolved, "missing": missing,
}
def set_llm_config(path: str, values: Mapping[str, Any],
changed_by: str) -> list[dict[str, Any]]:
"""写 DB 覆盖+审计;返回实际变更行(供 PUT 响应).
规则:未知键 ValueError;provider 须在 PROVIDERS;timeout 须 float∈
[1,600];enabled 归一 "1"/"0"(空=清键回落默认启用);api_key 含 '*'
(掩码回显)静默跳过;空串=清该键回落 env/默认(DELETE);值未变不写不审计。
"""
unknown = [k for k in values if k not in LLM_KEYS]
if unknown:
raise ValueError(f"未知 llm 配置键: {unknown}(合法集={list(LLM_KEYS)})")
if values.get("provider") and values["provider"] not in PROVIDERS:
raise ValueError(f"未知 provider: {values['provider']!r}"
f"(可用: {sorted(PROVIDERS)})")
if values.get("timeout") not in (None, ""):
try:
t = float(values["timeout"])
except (TypeError, ValueError) as e:
raise ValueError(
f"timeout 须为数字: {values['timeout']!r}") from e
if not 1.0 <= t <= 600.0:
raise ValueError("timeout 须在 [1, 600] 秒")
rows: list[dict[str, Any]] = []
with _connect(path) as conn:
cur = {r["key"]: r["value"] for r in conn.execute(
"SELECT key, value FROM llm_config").fetchall()}
for key in LLM_KEYS:
if key not in values:
continue
raw = values[key]
if key == "api_key" and isinstance(raw, str) and "*" in raw:
continue # 掩码回显不回写
if key == "enabled":
if raw is None or raw == "":
new = "" # 清键=回落默认(启用)
else:
new = "1" if raw in (True, 1, "1", "true", "True") else "0"
else:
new = "" if raw is None else str(raw).strip()
old = cur.get(key)
if old == new:
continue
masked = key == "api_key"
old_shown = mask_secret(old) if masked and old else old
new_shown = mask_secret(new) if masked else new
if new == "":
conn.execute("DELETE FROM llm_config WHERE key=?", (key,))
else:
conn.execute(
"INSERT INTO llm_config(key, value) VALUES(?, ?) "
"ON CONFLICT(key) DO UPDATE SET value=excluded.value",
(key, new))
conn.execute(
"INSERT INTO llm_config_audit(key, old_value, new_value, "
"changed_by) VALUES(?, ?, ?, ?)",
(key, old_shown, new_shown, changed_by))
rows.append({"key": key, "old_value": old_shown,
"new_value": new_shown})
return rows
def get_llm_audit(path: str, limit: int = 50) -> list[dict[str, Any]]:
"""配置页审计列表(倒序最新在前;api_key 行只含打码形)."""
with _connect(path) as conn:
rows = conn.execute(
"SELECT key, old_value, new_value, changed_by, changed_at "
"FROM llm_config_audit ORDER BY id DESC LIMIT ?",
(limit,)).fetchall()
return [dict(r) for r in rows]
+34 -7
View File
@@ -33,6 +33,9 @@ PROVIDERS: dict[str, ProviderSpec] = {
"openai_compat": ProviderSpec("openai_compat", "", "", "SANGUO_LLM_API_KEY"),
}
_TIMEOUT_ENV = "SANGUO_LLM_TIMEOUT"
_TIMEOUT_MIN, _TIMEOUT_MAX = 1.0, 600.0
@dataclass(frozen=True)
class LLMConfig:
@@ -43,17 +46,29 @@ class LLMConfig:
timeout: float = 60.0
def resolve_config(env: Mapping[str, str] = os.environ) -> LLMConfig:
"""优先级:env 覆盖 > provider 表默认;缺关键项 raise LLMConfigError(列 env 名)."""
provider_name = (env.get("SANGUO_LLM_PROVIDER") or "").strip() or "glm_cn"
def resolve_config(env: Mapping[str, str] = os.environ,
overrides: Mapping[str, str] | None = None) -> LLMConfig:
"""优先级: overrides(设置页 DB) > env > provider 表默认(P5.1 三层,
spec §13;overrides=None 保持两层,既有调用方/funnel 零变化).
缺关键项 raise LLMConfigError(列 env 名,env 仍是兜底通道)."""
ov = {k: str(v).strip() for k, v in (overrides or {}).items()
if str(v).strip()}
def pick(key: str, env_name: str, default: str = "") -> str:
return ov.get(key) or (env.get(env_name) or "").strip() or default
if ov.get("enabled") in ("0", "false", "False"):
raise LLMConfigError("LLM 已在设置页停用(llm_config.enabled=0)")
provider_name = pick("provider", "SANGUO_LLM_PROVIDER", "glm_cn")
spec = PROVIDERS.get(provider_name)
if spec is None:
raise LLMConfigError(
f"未知 SANGUO_LLM_PROVIDER={provider_name!r},可用: {sorted(PROVIDERS)}")
base_url = (env.get("SANGUO_LLM_BASE_URL") or "").strip() or spec.base_url
model = (env.get("SANGUO_LLM_MODEL") or "").strip() or spec.default_model
api_key = (env.get(spec.api_key_env) or "").strip()
base_url = pick("base_url", "SANGUO_LLM_BASE_URL", spec.base_url)
model = pick("model", "SANGUO_LLM_MODEL", spec.default_model)
api_key = pick("api_key", spec.api_key_env)
missing = [n for n, v in (
("SANGUO_LLM_BASE_URL", base_url),
@@ -62,7 +77,19 @@ def resolve_config(env: Mapping[str, str] = os.environ) -> LLMConfig:
if missing:
raise LLMConfigError(
f"LLM 未配置: {'/'.join(missing)} 未设置(provider={provider_name})")
return LLMConfig(provider_name, base_url, model, api_key)
timeout_raw = ov.get("timeout") or (env.get(_TIMEOUT_ENV) or "").strip()
timeout = 60.0
if timeout_raw:
try:
timeout = float(timeout_raw)
except ValueError as e:
raise LLMConfigError(
f"LLM timeout 非数字: {timeout_raw!r}") from e
if not _TIMEOUT_MIN <= timeout <= _TIMEOUT_MAX:
raise LLMConfigError(
f"LLM timeout 越界 [{_TIMEOUT_MIN:g},{_TIMEOUT_MAX:g}]: {timeout}")
return LLMConfig(provider_name, base_url, model, api_key, timeout)
def chat_url(base_url: str) -> str:
+111
View File
@@ -0,0 +1,111 @@
# tests/api/test_llm_config_store.py
"""config_store(P5.1):双表 CRUD+审计+打码纪律;tmp db,零真 key."""
import os
import pytest
from sanguo_api.llm import config_store as cs
REAL_KEY = "sk-test-1234567890abcd"
@pytest.fixture()
def db(tmp_path):
return str(tmp_path / "pipeline.db")
class TestMaskSecret:
def test_long_key_masked_head_tail(self):
assert cs.mask_secret(REAL_KEY) == "sk-t****abcd"
def test_short_key_fully_masked(self):
assert cs.mask_secret("short") == "****"
class TestOverridesRoundtrip:
def test_set_then_resolve_db_wins_over_env(self, db, monkeypatch):
monkeypatch.setenv("SANGUO_LLM_API_KEY", "sk-env-old")
monkeypatch.setenv("SANGUO_LLM_MODEL", "env-model")
cs.set_llm_config(db, {"api_key": REAL_KEY, "model": "db-model"},
"admin")
cfg = cs.resolve_effective_config(db)
assert cfg.api_key == REAL_KEY and cfg.model == "db-model"
def test_masked_config_view_never_leaks_key(self, db, monkeypatch):
monkeypatch.delenv("SANGUO_LLM_API_KEY", raising=False)
cs.set_llm_config(db, {"api_key": REAL_KEY}, "admin")
view = cs.get_masked_config(db)
assert view["config"]["api_key"] == "sk-t****abcd"
assert REAL_KEY not in str(view)
assert view["sources"]["api_key"] == "db"
assert view["resolved"] is True and view["missing"] == []
def test_unresolved_reports_missing(self, db, monkeypatch):
for k in ("SANGUO_LLM_API_KEY", "SANGUO_LLM_BASE_URL",
"SANGUO_LLM_MODEL"):
monkeypatch.delenv(k, raising=False)
view = cs.get_masked_config(db)
assert view["resolved"] is False and view["missing"]
def test_sources_default_when_empty(self, db, monkeypatch):
monkeypatch.setenv("SANGUO_LLM_API_KEY", "sk-env")
view = cs.get_masked_config(db)
assert view["sources"]["provider"] == "default"
assert view["sources"]["api_key"] == "env"
class TestSetRules:
def test_masked_apikey_not_written_back(self, db):
cs.set_llm_config(db, {"api_key": REAL_KEY}, "admin")
cs.set_llm_config(db, {"api_key": "sk-t****abcd"}, "admin")
assert cs.get_llm_overrides(db)["api_key"] == REAL_KEY
def test_empty_value_clears_override(self, db, monkeypatch):
monkeypatch.setenv("SANGUO_LLM_API_KEY", "sk-env")
cs.set_llm_config(db, {"api_key": REAL_KEY}, "admin")
cs.set_llm_config(db, {"api_key": ""}, "admin")
assert "api_key" not in cs.get_llm_overrides(db)
assert cs.resolve_effective_config(db).api_key == "sk-env"
def test_unknown_key_rejected(self, db):
with pytest.raises(ValueError, match="未知 llm 配置键"):
cs.set_llm_config(db, {"evil": "x"}, "admin")
def test_unknown_provider_rejected(self, db):
with pytest.raises(ValueError, match="未知 provider"):
cs.set_llm_config(db, {"provider": "nope"}, "admin")
def test_timeout_validated(self, db):
with pytest.raises(ValueError, match="timeout"):
cs.set_llm_config(db, {"timeout": "abc"}, "admin")
with pytest.raises(ValueError, match="timeout"):
cs.set_llm_config(db, {"timeout": "900"}, "admin")
def test_enabled_normalized_and_clearable(self, db):
cs.set_llm_config(db, {"enabled": False}, "admin")
assert cs.get_llm_overrides(db)["enabled"] == "0"
cs.set_llm_config(db, {"enabled": ""}, "admin")
assert "enabled" not in cs.get_llm_overrides(db)
def test_unchanged_value_no_audit_row(self, db):
cs.set_llm_config(db, {"model": "m1"}, "admin")
n1 = len(cs.get_llm_audit(db))
cs.set_llm_config(db, {"model": "m1"}, "admin")
assert len(cs.get_llm_audit(db)) == n1
class TestAudit:
def test_audit_masks_apikey_only(self, db):
cs.set_llm_config(db, {"api_key": REAL_KEY, "model": "m1"}, "admin")
rows = cs.get_llm_audit(db)
by = {r["key"]: r for r in rows}
assert by["api_key"]["new_value"] == "sk-t****abcd"
assert by["api_key"]["old_value"] is None
assert by["model"]["new_value"] == "m1"
assert all(REAL_KEY not in str(r) for r in rows)
def test_audit_desc_newest_first(self, db):
cs.set_llm_config(db, {"model": "m1"}, "a")
cs.set_llm_config(db, {"model": "m2"}, "b")
rows = cs.get_llm_audit(db)
assert rows[0]["new_value"] == "m2" and rows[0]["changed_by"] == "b"
+37
View File
@@ -62,3 +62,40 @@ class TestChatUrl:
"https://open.bigmodel.cn/api/paas/v4/chat/completions"
assert chat_url("http://10.0.0.1:8000/v4") == \
"http://10.0.0.1:8000/v4/chat/completions"
class TestResolveConfigOverrides:
"""P5.1 三层:DB overrides > env > 默认;overrides=None 零行为变化."""
def test_override_beats_env_and_default(self):
cfg = resolve_config(_env(), overrides={
"provider": "openai_compat", "base_url": "http://x/v1/",
"model": "m2", "api_key": "sk-db", "timeout": "30"})
assert (cfg.provider, cfg.base_url, cfg.model, cfg.api_key) == \
("openai_compat", "http://x/v1/", "m2", "sk-db")
assert cfg.timeout == 30.0
def test_env_fills_when_override_absent(self):
cfg = resolve_config(_env(SANGUO_LLM_MODEL="env-model"),
overrides={"api_key": "sk-db"})
assert cfg.model == "env-model" and cfg.api_key == "sk-db"
def test_empty_override_falls_through(self):
# 空串 override=无此层(设置页"清空"语义由 DELETE 行实现,不靠空串)
cfg = resolve_config(_env(), overrides={"model": " "})
assert cfg.model == "glm-5.3-flash"
def test_disabled_override_raises(self):
with pytest.raises(LLMConfigError, match="停用"):
resolve_config(_env(), overrides={"enabled": "0"})
def test_timeout_env_fallback_and_bounds(self):
assert resolve_config(
_env(SANGUO_LLM_TIMEOUT="15")).timeout == 15.0
with pytest.raises(LLMConfigError, match="timeout"):
resolve_config(_env(), overrides={"timeout": "900"})
with pytest.raises(LLMConfigError, match="非数字"):
resolve_config(_env(), overrides={"timeout": "abc"})
def test_no_overrides_backward_compatible(self):
assert resolve_config(_env()) == resolve_config(_env(), overrides=None)