410 lines
15 KiB
Python
410 lines
15 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""插件配置统一标准:schema 注册 + 存储 + 保存热刷新。
|
||
|
||
标准 schema 格式(每个插件注册一份):
|
||
{
|
||
"fields": [
|
||
{
|
||
"key": "web_password", # 配置字段名
|
||
"label": "Web 密码", # 表单显示名
|
||
"type": "string|text|password|int|float|bool|enum", # 类型
|
||
"default": ..., # 默认值
|
||
"description": "...", # 提示
|
||
"env": "WEB_PASSWORD", # 环境变量名(缺省=key.upper())
|
||
"secret": true, # 是否敏感(前端默认掩码显示,回传明文)
|
||
"options": [...], # enum 可选项
|
||
"required": false,
|
||
}
|
||
]
|
||
}
|
||
|
||
存储:hexi/config/plugin_config.json(Web 读写源)。
|
||
热刷新:
|
||
- 若注册时提供了 apply(values) 回调,保存后调用它(插件自定义热应用,推荐)。
|
||
- 否则写环境变量(os.environ + .env,保证 get_plugin_config 能读到),再按需 hot_reload 插件。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import os
|
||
import re
|
||
import threading
|
||
from pathlib import Path
|
||
from typing import Any, Callable, Optional
|
||
|
||
from nonebot import logger
|
||
|
||
_ROOT = Path(__file__).resolve().parents[1] # hexi/
|
||
DATA_DIR = _ROOT / "config"
|
||
STORE_PATH = DATA_DIR / "plugin_config.json"
|
||
ENV_PATH = _ROOT.parent / ".env" # 仓库根 .env
|
||
|
||
_TYPES = {"string", "text", "password", "int", "float", "bool", "enum", "json", "object", "path", "object_set", "array"}
|
||
|
||
# plugin_id -> schema
|
||
_schemas: dict[str, dict[str, Any]] = {}
|
||
# plugin_id -> apply(values) 回调
|
||
_appliers: dict[str, Callable[[dict[str, Any]], Any]] = {}
|
||
# plugin_id -> getter() 回调(返回当前生效值,Web 表单回填用)
|
||
_getters: dict[str, Callable[[], dict[str, Any]]] = {}
|
||
# plugin_id -> {key: value}
|
||
_values: dict[str, dict[str, Any]] = {}
|
||
_revisions: dict[str, int] = {}
|
||
_config_lock = threading.RLock()
|
||
MASKED_SECRET = "****"
|
||
|
||
|
||
def _load() -> None:
|
||
global _values
|
||
if STORE_PATH.exists():
|
||
try:
|
||
data = json.loads(STORE_PATH.read_text(encoding="utf-8"))
|
||
if isinstance(data, dict) and isinstance(data.get("plugins"), dict):
|
||
_values = {}
|
||
for pid, entry in data["plugins"].items():
|
||
if isinstance(entry, dict) and isinstance(entry.get("values", {}), dict):
|
||
_values[pid] = entry["values"]
|
||
_revisions[pid] = int(entry.get("revision", 0))
|
||
else:
|
||
_values = data if isinstance(data, dict) else {}
|
||
except Exception as e: # noqa: BLE001
|
||
logger.warning(f"加载插件配置存储失败: {type(e).__name__}: {e}")
|
||
_values = {}
|
||
else:
|
||
_values = {}
|
||
|
||
|
||
def _save() -> None:
|
||
DATA_DIR.mkdir(parents=True, exist_ok=True)
|
||
payload = {
|
||
"version": 1,
|
||
"plugins": {
|
||
pid: {"revision": _revisions.get(pid, 0), "values": values}
|
||
for pid, values in _values.items()
|
||
},
|
||
}
|
||
tmp = STORE_PATH.with_suffix(".tmp")
|
||
tmp.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")
|
||
os.replace(tmp, STORE_PATH)
|
||
|
||
|
||
def _field_env(field: dict[str, Any]) -> str:
|
||
return str(field.get("env") or str(field["key"]).upper())
|
||
|
||
|
||
def register_plugin_config(
|
||
plugin_id: str,
|
||
schema: dict[str, Any],
|
||
apply: Optional[Callable[[dict[str, Any]], Any]] = None,
|
||
getter: Optional[Callable[[], dict[str, Any]]] = None,
|
||
) -> None:
|
||
"""注册某插件的标准配置 schema;apply 可选,用于保存后热刷新;
|
||
getter 可选,返回插件当前生效值,供 Web 表单回填(否则回填 schema 默认值)。
|
||
|
||
同一 plugin_id 可多次调用:字段按 key 合并到一个 schema,
|
||
getter/apply 会被组合成「合并取值 / 依次应用」,避免后注册覆盖前注册。
|
||
"""
|
||
fields = schema.get("fields", [])
|
||
for f in fields:
|
||
typ = f.get("type", "string")
|
||
if typ not in _TYPES:
|
||
raise ValueError(f"插件 {plugin_id} 字段 {f.get('key')} 类型 {typ} 不受支持")
|
||
if "key" not in f:
|
||
raise ValueError(f"插件 {plugin_id} 存在缺 key 的字段")
|
||
f.setdefault("label", f["key"])
|
||
f.setdefault("default", None)
|
||
f.setdefault("secret", False)
|
||
|
||
existing = _schemas.get(plugin_id)
|
||
if existing is not None:
|
||
existing_keys = {f["key"] for f in existing.get("fields", [])}
|
||
for f in fields:
|
||
if f["key"] not in existing_keys:
|
||
existing["fields"].append(f)
|
||
existing_keys.add(f["key"])
|
||
schema = existing
|
||
_schemas[plugin_id] = schema
|
||
|
||
if apply is not None:
|
||
prev = _appliers.get(plugin_id)
|
||
if prev is None:
|
||
_appliers[plugin_id] = apply
|
||
else:
|
||
def _combo_apply(values, _prev=prev, _new=apply):
|
||
try:
|
||
_prev(values)
|
||
except Exception: # noqa: BLE001
|
||
logger.warning(f"插件配置 apply 前段失败({plugin_id}): …")
|
||
_new(values)
|
||
_appliers[plugin_id] = _combo_apply
|
||
|
||
if getter is not None:
|
||
prevg = _getters.get(plugin_id)
|
||
if prevg is None:
|
||
_getters[plugin_id] = getter
|
||
else:
|
||
def _combo_getter(_prev=prevg, _new=getter):
|
||
out: dict[str, Any] = {}
|
||
try:
|
||
out.update(_prev())
|
||
except Exception: # noqa: BLE001
|
||
pass
|
||
try:
|
||
out.update(_new())
|
||
except Exception: # noqa: BLE001
|
||
pass
|
||
return out
|
||
_getters[plugin_id] = _combo_getter
|
||
|
||
|
||
def has_schema(plugin_id: str) -> bool:
|
||
return plugin_id in _schemas
|
||
|
||
|
||
def get_effective_value(plugin_id: str, key: str, default: Any = None) -> Any:
|
||
"""统一读取「插件配置生效值」(供插件运行期调用,来源无关)。
|
||
|
||
优先级:
|
||
1. 用户已保存的值(plugin_config.json)
|
||
2. 插件 getter 返回的当前生效值
|
||
3. default
|
||
|
||
适合模块常量/配置文件/DB 内的配置项:插件在用到某配置时调用
|
||
get_effective_value(plugin_id, key, default) 读取,即可立刻反映 Web 修改。
|
||
"""
|
||
stored = _values.get(plugin_id, {})
|
||
if key in stored and stored[key] is not None:
|
||
return stored[key]
|
||
getter = _getters.get(plugin_id)
|
||
if getter:
|
||
try:
|
||
live = getter()
|
||
except Exception: # noqa: BLE001
|
||
live = None
|
||
if isinstance(live, dict) and key in live and live[key] is not None:
|
||
return live[key]
|
||
return default
|
||
|
||
|
||
def get_schema(plugin_id: str) -> Optional[dict[str, Any]]:
|
||
return _schemas.get(plugin_id)
|
||
|
||
|
||
def _coerce(field: dict[str, Any], raw: Any) -> Any:
|
||
typ = field.get("type", "string")
|
||
if typ == "bool":
|
||
if isinstance(raw, bool):
|
||
return raw
|
||
normalized = str(raw).strip().lower()
|
||
if normalized not in {"0", "1", "true", "false", "yes", "no", "on", "off"}:
|
||
raise ValueError("布尔值必须是 true/false")
|
||
return normalized in {"1", "true", "yes", "on"}
|
||
if typ == "int":
|
||
try:
|
||
return int(raw)
|
||
except (TypeError, ValueError) as e:
|
||
raise ValueError(f"字段 {field['key']} 必须是整数") from e
|
||
if typ == "float":
|
||
try:
|
||
return float(raw)
|
||
except (TypeError, ValueError) as e:
|
||
raise ValueError(f"字段 {field['key']} 必须是数字") from e
|
||
if typ == "array":
|
||
item_type = str(field.get("item_type") or "str")
|
||
if isinstance(raw, str):
|
||
# 防御:兼容旧前端提交的多行/逗号分隔文本
|
||
raw = [p.strip() for p in re.split(r"[,\n]+", raw) if p.strip()]
|
||
if not isinstance(raw, list):
|
||
raise ValueError(f"字段 {field['key']} 必须是数组")
|
||
out = []
|
||
for item in raw:
|
||
if item_type in {"int", "float"}:
|
||
try:
|
||
out.append(int(item) if item_type == "int" else float(item))
|
||
except (TypeError, ValueError) as e:
|
||
raise ValueError(f"字段 {field['key']} 含非法数字: {item}") from e
|
||
else:
|
||
out.append(str(item))
|
||
return out
|
||
if typ in {"string", "text", "password"}:
|
||
return "" if raw is None else str(raw)
|
||
if typ == "enum":
|
||
return "" if raw is None else str(raw)
|
||
return raw
|
||
|
||
|
||
def get_config(plugin_id: str) -> Optional[dict[str, Any]]:
|
||
"""返回 {schema, values};未注册 schema 返回 None。"""
|
||
schema = _schemas.get(plugin_id)
|
||
if schema is None:
|
||
return None
|
||
stored = _values.get(plugin_id, {})
|
||
live: dict[str, Any] = {}
|
||
getter = _getters.get(plugin_id)
|
||
if getter:
|
||
try:
|
||
live = getter()
|
||
except Exception as e: # noqa: BLE001
|
||
logger.warning(f"插件配置 getter 回填失败({plugin_id}): {type(e).__name__}: {e}")
|
||
values: dict[str, Any] = {}
|
||
for f in schema.get("fields", []):
|
||
key = f["key"]
|
||
if key in stored:
|
||
values[key] = stored[key]
|
||
elif key in live:
|
||
values[key] = live[key]
|
||
elif "default" in f:
|
||
values[key] = f["default"]
|
||
else:
|
||
values[key] = None
|
||
# 密钥类字段回传明文:该管理台仅本地局域网/本机使用,无外泄面;
|
||
# 明文由前端默认以掩码(小圆点)形式展示,用户点「眼睛」才可查看。
|
||
# array 字段兼容旧存储:text 时期保存的是换行/逗号分隔字符串,归一化为数组
|
||
for f in schema.get("fields", []):
|
||
key = f["key"]
|
||
if f.get("type") == "array" and isinstance(values.get(key), str):
|
||
try:
|
||
values[key] = _coerce(f, values[key])
|
||
except ValueError:
|
||
values[key] = []
|
||
return {"schema": schema, "values": values, "revision": _revisions.get(plugin_id, 0)}
|
||
|
||
|
||
def _validate_field(field: dict[str, Any], value: Any) -> None:
|
||
key = field["key"]
|
||
if value in (None, ""):
|
||
if field.get("required"):
|
||
raise ValueError(f"字段 {key} 不能为空")
|
||
return
|
||
typ = field.get("type", "string")
|
||
if typ == "enum":
|
||
options = {str(o.get("value") if isinstance(o, dict) else o) for o in field.get("options", [])}
|
||
if str(value) not in options:
|
||
raise ValueError(f"字段 {key} 不是有效选项")
|
||
if typ in {"int", "float"}:
|
||
number = float(value)
|
||
if field.get("min") is not None and number < field["min"]:
|
||
raise ValueError(f"字段 {key} 小于最小值")
|
||
if field.get("max") is not None and number > field["max"]:
|
||
raise ValueError(f"字段 {key} 大于最大值")
|
||
|
||
|
||
def _update_env(plugin_id: str, values: dict[str, Any]) -> None:
|
||
"""写 os.environ(即时生效)+ 写 .env(重启保留)。"""
|
||
fields = _schemas.get(plugin_id, {}).get("fields", [])
|
||
for f in fields:
|
||
key = f["key"]
|
||
if key not in values:
|
||
continue
|
||
# text/json/object/path/array 这类非标量不适合写进 .env(重启解析会出错),
|
||
# 这类字段仅在运行期热更新,不落盘 .env(见 config_standard 说明)。
|
||
if f.get("type") in {"text", "json", "object", "path", "array"}:
|
||
continue
|
||
env_name = _field_env(f)
|
||
raw = values[key]
|
||
text = "" if raw is None else str(raw)
|
||
os.environ[env_name] = text
|
||
_append_env(env_name, text)
|
||
|
||
|
||
def _append_env(env_name: str, value: str) -> None:
|
||
"""把 KEY=VALUE 写进 .env(已存在则在原行更新,否则追加)。"""
|
||
try:
|
||
lines: list[str] = []
|
||
if ENV_PATH.exists():
|
||
lines = ENV_PATH.read_text(encoding="utf-8").splitlines()
|
||
replaced = False
|
||
out: list[str] = []
|
||
for line in lines:
|
||
stripped = line.strip()
|
||
if stripped.startswith(env_name + "="):
|
||
if not replaced:
|
||
out.append(f"{env_name}={value}")
|
||
replaced = True
|
||
# 重复行跳过
|
||
else:
|
||
out.append(line)
|
||
if not replaced:
|
||
out.append(f"{env_name}={value}")
|
||
ENV_PATH.write_text("\n".join(out).rstrip("\n") + "\n", encoding="utf-8")
|
||
except Exception as e: # noqa: BLE001
|
||
logger.warning(f"写入 .env({env_name}) 失败: {type(e).__name__}: {e}")
|
||
|
||
|
||
def _hot_reload(plugin_id: str) -> None:
|
||
try:
|
||
from hexi.core.plugin_manager import hot_reload
|
||
|
||
ok = hot_reload(plugin_id)
|
||
if ok:
|
||
logger.info(f"插件配置热刷新: hot_reload({plugin_id}) 成功")
|
||
else:
|
||
logger.warning(
|
||
f"插件配置热刷新跳过 {plugin_id}: 非 application(library/未声明)插件不可热重载,"
|
||
"请检查是否已提供 apply 回调或直接重启"
|
||
)
|
||
except Exception as e: # noqa: BLE001
|
||
logger.warning(f"插件配置热刷新失败({plugin_id}): {type(e).__name__}: {e}")
|
||
|
||
|
||
def save_config(
|
||
plugin_id: str,
|
||
values: dict[str, Any],
|
||
reload: bool = True,
|
||
expected_revision: int | None = None,
|
||
) -> dict[str, Any]:
|
||
"""保存配置并触发热刷新。返回最新 {schema, values}。"""
|
||
schema = _schemas.get(plugin_id)
|
||
if schema is None:
|
||
raise ValueError(f"插件 {plugin_id} 未注册配置 schema")
|
||
|
||
fields = schema.get("fields", [])
|
||
with _config_lock:
|
||
current_revision = _revisions.get(plugin_id, 0)
|
||
if expected_revision is not None and expected_revision != current_revision:
|
||
raise RuntimeError(f"配置已被其他请求修改,当前版本为 {current_revision}")
|
||
merged = dict(_values.get(plugin_id, {}))
|
||
persist_merged = dict(_values.get(plugin_id, {}))
|
||
for f in fields:
|
||
key = f["key"]
|
||
if key not in values:
|
||
continue
|
||
raw = values[key]
|
||
# 兼容旧版本前端回传的掩码占位:不覆盖已保存的明文
|
||
if f.get("secret") and raw == MASKED_SECRET:
|
||
continue
|
||
if raw in (None, ""):
|
||
val: Any = None if raw is None else ""
|
||
else:
|
||
val = _coerce(f, raw)
|
||
_validate_field(f, val)
|
||
merged[key] = val
|
||
# nosave fields apply to their authority source but skip the shared value store.
|
||
if not f.get("nosave"):
|
||
persist_merged[key] = val
|
||
_values[plugin_id] = persist_merged
|
||
_revisions[plugin_id] = current_revision + 1
|
||
_save()
|
||
|
||
# 统一落盘到 .env,保证重启后依然生效(与 NoneBot get_plugin_config 约定一致)
|
||
_update_env(plugin_id, merged)
|
||
|
||
applier = _appliers.get(plugin_id)
|
||
if applier is not None:
|
||
# apply 负责把值热应用到插件运行态对象(env 变更不会自动刷新 NoneBot driver.config)
|
||
try:
|
||
applier(dict(merged))
|
||
except Exception as e: # noqa: BLE001
|
||
logger.warning(f"插件配置 apply 回调失败({plugin_id}): {type(e).__name__}: {e}")
|
||
else:
|
||
if reload:
|
||
_hot_reload(plugin_id)
|
||
|
||
result = get_config(plugin_id)
|
||
return result
|
||
|
||
|
||
# 导入时加载存储
|
||
_load()
|