Files
zhenxun_bot/zhenxun/builtin_plugins/hooks/auth_checker.py
T

462 lines
17 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.
"""
优化后的权限检查系统设计
主要改进:
1. 优先级机制:将检查分为多个优先级阶段,高优先级检查失败时立即退出
2. 早期退出:避免不必要的检查执行,提高性能
3. 统一数据上下文:在开始前统一获取所有需要的数据
4. 检查结果缓存:对相同请求缓存检查结果
5. 统一的错误处理:所有检查使用统一的超时和错误处理机制
"""
import asyncio
from collections.abc import Callable
from dataclasses import dataclass
from enum import IntEnum
import time
from typing import Any
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 tortoise.exceptions import IntegrityError
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.user_console import UserConsole
from zhenxun.services.data_access import DataAccess
from zhenxun.services.log import logger
from zhenxun.utils.enum import GoldHandle, PluginType
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
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,
)
# 超时设置(秒)—— DataAccess 内部已对单次 DB / 缓存访问做了自己的超时控制;
# 这里主要用于控制单个权限检查步骤的上限时间。
TIMEOUT_SECONDS = 5.0
# 检查优先级
class CheckPriority(IntEnum):
"""检查优先级,数值越小优先级越高"""
CRITICAL = 1 # 关键检查:ban、bot状态
HIGH = 2 # 高优先级:插件状态、群组状态
MEDIUM = 3 # 中等优先级:管理员权限、限制
LOW = 4 # 低优先级:金币检查
@dataclass
class AuthContext:
"""权限检查上下文,统一管理所有需要的数据"""
plugin: PluginInfo
user: UserConsole
session: Uninfo
matcher: Matcher
bot: Bot
message: UniMsg
group: GroupConsole | None = None
bot_id: str = ""
entity: Any = None
@dataclass
class CheckResult:
"""检查结果"""
success: bool
error: Exception | None = None
execution_time: float = 0.0
cached: bool = False
class AuthChecker:
"""优化的权限检查器"""
async def _execute_check(
self,
check_func: Callable,
check_name: str,
priority: CheckPriority,
context: AuthContext,
**kwargs,
) -> CheckResult:
"""执行单个检查"""
start_time = time.time()
try:
# 执行检查函数
await asyncio.wait_for(check_func(**kwargs), timeout=TIMEOUT_SECONDS)
result = CheckResult(success=True, execution_time=time.time() - start_time)
except SkipPluginException as e:
result = CheckResult(
success=False,
error=e,
execution_time=time.time() - start_time,
)
except asyncio.TimeoutError:
logger.error(
f"{check_name} 检查超时", LOGGER_COMMAND, session=context.session
)
# 超时时根据优先级决定是否继续
if priority <= CheckPriority.HIGH:
result = CheckResult(
success=False,
error=PermissionExemption(f"{check_name} 检查超时"),
execution_time=time.time() - start_time,
)
else:
# 低优先级检查超时,允许继续
result = CheckResult(
success=True, execution_time=time.time() - start_time
)
except Exception as e:
logger.error(
f"{check_name} 检查失败: {e}", LOGGER_COMMAND, session=context.session
)
result = CheckResult(
success=False, error=e, execution_time=time.time() - start_time
)
return result
async def _load_context(
self, matcher: Matcher, bot: Bot, session: Uninfo, message: UniMsg
) -> AuthContext:
"""加载权限检查上下文数据"""
entity = get_entity_ids(session)
module = matcher.plugin_name or ""
if not module:
raise PermissionExemption("Matcher插件名称不存在...")
# 并行获取所有需要的数据。
# DataAccess 内部已经有 Redis 缓存和 DB 超时控制,这里只做一次整体超时保护,
# 不再额外手动走 CacheRoot 之类的二级 fallback,避免重复访问 Redis。
user_dao = DataAccess(UserConsole)
plugin_dao = DataAccess(PluginInfo)
group_dao = DataAccess(GroupConsole) if entity.group_id else None
# 为了更清晰地定位超时来源,创建具名 task
task_items = [
(
"plugin",
asyncio.create_task(
plugin_dao.safe_get_or_none(module=module),
name="authctx:plugin",
),
),
(
"user",
asyncio.create_task(
user_dao.get_by_func_or_none(
UserConsole.get_user, False, user_id=entity.user_id
),
name="authctx:user",
),
),
]
if entity.group_id and group_dao:
task_items.append(
(
"group",
asyncio.create_task(
group_dao.safe_get_or_none(
group_id=entity.group_id, channel_id__isnull=True
),
name="authctx:group",
),
)
)
task_list = [item[1] for item in task_items]
start_ts = time.monotonic()
try:
results = await asyncio.wait_for(
asyncio.gather(*task_list), timeout=TIMEOUT_SECONDS
)
except asyncio.TimeoutError:
# DataAccess 本身已经利用了 Redis / DB 缓存,这里整体超时直接视为失败,
# 避免在 Redis 也不稳定时再叠加一层「从缓存再试一次」的复杂 fallback。
elapsed = time.monotonic() - start_ts
def _describe(name: str, task: asyncio.Task) -> str:
if not task.done():
return f"{name}=pending"
if task.cancelled():
return f"{name}=cancelled"
exc = task.exception()
if exc:
return f"{name}=error({exc})"
return f"{name}=ok"
states = [_describe(name, task) for name, task in task_items]
timeout_msg = (
f"加载权限检查所需数据超时,模块: {module},"
f"耗时: {elapsed:.2f}s,状态: {states}"
)
logger.error(timeout_msg, LOGGER_COMMAND, session=session)
for _, task in task_items:
task.cancel()
raise PermissionExemption("获取权限检查所需数据超时,请稍后再试...")
except IntegrityError:
# 获取用户时可能因为 uid 竞争导致唯一约束冲突,稍作等待并重试多次
logger.warning(
f"检测到重复创建用户,准备重试获取用户,模块: {module}",
LOGGER_COMMAND,
session=session,
)
# 重新获取插件(走 DataAccess 缓存,代价很小)
plugin = await plugin_dao.safe_get_or_none(module=module)
user = None
group = None
# 最多重试 3 次,逐步增加等待时间
for attempt in range(3):
try:
user = await user_dao.get_by_func_or_none(
UserConsole.get_user, False, user_id=entity.user_id
)
if entity.group_id and group_dao:
group = await group_dao.safe_get_or_none(
group_id=entity.group_id, channel_id__isnull=True
)
break
except IntegrityError as e:
if attempt == 2:
logger.error(
f"多次尝试创建用户仍然出现唯一约束冲突,模块: {module}",
LOGGER_COMMAND,
session=session,
e=e,
)
raise PermissionExemption("重复创建用户,请稍后再试...") from e
await asyncio.sleep(0.5 * (attempt + 1))
else:
plugin = results[0]
user = results[1]
group = results[2] if len(results) > 2 else None
if not plugin:
raise PermissionExemption(f"插件:{module} 数据不存在...")
if plugin.plugin_type == PluginType.HIDDEN:
raise PermissionExemption(
f"插件: {plugin.name}:{plugin.module} 为HIDDEN..."
)
if not user:
raise PermissionExemption("用户数据不存在...")
return AuthContext(
plugin=plugin,
user=user,
group=group,
bot_id=bot.self_id,
session=session,
matcher=matcher,
bot=bot,
message=message,
entity=entity,
)
async def check(
self, matcher: Matcher, event: Event, bot: Bot, session: Uninfo, message: UniMsg
):
"""执行权限检查(优化版本)"""
start_time = time.time()
cost_gold = 0
ignore_flag = False
hook_times = {}
try:
bot_filter(session)
# 1. 加载上下文数据(可能包含多次数据库/缓存访问,这里单独计时)
ctx_start = time.time()
context = await self._load_context(matcher, bot, session, message)
hook_times["load_context"] = f"{time.time() - ctx_start:.3f}s"
# 2. 按优先级执行检查
# 阶段1:关键检查(ban、bot状态)
critical_checks = [
(
"auth_ban",
CheckPriority.CRITICAL,
lambda: auth_ban(
context.matcher, context.bot, context.session, context.plugin
),
),
(
"auth_bot",
CheckPriority.CRITICAL,
lambda: auth_bot(context.plugin, context.bot_id),
),
]
for check_name, priority, check_func in critical_checks:
result = await self._execute_check(
check_func, check_name, priority, context
)
hook_times[check_name] = f"{result.execution_time:.3f}s"
if not result.success:
if isinstance(result.error, SkipPluginException):
ignore_flag = True
raise result.error
raise result.error or PermissionExemption(f"{check_name} 检查失败")
# 4. 阶段2:高优先级检查(插件状态、群组状态)
high_priority_checks = [
(
"auth_plugin",
CheckPriority.HIGH,
lambda: auth_plugin(
context.plugin, context.group, context.session, event
),
),
(
"auth_group",
CheckPriority.HIGH,
lambda: auth_group(
context.plugin,
context.group,
context.message,
context.entity.group_id,
),
),
]
# 并行执行高优先级检查
high_tasks = []
for check_name, priority, check_func in high_priority_checks:
task = self._execute_check(check_func, check_name, priority, context)
high_tasks.append((check_name, task))
high_results = await asyncio.gather(*[task for _, task in high_tasks])
for (check_name, _), result in zip(high_tasks, high_results):
hook_times[check_name] = f"{result.execution_time:.3f}s"
if not result.success:
if isinstance(result.error, SkipPluginException):
ignore_flag = True
raise result.error
raise result.error or PermissionExemption(f"{check_name} 检查失败")
# 5. 阶段3:中等优先级检查(管理员权限、限制)
medium_checks = [
(
"auth_admin",
CheckPriority.MEDIUM,
lambda: auth_admin(context.plugin, context.session),
),
(
"auth_limit",
CheckPriority.MEDIUM,
lambda: auth_limit(context.plugin, context.session),
),
]
# 并行执行中等优先级检查
medium_tasks = []
for check_name, priority, check_func in medium_checks:
task = self._execute_check(check_func, check_name, priority, context)
medium_tasks.append((check_name, task))
medium_results = await asyncio.gather(*[task for _, task in medium_tasks])
for (check_name, _), result in zip(medium_tasks, medium_results):
hook_times[check_name] = f"{result.execution_time:.3f}s"
if not result.success:
if isinstance(result.error, SkipPluginException):
ignore_flag = True
raise result.error
raise result.error or PermissionExemption(f"{check_name} 检查失败")
# 6. 阶段4:低优先级检查(金币检查)
try:
cost_start = time.time()
cost_gold = await asyncio.wait_for(
auth_cost(context.user, context.plugin, context.session),
timeout=TIMEOUT_SECONDS,
)
if context.session.user.id in bot.config.superusers:
if context.plugin.plugin_type == PluginType.SUPERUSER:
raise IsSuperuserException()
if not context.plugin.limit_superuser:
raise IsSuperuserException()
hook_times["cost_gold"] = f"{time.time() - cost_start:.3f}s"
except asyncio.TimeoutError:
logger.error(
f"获取插件费用超时,模块: {context.plugin.module}",
LOGGER_COMMAND,
session=session,
)
except SkipPluginException as e:
LimitManager.unblock(
matcher.plugin_name or "",
get_entity_ids(session).user_id,
get_entity_ids(session).group_id,
get_entity_ids(session).channel_id,
)
logger.info(str(e), LOGGER_COMMAND, session=session)
ignore_flag = True
except IsSuperuserException:
logger.debug("超级用户跳过权限检测...", LOGGER_COMMAND, session=session)
except PermissionExemption as e:
logger.info(str(e), LOGGER_COMMAND, session=session)
# 扣除金币
if not ignore_flag and cost_gold > 0:
try:
await asyncio.wait_for(
UserConsole.reduce_gold(
get_entity_ids(session).user_id,
cost_gold,
GoldHandle.PLUGIN,
matcher.plugin_name or "",
PlatformUtils.get_platform(session),
),
timeout=TIMEOUT_SECONDS,
)
hook_times["reduce_gold"] = f"{time.time() - start_time:.3f}s"
except asyncio.TimeoutError:
logger.error(
f"扣除金币超时,模块: {matcher.plugin_name}",
LOGGER_COMMAND,
session=session,
)
# 记录总执行时间
total_time = time.time() - start_time
if total_time > WARNING_THRESHOLD:
logger.warning(
f"权限检查耗时过长: {total_time:.3f}s, "
f"模块: {matcher.plugin_name}, 详情: {hook_times}",
LOGGER_COMMAND,
session=session,
)
if ignore_flag:
raise IgnoredException("权限检测 ignore")
_auth_checker = AuthChecker()