diff --git a/zhenxun/builtin_plugins/hooks/auth_checker.py b/zhenxun/builtin_plugins/hooks/auth_checker.py deleted file mode 100644 index 4b38c16c..00000000 --- a/zhenxun/builtin_plugins/hooks/auth_checker.py +++ /dev/null @@ -1,461 +0,0 @@ -""" -优化后的权限检查系统设计 - -主要改进: -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() diff --git a/zhenxun/services/auth_snapshot/builder.py b/zhenxun/services/auth_snapshot/builder.py index 59ac3ed8..39b8b77d 100644 --- a/zhenxun/services/auth_snapshot/builder.py +++ b/zhenxun/services/auth_snapshot/builder.py @@ -231,6 +231,26 @@ class SnapshotBuilder: else: # sqlite return "?" + @classmethod + def _get_null_cast(cls, db_type: str, col_type: str) -> str: + """获取 NULL 的类型转换语法 + + 参数: + db_type: 数据库类型 + col_type: 目标列类型 (bigint, int, etc.) + + 返回: + str: 带类型转换的 NULL + """ + if db_type == DB_TYPE_POSTGRES: + return f"NULL::{col_type}" + elif db_type == DB_TYPE_MYSQL: + # MySQL UNION 会自动推断类型,但显式转换更安全 + return "CAST(NULL AS SIGNED)" + else: # sqlite + # SQLite 是动态类型,NULL 不需要转换 + return "NULL" + @classmethod def _build_user_data_sql( cls, user_id: str, group_id: str | None, db_type: str @@ -259,18 +279,22 @@ class SnapshotBuilder: param_idx += 1 return placeholder + # 获取类型转换的 NULL(PostgreSQL 需要显式类型) + null_bigint = cls._get_null_cast(db_type, "bigint") + null_int = cls._get_null_cast(db_type, "integer") + # 1. 用户金币 queries.append(f""" - SELECT 'user' as query_type, gold, NULL as user_level, - NULL as ban_time, NULL as duration + SELECT 'user' as query_type, gold, {null_int} as user_level, + {null_bigint} as ban_time, {null_int} as duration FROM user_console WHERE user_id = {ph()} """) params.append(user_id) # 2. 全局权限等级 queries.append(f""" - SELECT 'level_global' as query_type, NULL as gold, user_level, - NULL as ban_time, NULL as duration + SELECT 'level_global' as query_type, {null_int} as gold, user_level, + {null_bigint} as ban_time, {null_int} as duration FROM level_users WHERE user_id = {ph()} AND group_id IS NULL """) params.append(user_id) @@ -278,8 +302,8 @@ class SnapshotBuilder: # 3. 群组权限等级 if group_id: queries.append(f""" - SELECT 'level_group' as query_type, NULL as gold, user_level, - NULL as ban_time, NULL as duration + SELECT 'level_group' as query_type, {null_int} as gold, user_level, + {null_bigint} as ban_time, {null_int} as duration FROM level_users WHERE user_id = {ph()} AND group_id = {ph()} """) @@ -287,7 +311,8 @@ class SnapshotBuilder: # 4. 用户全局 ban queries.append(f""" - SELECT 'ban_user_global' as query_type, NULL as gold, NULL as user_level, + SELECT 'ban_user_global' as query_type, + {null_int} as gold, {null_int} as user_level, ban_time, duration FROM ban_console WHERE user_id = {ph()} AND group_id IS NULL @@ -297,7 +322,8 @@ class SnapshotBuilder: # 5. 用户群组 ban if group_id: queries.append(f""" - SELECT 'ban_user_group' as query_type, NULL as gold, NULL as user_level, + SELECT 'ban_user_group' as query_type, + {null_int} as gold, {null_int} as user_level, ban_time, duration FROM ban_console WHERE user_id = {ph()} AND group_id = {ph()} @@ -306,7 +332,8 @@ class SnapshotBuilder: # 6. 群组 ban queries.append(f""" - SELECT 'ban_group' as query_type, NULL as gold, NULL as user_level, + SELECT 'ban_group' as query_type, + {null_int} as gold, {null_int} as user_level, ban_time, duration FROM ban_console WHERE user_id = {ph()} AND group_id = {ph()}