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:
@@ -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]
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user