mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-29 00:32:06 +08:00
* 添加测试插件 * ✨ 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>
1077 lines
36 KiB
Python
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")
|