252 lines
8.1 KiB
Python
252 lines
8.1 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""所有插件上层的 filter 规则:全局启用 + 分群控制 + 群聊/私聊使用控制。
|
||
|
||
设计:
|
||
- 以「插件模块名(含 __plugin_meta__ 的包根)」为稳定 identifier。
|
||
- 两层状态存 hexi/data/plugin_control.json:
|
||
{
|
||
"<plugin_id>": {
|
||
"global": {"enabled": true, "chat": ["group","private"]},
|
||
"groups": {"<group_id>": {"enabled": true, "chat": ["group"]}}
|
||
}
|
||
}
|
||
- 默认(未受管插件)全部放行,仅受管的做过滤,避免误伤。
|
||
- 通过往 NoneBot matcher 注册表里每个 application 插件 matcher 的
|
||
rule 追加一个 gate checker(Rule &),实现「在所有插件之上」的统一闸门。
|
||
- 启动时 instrument_plugin_gate() 扫一次;热加载/热重载后再扫一次。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import sys
|
||
from pathlib import Path
|
||
from typing import Any, Optional
|
||
|
||
from nonebot import get_loaded_plugins, logger
|
||
from nonebot.internal.matcher import matchers as matchers_registry
|
||
from nonebot.rule import Rule
|
||
|
||
_PLUGIN_ROOT = Path(__file__).resolve().parents[1] # hexi/
|
||
DATA_DIR = _PLUGIN_ROOT / "data"
|
||
STORE_PATH = DATA_DIR / "plugin_control.json"
|
||
|
||
DEFAULT_CHAT = ["group", "private"]
|
||
|
||
# matcher 被注入后打标,避免重复叠加 gate checker
|
||
_GATE_ATTR = "_hexi_plugin_gate_applied"
|
||
|
||
|
||
def _default_level() -> dict[str, Any]:
|
||
return {"enabled": True, "chat": list(DEFAULT_CHAT)}
|
||
|
||
|
||
def _default_entry() -> dict[str, Any]:
|
||
return {"global": _default_level(), "groups": {}}
|
||
|
||
|
||
_controls: dict[str, dict[str, Any]] = {}
|
||
|
||
|
||
def load() -> None:
|
||
"""从磁盘加载控制面状态(幂等)。"""
|
||
global _controls
|
||
if STORE_PATH.exists():
|
||
try:
|
||
data = json.loads(STORE_PATH.read_text(encoding="utf-8"))
|
||
_controls = data if isinstance(data, dict) else {}
|
||
except Exception as e: # noqa: BLE001
|
||
logger.warning(f"加载插件控制配置失败: {type(e).__name__}: {e}")
|
||
_controls = {}
|
||
else:
|
||
_controls = {}
|
||
|
||
|
||
def save() -> None:
|
||
"""持久化控制面状态。"""
|
||
DATA_DIR.mkdir(parents=True, exist_ok=True)
|
||
STORE_PATH.write_text(
|
||
json.dumps(_controls, ensure_ascii=False, indent=2), encoding="utf-8"
|
||
)
|
||
|
||
|
||
def _norm_level(level: Optional[dict[str, Any]]) -> dict[str, Any]:
|
||
lvl = _default_level()
|
||
if level:
|
||
if "enabled" in level:
|
||
lvl["enabled"] = bool(level["enabled"])
|
||
chat = level.get("chat")
|
||
if isinstance(chat, list) and chat:
|
||
# 只保留合法值
|
||
lvl["chat"] = [c for c in chat if c in {"group", "private"}] or DEFAULT_CHAT
|
||
return lvl
|
||
|
||
|
||
def get_plugin_control(plugin_id: str) -> dict[str, Any]:
|
||
"""返回某插件的控制配置(含 global + groups),未受管返回默认全放行。"""
|
||
entry = _controls.get(plugin_id)
|
||
if not entry:
|
||
return {
|
||
"global": _default_level(),
|
||
"groups": {},
|
||
"managed": False,
|
||
}
|
||
return {
|
||
"global": _norm_level(entry.get("global")),
|
||
"groups": {
|
||
str(gid): _norm_level(level) for gid, level in (entry.get("groups") or {}).items()
|
||
},
|
||
"managed": True,
|
||
}
|
||
|
||
|
||
def set_global(
|
||
plugin_id: str, enabled: Optional[bool] = None, chat: Optional[list[str]] = None
|
||
) -> dict[str, Any]:
|
||
"""设置全局开关/聊天类型,返回最新控制配置。"""
|
||
entry = _controls.setdefault(plugin_id, _default_entry())
|
||
level = _norm_level(entry.get("global"))
|
||
if enabled is not None:
|
||
level["enabled"] = bool(enabled)
|
||
if chat is not None:
|
||
level["chat"] = [c for c in chat if c in {"group", "private"}] or DEFAULT_CHAT
|
||
entry["global"] = level
|
||
save()
|
||
return get_plugin_control(plugin_id)
|
||
|
||
|
||
def set_group(
|
||
plugin_id: str,
|
||
group_id: str,
|
||
enabled: Optional[bool] = None,
|
||
chat: Optional[list[str]] = None,
|
||
) -> dict[str, Any]:
|
||
"""设置某群对某插件的开关/聊天类型,返回最新控制配置。"""
|
||
entry = _controls.setdefault(plugin_id, _default_entry())
|
||
groups = entry.setdefault("groups", {})
|
||
gid = str(group_id)
|
||
level = _norm_level(groups.get(gid))
|
||
if enabled is not None:
|
||
level["enabled"] = bool(enabled)
|
||
if chat is not None:
|
||
level["chat"] = [c for c in chat if c in {"group", "private"}] or DEFAULT_CHAT
|
||
groups[gid] = level
|
||
save()
|
||
return get_plugin_control(plugin_id)
|
||
|
||
|
||
def remove_group(plugin_id: str, group_id: str) -> dict[str, Any]:
|
||
"""移除某群覆盖(回到继承全局)。"""
|
||
entry = _controls.get(plugin_id)
|
||
if entry:
|
||
entry.get("groups", {}).pop(str(group_id), None)
|
||
save()
|
||
return get_plugin_control(plugin_id)
|
||
|
||
|
||
def remove_plugin(plugin_id: str) -> None:
|
||
"""清空某插件所有覆盖,恢复默认放行。"""
|
||
_controls.pop(plugin_id, None)
|
||
save()
|
||
|
||
|
||
def list_plugins() -> list[dict[str, Any]]:
|
||
"""枚举所有 application 插件(或注册了配置 schema 的插件)及其控制面状态。"""
|
||
from nonebot.plugin import get_loaded_plugins
|
||
from hexi.web_hub.web_config import has_schema
|
||
|
||
result: list[dict[str, Any]] = []
|
||
seen: set[str] = set()
|
||
for p in get_loaded_plugins():
|
||
mod = p.module_name
|
||
if mod in seen:
|
||
continue
|
||
seen.add(mod)
|
||
meta = p.metadata
|
||
if not meta:
|
||
continue
|
||
if meta.type != "application" and not has_schema(mod):
|
||
continue
|
||
ctl = get_plugin_control(mod)
|
||
result.append(
|
||
{
|
||
"id": mod,
|
||
"name": meta.name,
|
||
"description": meta.description or "",
|
||
"usage": meta.usage or "",
|
||
"type": meta.type,
|
||
"enabled": ctl["global"]["enabled"],
|
||
"chat": ctl["global"]["chat"],
|
||
"groups": {gid: lvl for gid, lvl in ctl["groups"].items()},
|
||
"managed": ctl["managed"],
|
||
}
|
||
)
|
||
result.sort(key=lambda x: x["name"])
|
||
return result
|
||
|
||
|
||
def _resolve_plugin_id(module_name: str) -> Optional[str]:
|
||
"""把 matcher 的模块名回溯到带 __plugin_meta__ / NoneBot 注册表的插件根模块。"""
|
||
parts = module_name.split(".")
|
||
for i in range(len(parts), 0, -1):
|
||
cand = ".".join(parts[:i])
|
||
for p in get_loaded_plugins():
|
||
if p.module_name == cand:
|
||
return cand
|
||
mod = sys.modules.get(cand)
|
||
if mod is not None and getattr(mod, "__plugin_meta__", None) is not None:
|
||
return cand
|
||
return None
|
||
|
||
|
||
def _gate_checker(plugin_id: str):
|
||
"""为单个插件生成一个异步 rule checker(读控制面状态)。"""
|
||
|
||
async def _check(bot, event, state) -> bool: # noqa: ANN001
|
||
return _allowed(plugin_id, event)
|
||
|
||
return _check
|
||
|
||
|
||
def _allowed(plugin_id: str, event) -> bool:
|
||
"""决定某事件是否允许进入该插件。"""
|
||
cfg = get_plugin_control(plugin_id)
|
||
is_group = getattr(event, "message_type", "") == "group" or (
|
||
getattr(event, "group_id", None) is not None
|
||
)
|
||
if is_group:
|
||
gid = str(getattr(event, "group_id", "") or "")
|
||
gcfg = cfg["groups"].get(gid)
|
||
level = gcfg if gcfg is not None else cfg["global"]
|
||
else:
|
||
level = cfg["global"]
|
||
if not level["enabled"]:
|
||
return False
|
||
chat_type = "group" if is_group else "private"
|
||
return chat_type in level["chat"]
|
||
|
||
|
||
def instrument_plugin_gate() -> int:
|
||
"""给所有 application 插件 matcher 追加统一 gate 规则,返回注入数量。"""
|
||
count = 0
|
||
for group in matchers_registry.values():
|
||
for matcher in group:
|
||
if getattr(matcher, _GATE_ATTR, False):
|
||
continue
|
||
mod = getattr(matcher, "module_name", None)
|
||
if not mod:
|
||
continue
|
||
pid = _resolve_plugin_id(mod)
|
||
if not pid:
|
||
continue
|
||
checker = Rule(_gate_checker(pid))
|
||
current = getattr(matcher, "rule", None)
|
||
matcher.rule = checker if current is None else current & checker
|
||
setattr(matcher, _GATE_ATTR, pid)
|
||
count += 1
|
||
return count
|
||
|
||
|
||
# 导入时加载状态(幂等,可重复调用 load() 刷新)
|
||
load()
|