Files
zhenxun_bot/zhenxun/builtin_plugins/hooks/auth_checker.py
T
837330e30a 🎉🍾♻️refactor(core): 重构核心鉴权与缓存机制,引入异步发送队列以提升性能 🚀 (#2088)
* 添加测试插件

* ✨ feat(auth): 添加缓存机制以优化用户和插件数据查询性能

* 🗑️ chore(help): 删除笨蛋检测插件代码

* ✨ feat(bot): 添加对扩展插件的加载支持

* ✨ feat(auth): 优化权限检查和缓存机制,增加用户和插件数据的并行查询

* ✨ feat(cache): 引入运行时缓存机制,优化用户和插件的ban记录管理

* ✨ feat(auth): 更新is_ban函数文档,添加参数和返回值说明

* ```
feat(auth): 使用内存缓存优化权限验证性能

- 移除数据库查询超时控制,改用 LevelUserMemoryCache、BotMemoryCache
  和 PluginLimitMemoryCache 进行缓存查询
- 优化 auth_admin、auth_bot、auth_limit 权限验证逻辑,提升响应速度
- 添加 background 参数支持异步发送权限不足提示消息
- 移除 asyncio 依赖,简化代码结构

fix(auth): 修复限制通知频率控制问题

- 实现限制通知冷却机制,避免重复发送相同限制消息
- 添加 AUTH_LIMIT_NOTICE_CD 配置项,默认值为 2 秒
- 使用 FreqLimiter 控制限制通知发送频率

refactor(models): 增强模型数据变更时的缓存同步

- 在 BotConsole、GroupConsole、LevelUser、PluginLimit 模型的
  create、update_or_create、save、delete 方法中自动更新对应缓存
- 确保数据库和内存缓存数据一致性

docs(ban_console): 修正文档注释并优化日志信息

- 修正 BanConsole 类中方法的文档字符串,使用标准参数和返回值格式
- 优化调试日志信息,使描述更加清晰准确
```

* ✨ feat(auth): 更新Limit类以支持PluginLimitSnapshot,优化限制信息处理

* ✨ feat(bot_manage): 优化Bot控制台初始化逻辑,处理IntegrityError异常

* ✨ feat(mmm1): 新增消息推送功能,支持私聊和群聊事件处理

* ✨ feat(group_member_update): 优化群组成员更新逻辑,增加活动跟踪和消息记录功能

* ✨ feat(mmm1): 删除冗余的消息推送功能代码

* ✨ feat(auth): 优化权限检查逻辑,增加模块阻止功能和缓存处理

* ✨ feat(auth): 优化权限检查逻辑,增加快速ban检测和前置检查功能

* ✨ feat(chat_history): 增强消息处理规则,添加时间间隔限制以防止重复消息
✨ feat(data_source): 引入异步获取群成员信息的功能,优化用户信息更新逻辑

* ✨ feat(send_queue): 添加异步发送队列以优化API调用和速率限制

* ✨ feat(auth): 添加缓存就绪检查以优化权限处理逻辑

* ✨ feat(plugins): 移除不必要的插件加载以简化插件管理

* ✨ feat(group_console): 优化群组获取逻辑,添加缓存检查以提升性能
✨ 只接收缓存完成之后时间的消息

* ✨ feat(auth): 添加异步任务管理和超载检测,优化权限处理逻辑
✨ feat(chat_history): 修改规则函数为异步,提升消息处理效率
✨ feat(group_handle): 增加安全获取群组信息的异步方法,添加超时处理
✨ feat(record_request): 引入安全获取群组信息的异步方法,优化群邀请处理
✨ feat(ban_memory_cache): 增强禁言内存缓存,添加负缓存机制
✨ feat(message_load): 新增消息负载检测功能,优化任务调度
✨ feat(scheduler): 在调度器中集成消息压力检测,优化任务执行
✨ feat(send_queue): 引入异步任务管理,优化发送队列处理

* ✨ feat(auth): 添加对 LevelUserSnapshot 和 BotSnapshot 的支持,优化权限检查逻辑
✨ 格式化

* ✨ feat(db_context): 增强 get_or_create 方法,处理并发创建冲突并回退查询已存在记录

* ✨ feat(bot_manage): 增强 init_bot_console 方法,处理并发创建冲突并回退查询已存在的 bot 数据

* 🚨 auto fix by pre-commit hooks

* ✨ feat(runtime_cache): 优化消息处理逻辑,支持 bytes 和 bytearray 类型的联合判断

* ✨ style(runtime_cache): 格式化代码,优化多行表达式的可读性

* 🚨 auto fix by pre-commit hooks

* ✨ refactor: 优化类型注解,改进群组成员更新逻辑和平台处理
✨ 格式化

* 🚨 auto fix by pre-commit hooks

* ✨ refactor: 优化代码格式,增强可读性并修复类型注解

* 🚨 auto fix by pre-commit hooks

* ✨ refactor: 更新文档注释,增强is_ban函数的可读性

* ✨ refactor: 优化代码结构,移除冗余函数,增强可读性并改进任务调度逻辑

* ✨ refactor: 调整定时任务时间,优化渲染服务的初始化逻辑,增强代码可读性

---------

Co-authored-by: ATTomatoo <1126160939@qq.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
2026-01-28 10:15:24 +08:00

1077 lines
36 KiB
Python

import asyncio
import contextlib
import time
from typing import cast
from nonebot import get_loaded_plugins
from nonebot.adapters import Bot, Event
from nonebot.exception import IgnoredException
from nonebot.matcher import Matcher
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config
from zhenxun.configs.utils import PluginExtraData
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.user_console import UserConsole
from zhenxun.services.cache.cache_containers import CacheDict
from zhenxun.services.cache.runtime_cache import (
BotMemoryCache,
BotSnapshot,
GroupMemoryCache,
GroupSnapshot,
LevelUserMemoryCache,
LevelUserSnapshot,
PluginInfoMemoryCache,
)
from zhenxun.services.data_access import DataAccess
from zhenxun.services.log import logger
from zhenxun.services.message_load import is_overloaded
from zhenxun.utils.enum import BlockType, GoldHandle, PluginType
from zhenxun.utils.exception import InsufficientGold
from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.utils import get_entity_ids
from .auth.auth_admin import auth_admin
from .auth.auth_ban import auth_ban, is_ban
from .auth.auth_bot import auth_bot
from .auth.auth_cost import auth_cost
from .auth.auth_group import auth_group
from .auth.auth_limit import LimitManager, auth_limit
from .auth.auth_plugin import auth_plugin
from .auth.bot_filter import bot_filter
from .auth.config import LOGGER_COMMAND, WARNING_THRESHOLD
from .auth.exception import (
IsSuperuserException,
PermissionExemption,
SkipPluginException,
)
from .auth.utils import base_config
Config.add_plugin_config(
"hook",
"AUTH_HOOKS_CONCURRENCY_LIMIT",
6,
help="auth hooks concurrency limit",
)
Config.add_plugin_config(
"hook",
"AUTH_DB_CONCURRENCY_LIMIT",
6,
help="auth db concurrency limit",
)
Config.add_plugin_config(
"hook",
"AUTH_PLUGIN_CACHE_TTL",
30,
help="plugin info cache ttl seconds",
)
Config.add_plugin_config(
"hook",
"AUTH_USER_CACHE_TTL",
5,
help="user cache ttl seconds",
)
Config.add_plugin_config(
"hook",
"AUTH_EVENT_CACHE_TTL",
2,
help="event auth cache ttl seconds",
)
def _coerce_positive_int(value, default):
try:
value_int = int(value)
except (TypeError, ValueError):
return default
return value_int if value_int > 0 else default
def _coerce_cache_ttl(value, default):
try:
value_int = int(value)
except (TypeError, ValueError):
return default
return value_int if value_int >= 0 else default
# 超时设置(秒)
TIMEOUT_SECONDS = 5.0
# 熔断计数器
CIRCUIT_BREAKERS = {
"auth_ban": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_bot": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_group": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_admin": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_plugin": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_limit": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
}
# 熔断重置时间(秒)
CIRCUIT_RESET_TIME = 300 # 5分钟
# 并发控制:限制同时进入 hooks 并行检查的协程数
# 默认为 6,可通过环境变量 AUTH_HOOKS_CONCURRENCY_LIMIT 调整
HOOKS_CONCURRENCY_LIMIT = _coerce_positive_int(
base_config.get("AUTH_HOOKS_CONCURRENCY_LIMIT", 6), 6
)
DB_CONCURRENCY_LIMIT = _coerce_positive_int(
base_config.get("AUTH_DB_CONCURRENCY_LIMIT", HOOKS_CONCURRENCY_LIMIT),
HOOKS_CONCURRENCY_LIMIT,
)
PLUGIN_CACHE_TTL = _coerce_cache_ttl(base_config.get("AUTH_PLUGIN_CACHE_TTL", 30), 30)
USER_CACHE_TTL = _coerce_cache_ttl(base_config.get("AUTH_USER_CACHE_TTL", 5), 5)
PLUGIN_CACHE = (
CacheDict("AUTH_PLUGIN_CACHE", expire=PLUGIN_CACHE_TTL)
if PLUGIN_CACHE_TTL > 0
else None
)
USER_CACHE = (
CacheDict("AUTH_USER_CACHE", expire=USER_CACHE_TTL) if USER_CACHE_TTL > 0 else None
)
EVENT_CACHE_TTL = _coerce_cache_ttl(base_config.get("AUTH_EVENT_CACHE_TTL", 2), 2)
EVENT_CACHE = (
CacheDict("AUTH_EVENT_CACHE", expire=EVENT_CACHE_TTL)
if EVENT_CACHE_TTL > 0
else None
)
# 路由索引缓存
_ROUTE_INDEX_LOCK = asyncio.Lock()
_ROUTE_INDEX_READY = False
_ROUTE_COMMAND_MAP: dict[str, set[str]] = {}
_ROUTE_PREFIX_MAP: dict[str, set[str]] = {}
_ROUTE_MODULES_WITH_COMMANDS: set[str] = set()
# 全局信号量与计数器
HOOKS_SEMAPHORE = asyncio.Semaphore(HOOKS_CONCURRENCY_LIMIT)
HOOKS_ACTIVE_COUNT = 0
HOOKS_ACTIVE_LOCK = asyncio.Lock()
DB_SEMAPHORE = asyncio.Semaphore(DB_CONCURRENCY_LIMIT)
DB_ACTIVE_COUNT = 0
DB_ACTIVE_LOCK = asyncio.Lock()
def _cache_get(cache: CacheDict | None, key: str):
if not cache:
return None
try:
return cache[key]
except KeyError:
return None
def _cache_set(cache: CacheDict | None, key: str, value):
if cache:
cache[key] = value
def _debug_log(message: str, *args, **kwargs) -> None:
if is_overloaded():
return
logger.debug(message, *args, **kwargs)
def _event_cache_key(event: Event, session: Uninfo, entity) -> str:
msg_id = getattr(event, "message_id", None)
if msg_id is None:
msg_id = getattr(event, "id", None)
if msg_id is None:
msg_id = id(event)
platform = PlatformUtils.get_platform(session)
group_id = entity.group_id or ""
channel_id = entity.channel_id or ""
return (
f"{platform}:{session.self_id}:{entity.user_id}:"
f"{group_id}:{channel_id}:{msg_id}"
)
def _get_event_cache(event: Event, session: Uninfo, entity):
if not EVENT_CACHE:
return None
key = _event_cache_key(event, session, entity)
try:
return EVENT_CACHE[key]
except KeyError:
cache = {}
EVENT_CACHE[key] = cache
return cache
def _normalize_command(command: str) -> str:
return command.strip()
def _extract_commands(extra: PluginExtraData | None) -> set[str]:
if not extra:
return set()
commands = {c.command for c in extra.commands if c.command}
commands.update(extra.aliases or set())
return {cmd.strip() for cmd in commands if cmd and cmd.strip()}
async def _ensure_route_index():
global _ROUTE_INDEX_READY
if _ROUTE_INDEX_READY:
return
async with _ROUTE_INDEX_LOCK:
if _ROUTE_INDEX_READY:
return
_ROUTE_COMMAND_MAP.clear()
_ROUTE_PREFIX_MAP.clear()
_ROUTE_MODULES_WITH_COMMANDS.clear()
for plugin in get_loaded_plugins():
if not plugin.metadata:
continue
extra = plugin.metadata.extra or {}
try:
extra_data = PluginExtraData(**extra)
except Exception:
continue
command_set = _extract_commands(extra_data)
if not command_set:
continue
module = plugin.name
_ROUTE_MODULES_WITH_COMMANDS.add(module)
for command in command_set:
normalized = _normalize_command(command)
if not normalized:
continue
_ROUTE_COMMAND_MAP.setdefault(normalized, set()).add(module)
_ROUTE_PREFIX_MAP.setdefault(normalized[0], set()).add(normalized)
_ROUTE_INDEX_READY = True
def _command_matches(text: str, command: str) -> bool:
if not text or not command:
return False
if text == command:
return True
if text.startswith(command):
if len(text) == len(command):
return True
next_char = text[len(command)]
return next_char.isspace()
return False
def _match_route_modules(text: str) -> set[str]:
text = text.strip()
if not text:
return set()
commands = _ROUTE_PREFIX_MAP.get(text[0])
if not commands:
return set()
matched_modules: set[str] = set()
for command in commands:
if _command_matches(text, command):
modules = _ROUTE_COMMAND_MAP.get(command)
if modules:
matched_modules.update(modules)
return matched_modules
def _get_message_text(message: UniMsg, event_cache: dict | None) -> str:
if event_cache is None:
return message.extract_plain_text()
cached = event_cache.get("plain_text")
if cached is None:
cached = message.extract_plain_text()
event_cache["plain_text"] = cached
return cached
async def _get_route_context(text: str, event_cache: dict | None) -> set[str]:
if not text:
return set()
if event_cache is not None and "route_modules" in event_cache:
return event_cache["route_modules"]
await _ensure_route_index()
matched = _match_route_modules(text)
if event_cache is not None:
event_cache["route_modules"] = matched
return matched
async def _has_limits_cached(module: str, event_cache: dict | None) -> bool:
module_limit_cache: dict[str, bool] = {}
if event_cache is not None:
module_limit_cache = event_cache.setdefault("module_limits", {})
if module in module_limit_cache:
return module_limit_cache[module]
limits = await LimitManager.get_module_limits(module)
has_limits = bool(limits)
module_limit_cache[module] = has_limits
return has_limits
@contextlib.asynccontextmanager
async def _db_section():
global DB_ACTIVE_COUNT
await DB_SEMAPHORE.acquire()
async with DB_ACTIVE_LOCK:
DB_ACTIVE_COUNT += 1
_debug_log(f"current db auth concurrency: {DB_ACTIVE_COUNT}", LOGGER_COMMAND)
try:
yield
finally:
with contextlib.suppress(Exception):
DB_SEMAPHORE.release()
async with DB_ACTIVE_LOCK:
DB_ACTIVE_COUNT = max(DB_ACTIVE_COUNT - 1, 0)
_debug_log(
f"current db auth concurrency: {DB_ACTIVE_COUNT}", LOGGER_COMMAND
)
async def _get_group_cached(entity, event_cache) -> GroupSnapshot | None:
if not entity.group_id:
return None
if event_cache is not None and "group" in event_cache:
return event_cache["group"]
group = GroupMemoryCache.get_if_ready(entity.group_id, entity.channel_id)
if event_cache is not None:
event_cache["group"] = group
return group
def _module_in_block_string(module: str, value: str | None) -> bool:
if not value:
return False
return f"<{module}," in value
def _group_has_plugin_block(group, module: str) -> bool:
if not group:
return False
block_set = getattr(group, "block_plugin_set", None)
super_block_set = getattr(group, "superuser_block_plugin_set", None)
if block_set is not None or super_block_set is not None:
if block_set and module in block_set:
return True
if super_block_set and module in super_block_set:
return True
return False
block_plugin = getattr(group, "block_plugin", "") or ""
super_block_plugin = getattr(group, "superuser_block_plugin", "") or ""
return _module_in_block_string(module, block_plugin) or _module_in_block_string(
module, super_block_plugin
)
def _needs_auth_plugin(plugin: PluginInfo, group, entity) -> bool:
if plugin.block_type == BlockType.ALL and not plugin.status:
if group and getattr(group, "is_super", False):
return False
return True
if entity.group_id:
if plugin.block_type == BlockType.GROUP:
return True
return _group_has_plugin_block(group, plugin.module)
return plugin.block_type == BlockType.PRIVATE
def _needs_admin_check(plugin: PluginInfo) -> bool:
if plugin.admin_level and plugin.admin_level > 0:
return True
return plugin.plugin_type in {
PluginType.ADMIN,
PluginType.SUPERUSER,
PluginType.SUPER_AND_ADMIN,
}
async def _get_bot_data_cached(
bot_id: str, event_cache
) -> tuple[BotSnapshot | None, bool]:
if event_cache is not None and "bot_data" in event_cache:
return event_cache.get("bot_data"), event_cache.get("bot_timeout", False)
bot = await BotMemoryCache.get(bot_id)
if event_cache is not None:
event_cache["bot_data"] = bot
event_cache["bot_timeout"] = False
return bot, False
async def _get_admin_levels_cached(
session: Uninfo, entity, event_cache
) -> tuple[tuple[LevelUserSnapshot | None, LevelUserSnapshot | None] | None, bool]:
if event_cache is not None and "admin_levels" in event_cache:
return event_cache.get("admin_levels"), event_cache.get("admin_timeout", False)
levels = await LevelUserMemoryCache.get_levels(session.user.id, entity.group_id)
if event_cache is not None:
event_cache["admin_levels"] = levels
event_cache["admin_timeout"] = False
return levels, False
# 超时装饰器
async def with_timeout(coro, timeout=TIMEOUT_SECONDS, name=None):
"""带超时控制的协程执行
参数:
coro: 要执行的协程
timeout: 超时时间(秒)
name: 操作名称,用于日志记录
返回:
协程的返回值,或者在超时时抛出 TimeoutError
"""
try:
return await asyncio.wait_for(coro, timeout=timeout)
except asyncio.TimeoutError:
if name:
logger.error(f"{name} 操作超时 (>{timeout}s)", LOGGER_COMMAND)
# 更新熔断计数器
if name in CIRCUIT_BREAKERS:
CIRCUIT_BREAKERS[name]["failures"] += 1
if (
CIRCUIT_BREAKERS[name]["failures"]
>= CIRCUIT_BREAKERS[name]["threshold"]
and not CIRCUIT_BREAKERS[name]["active"]
):
CIRCUIT_BREAKERS[name]["active"] = True
CIRCUIT_BREAKERS[name]["reset_time"] = (
time.time() + CIRCUIT_RESET_TIME
)
logger.warning(
f"{name} 熔断器已激活,将在 {CIRCUIT_RESET_TIME} 秒后重置",
LOGGER_COMMAND,
)
raise
# 检查熔断状态
def check_circuit_breaker(name):
"""检查熔断器状态
参数:
name: 操作名称
返回:
bool: 是否已熔断
"""
if name not in CIRCUIT_BREAKERS:
return False
# 检查是否需要重置熔断器
if (
CIRCUIT_BREAKERS[name]["active"]
and time.time() > CIRCUIT_BREAKERS[name]["reset_time"]
):
CIRCUIT_BREAKERS[name]["active"] = False
CIRCUIT_BREAKERS[name]["failures"] = 0
logger.info(f"{name} 熔断器已重置", LOGGER_COMMAND)
return CIRCUIT_BREAKERS[name]["active"]
def _is_hidden_plugin(matcher: Matcher) -> bool:
plugin = matcher.plugin
if not plugin or not plugin.metadata:
return False
extra = plugin.metadata.extra or {}
return extra.get("plugin_type") == PluginType.HIDDEN
async def _fetch_user_readonly(
user_dao: DataAccess, user_id: str
) -> UserConsole | None:
return await with_timeout(
user_dao.safe_get_or_none(user_id=user_id), name="get_user"
)
async def _fetch_plugin(plugin_dao: DataAccess, module: str) -> PluginInfo | None:
return await with_timeout(
plugin_dao.safe_get_or_none(module=module), name="get_plugin"
)
async def get_plugin_and_user(
module: str,
user_id: str,
platform: str | None = None,
event_cache: dict | None = None,
need_user: bool = True,
) -> tuple[PluginInfo, UserConsole | None]:
"""Fetch plugin info and read user only when cost is required."""
user_dao = DataAccess(UserConsole)
plugin = None
if event_cache is not None:
plugin_cache = event_cache.setdefault("plugin_cache", {})
if module in plugin_cache:
plugin = plugin_cache[module]
if plugin is None:
plugin = await PluginInfoMemoryCache.get_by_module(module)
if event_cache is not None:
event_cache.setdefault("plugin_cache", {})[module] = plugin
plugin = cast(PluginInfo | None, plugin)
if not plugin:
raise PermissionExemption(f"plugin:{module} not found, skip permission check")
if plugin.plugin_type == PluginType.HIDDEN:
raise PermissionExemption(f"plugin {plugin.name}:{plugin.module} hidden, skip")
user = None
if need_user and plugin.cost_gold > 0:
if event_cache is not None:
user_cache = event_cache.setdefault("user_cache", {})
if user_id in user_cache:
user = user_cache[user_id]
else:
async with _db_section():
user = await _fetch_user_readonly(user_dao, user_id)
user_cache[user_id] = user
else:
async with _db_section():
user = await _fetch_user_readonly(user_dao, user_id)
return plugin, user
async def get_plugin_cost(
bot: Bot, user: UserConsole | None, plugin: PluginInfo, session: Uninfo
) -> int:
"""获取插件费用
参数:
bot: Bot
user: 用户数据
plugin: 插件数据
session: Uninfo
异常:
IsSuperuserException: 超级用户
IsSuperuserException: 超级用户
返回:
int: 调用插件金币费用
"""
cost_gold = await with_timeout(auth_cost(user, plugin, session), name="auth_cost")
if session.user.id in bot.config.superusers:
if plugin.plugin_type == PluginType.SUPERUSER:
raise IsSuperuserException()
if not plugin.limit_superuser:
raise IsSuperuserException()
return cost_gold
async def reduce_gold(user_id: str, module: str, cost_gold: int, session: Uninfo):
"""扣除用户金币
参数:
user_id: 用户id
module: 插件模块名称
cost_gold: 消耗金币
session: Uninfo
"""
user_dao = DataAccess(UserConsole)
try:
await with_timeout(
UserConsole.reduce_gold(
user_id,
cost_gold,
GoldHandle.PLUGIN,
module,
PlatformUtils.get_platform(session),
),
name="reduce_gold",
)
except InsufficientGold:
if u := await UserConsole.get_user(user_id):
u.gold = 0
await u.save(update_fields=["gold"])
except asyncio.TimeoutError:
logger.error(
f"扣除金币超时,用户: {user_id}, 金币: {cost_gold}",
LOGGER_COMMAND,
session=session,
)
# 清除缓存,使下次查询时从数据库获取最新数据
await user_dao.clear_cache(user_id=user_id)
logger.debug(f"调用功能花费金币: {cost_gold}", LOGGER_COMMAND, session=session)
# 辅助函数,用于记录每个 hook 的执行时间
async def time_hook(coro, name, time_dict):
start = time.time()
try:
# 检查熔断状态
if check_circuit_breaker(name):
logger.info(f"{name} 熔断器激活中,跳过执行", LOGGER_COMMAND)
time_dict[name] = "熔断跳过"
return
# 添加超时控制
return await with_timeout(coro, name=name)
except asyncio.TimeoutError:
time_dict[name] = f"超时 (>{TIMEOUT_SECONDS}s)"
finally:
if name not in time_dict:
time_dict[name] = f"{time.time() - start:.3f}s"
async def _enter_hooks_section():
"""尝试获取全局信号量并更新计数器,超时则抛出 PermissionExemption。"""
global HOOKS_ACTIVE_COUNT
# 队列模式:如果达到上限,协程将排队等待直到获取到信号量
await HOOKS_SEMAPHORE.acquire()
async with HOOKS_ACTIVE_LOCK:
HOOKS_ACTIVE_COUNT += 1
_debug_log(f"当前并发权限检查数量: {HOOKS_ACTIVE_COUNT}", LOGGER_COMMAND)
async def _leave_hooks_section():
"""释放信号量并更新计数器。"""
global HOOKS_ACTIVE_COUNT
from contextlib import suppress
with suppress(Exception):
HOOKS_SEMAPHORE.release()
async with HOOKS_ACTIVE_LOCK:
HOOKS_ACTIVE_COUNT -= 1
# 保证计数不为负
HOOKS_ACTIVE_COUNT = max(HOOKS_ACTIVE_COUNT, 0)
_debug_log(f"当前并发权限检查数量: {HOOKS_ACTIVE_COUNT}", LOGGER_COMMAND)
async def auth_ban_fast(
matcher: Matcher, event: Event, bot: Bot, session: Uninfo
) -> None:
"""快速 ban 检测(仅使用内存缓存),用于前置快速裁决。"""
entity = get_entity_ids(session)
event_cache = _get_event_cache(event, session, entity)
if event_cache is not None and event_cache.get("ban_state") is True:
raise SkipPluginException("user or group banned (cached)")
if entity.user_id in bot.config.superusers:
if event_cache is not None:
event_cache["ban_state"] = False
return
if entity.group_id and await is_ban(None, entity.group_id):
if event_cache is not None:
event_cache["ban_state"] = True
raise SkipPluginException("group banned (fast)")
if entity.user_id and await is_ban(entity.user_id, entity.group_id):
if event_cache is not None:
event_cache["ban_state"] = True
raise SkipPluginException("user banned (fast)")
if event_cache is not None:
event_cache["ban_state"] = False
async def route_precheck(
matcher: Matcher,
event: Event,
session: Uninfo,
message: UniMsg,
) -> bool:
module = matcher.plugin_name or ""
if not module:
return False
if _is_hidden_plugin(matcher):
return False
entity = get_entity_ids(session)
event_cache = _get_event_cache(event, session, entity)
text = _get_message_text(message, event_cache)
route_modules = await _get_route_context(text, event_cache)
await _ensure_route_index()
if module in _ROUTE_MODULES_WITH_COMMANDS and module not in route_modules:
if event_cache is not None:
event_cache["route_skip"] = True
return True
return False
async def auth_precheck(
matcher: Matcher,
event: Event,
bot: Bot,
session: Uninfo,
message: UniMsg,
) -> None:
"""轻量前置检查:命令路由 + 必要管理员权限。"""
module = matcher.plugin_name or ""
if not module:
return
if _is_hidden_plugin(matcher):
return
entity = get_entity_ids(session)
if session.user.id in bot.config.superusers:
return
plugin = cast(PluginInfo | None, await PluginInfoMemoryCache.get_by_module(module))
if not plugin:
return
if plugin.plugin_type == PluginType.SUPERUSER:
raise SkipPluginException("超级管理员权限不足...")
if _needs_admin_check(plugin):
await LevelUserMemoryCache.ensure_fresh()
levels = await LevelUserMemoryCache.get_levels(session.user.id, entity.group_id)
await auth_admin(plugin, session, cached_levels=levels)
async def auth(
matcher: Matcher,
event: Event,
bot: Bot,
session: Uninfo,
message: UniMsg,
*,
skip_ban: bool = False,
):
"""权限检查
参数:
matcher: matcher
event: Event
bot: bot
session: Uninfo
message: UniMsg
"""
start_time = time.time()
cost_gold = 0
ignore_flag = False
entity = get_entity_ids(session)
module = matcher.plugin_name or ""
event_cache = _get_event_cache(event, session, entity)
auth_allowed = None
auth_result_cache = None
admin_checked_pre = False
# 用于记录各个 hook 的执行时间
hook_times = {}
hooks_time = 0 # 初始化 hooks_time 变量
# 记录是否已进入 hooks 区域(用于 finally 中释放)
entered_hooks = False
try:
if not module:
raise PermissionExemption("Matcher插件名称不存在...")
if event_cache is not None:
auth_result_cache = event_cache.setdefault("auth_result", {})
cached_result = auth_result_cache.get(module)
if cached_result is not None:
allowed, reason = cached_result
if not allowed:
raise SkipPluginException(reason or "auth cached skip")
return
if _is_hidden_plugin(matcher):
raise PermissionExemption(f"plugin {module} hidden, skip")
if event_cache is not None and event_cache.get("ban_state") is True:
raise SkipPluginException("user or group banned (cached)")
text = _get_message_text(message, event_cache)
route_modules = await _get_route_context(text, event_cache)
await _ensure_route_index()
route_skip_checks = (
module in _ROUTE_MODULES_WITH_COMMANDS and module not in route_modules
)
if route_skip_checks:
if event_cache is not None:
event_cache["route_skip"] = True
hook_times["route"] = "miss"
auth_allowed = True
return
platform = PlatformUtils.get_platform(session)
# 获取插件和用户数据
plugin_user_start = time.time()
try:
plugin, user = await with_timeout(
get_plugin_and_user(
module,
entity.user_id,
platform,
event_cache=event_cache,
need_user=not route_skip_checks,
),
name="get_plugin_and_user",
)
hook_times["get_plugin_user"] = f"{time.time() - plugin_user_start:.3f}s"
except asyncio.TimeoutError:
logger.error(
f"获取插件和用户数据超时,模块: {module}",
LOGGER_COMMAND,
session=session,
)
raise PermissionExemption("获取插件和用户数据超时,请稍后再试...")
if not route_skip_checks and _needs_admin_check(plugin):
if plugin.plugin_type in {
PluginType.SUPERUSER,
PluginType.SUPER_AND_ADMIN,
}:
if session.user.id in bot.config.superusers:
hook_times["auth_admin"] = "superuser"
admin_checked_pre = True
elif plugin.plugin_type == PluginType.SUPERUSER:
raise SkipPluginException("超级管理员权限不足...")
if not admin_checked_pre:
await LevelUserMemoryCache.ensure_fresh()
admin_levels = None
admin_timeout = False
if event_cache is not None:
admin_levels, admin_timeout = await _get_admin_levels_cached(
session, entity, event_cache
)
if admin_timeout:
hook_times["auth_admin"] = "timeout"
else:
admin_start = time.time()
await auth_admin(plugin, session, cached_levels=admin_levels)
hook_times["auth_admin"] = f"{time.time() - admin_start:.3f}s(pre)"
admin_checked_pre = True
ban_cache_state = None
if event_cache is not None:
ban_cache_state = event_cache.get("ban_state")
if skip_ban:
if ban_cache_state is True:
hook_times["auth_ban"] = "cached"
raise SkipPluginException("user or group banned (cached)")
if ban_cache_state is None:
ban_start = time.time()
try:
await auth_ban(matcher, bot, session, plugin)
hook_times["auth_ban"] = f"{time.time() - ban_start:.3f}s"
if event_cache is not None:
event_cache["ban_state"] = False
except SkipPluginException:
hook_times["auth_ban"] = f"{time.time() - ban_start:.3f}s"
if event_cache is not None:
event_cache["ban_state"] = True
raise
else:
hook_times["auth_ban"] = "skipped"
else:
if ban_cache_state is True:
hook_times["auth_ban"] = "cached"
raise SkipPluginException("user or group banned (cached)")
if ban_cache_state is None:
ban_start = time.time()
try:
await auth_ban(matcher, bot, session, plugin)
hook_times["auth_ban"] = f"{time.time() - ban_start:.3f}s"
if event_cache is not None:
event_cache["ban_state"] = False
except SkipPluginException:
hook_times["auth_ban"] = f"{time.time() - ban_start:.3f}s"
if event_cache is not None:
event_cache["ban_state"] = True
raise
else:
hook_times["auth_ban"] = "cached"
# 获取插件费用
if not route_skip_checks and plugin.cost_gold > 0:
cost_start = time.time()
try:
cost_gold = await with_timeout(
get_plugin_cost(bot, user, plugin, session), name="get_plugin_cost"
)
hook_times["cost_gold"] = f"{time.time() - cost_start:.3f}s"
except asyncio.TimeoutError:
logger.error(
f"获取插件费用超时,模块: {module}", LOGGER_COMMAND, session=session
)
# 继续执行,不阻止权限检查
else:
hook_times["cost_gold"] = "skipped"
# 执行 bot_filter
bot_filter(session)
group = await _get_group_cached(entity, event_cache)
bot_data = None
bot_timeout = False
if event_cache is not None:
bot_data, bot_timeout = await _get_bot_data_cached(bot.self_id, event_cache)
admin_levels = None
admin_timeout = False
if (
not admin_checked_pre
and plugin.admin_level
and event_cache is not None
and not route_skip_checks
):
admin_levels, admin_timeout = await _get_admin_levels_cached(
session, entity, event_cache
)
# 并行执行所有 hook 检查,并记录执行时间
hooks_start = time.time()
# 创建所有 hook 任务
hook_tasks = []
if event_cache is None:
hook_tasks.append(
time_hook(auth_bot(plugin, bot.self_id), "auth_bot", hook_times)
)
else:
if bot_timeout:
hook_times["auth_bot"] = "timeout"
else:
hook_tasks.append(
time_hook(
auth_bot(
plugin,
bot.self_id,
bot_data=bot_data,
skip_fetch=True,
),
"auth_bot",
hook_times,
)
)
if session.user.id in bot.config.superusers:
hook_times["auth_group"] = "superuser"
else:
hook_tasks.append(
time_hook(
auth_group(plugin, group, text, entity.group_id),
"auth_group",
hook_times,
)
)
if not route_skip_checks and plugin.admin_level and not admin_checked_pre:
if event_cache is None:
hook_tasks.append(
time_hook(auth_admin(plugin, session), "auth_admin", hook_times)
)
else:
if admin_timeout:
hook_times["auth_admin"] = "timeout"
else:
hook_tasks.append(
time_hook(
auth_admin(plugin, session, cached_levels=admin_levels),
"auth_admin",
hook_times,
)
)
else:
hook_times.setdefault("auth_admin", "skipped")
if session.user.id in bot.config.superusers:
hook_times["auth_plugin"] = "superuser"
elif not route_skip_checks and _needs_auth_plugin(plugin, group, entity):
hook_tasks.append(
time_hook(
auth_plugin(
plugin,
group,
session,
event,
skip_group_block=session.user.id in bot.config.superusers,
),
"auth_plugin",
hook_times,
)
)
else:
hook_times["auth_plugin"] = "skipped"
if not route_skip_checks:
has_limits = await _has_limits_cached(module, event_cache)
if has_limits:
hook_tasks.append(
time_hook(auth_limit(plugin, session), "auth_limit", hook_times)
)
else:
hook_times["auth_limit"] = "skipped"
else:
hook_times["auth_limit"] = "skipped"
if hook_tasks:
# 进入 hooks 并行检查区域(会在高并发时排队)
await _enter_hooks_section()
entered_hooks = True
# 使用 gather 并行执行所有 hook,但添加总体超时控制
try:
await with_timeout(
asyncio.gather(*hook_tasks),
timeout=TIMEOUT_SECONDS * 2, # 给总体执行更多时间
name="auth_hooks_gather",
)
except asyncio.TimeoutError:
logger.error(
f"权限检查 hooks 总体执行超时,模块: {module}",
LOGGER_COMMAND,
session=session,
)
# 不抛出异常,允许继续执行
hooks_time = time.time() - hooks_start
auth_allowed = True
except SkipPluginException as e:
LimitManager.unblock(module, entity.user_id, entity.group_id, entity.channel_id)
logger.info(str(e), LOGGER_COMMAND, session=session)
ignore_flag = True
auth_allowed = False
except IsSuperuserException:
logger.debug("超级用户跳过权限检测...", LOGGER_COMMAND, session=session)
auth_allowed = True
except PermissionExemption as e:
logger.info(str(e), LOGGER_COMMAND, session=session)
auth_allowed = True
finally:
# 如果进入过 hooks 区域,确保释放信号量(即使上层处理抛出了异常)
if entered_hooks:
try:
await _leave_hooks_section()
except Exception:
logger.error(
"释放 hooks 信号量时出错",
LOGGER_COMMAND,
session=session,
)
if auth_result_cache is not None and auth_allowed is not None:
auth_result_cache[module] = (auth_allowed, None)
# 扣除金币
if not ignore_flag and cost_gold > 0:
gold_start = time.time()
try:
await with_timeout(
reduce_gold(entity.user_id, module, cost_gold, session),
name="reduce_gold",
)
hook_times["reduce_gold"] = f"{time.time() - gold_start:.3f}s"
except asyncio.TimeoutError:
logger.error(
f"扣除金币超时,模块: {module}", LOGGER_COMMAND, session=session
)
# 记录总执行时间
total_time = time.time() - start_time
if total_time > WARNING_THRESHOLD: # 如果总时间超过500ms,记录详细信息
logger.warning(
f"权限检查耗时过长: {total_time:.3f}s, 模块: {module}, "
f"hooks时间: {hooks_time:.3f}s, "
f"详情: {hook_times}",
LOGGER_COMMAND,
session=session,
)
if ignore_flag:
raise IgnoredException("权限检测 ignore")