Files
HeXi/hexi/core/plugin_control.py
T
sansenhoshi 131b92b319 结构调整
视频解析多图/多媒体结构 消息体适配
2026-09-08 14:25:32 +08:00

252 lines
8.1 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.
# -*- 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()