From 47ec5bc7b99a3520ebd2856475f0f32d4bbef43e Mon Sep 17 00:00:00 2001 From: HibiKier <775757368@qq.com> Date: Tue, 23 Dec 2025 15:15:35 +0800 Subject: [PATCH] refactor: enhance ban handling with caching and optimize user/group ban checks --- .../builtin_plugins/hooks/auth/auth_ban.py | 178 ++++-------------- zhenxun/models/ban_console.py | 65 +++++-- 2 files changed, 85 insertions(+), 158 deletions(-) diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_ban.py b/zhenxun/builtin_plugins/hooks/auth/auth_ban.py index 0ab3cb70..de78eb18 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_ban.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_ban.py @@ -9,10 +9,11 @@ from nonebot_plugin_uninfo import Uninfo from zhenxun.configs.config import Config from zhenxun.models.ban_console import BanConsole from zhenxun.models.plugin_info import PluginInfo +from zhenxun.services.cache import CacheRoot from zhenxun.services.data_access import DataAccess from zhenxun.services.db_context import DB_TIMEOUT_SECONDS from zhenxun.services.log import logger -from zhenxun.utils.enum import PluginType +from zhenxun.utils.enum import CacheType, PluginType from zhenxun.utils.utils import EntityIDs, get_entity_ids from .config import LOGGER_COMMAND, WARNING_THRESHOLD @@ -49,105 +50,6 @@ async def calculate_ban_time(ban_record: BanConsole | None) -> int: return 0 -async def is_ban(user_id: str | None, group_id: str | None) -> int: - """检查用户或群组是否被ban - - 参数: - user_id: 用户ID - group_id: 群组ID - - 返回: - int: ban的剩余时间,0表示未被ban - """ - if not user_id and not group_id: - return 0 - - start_time = time.time() - ban_dao = DataAccess(BanConsole) - - # 分别获取用户在群组中的ban记录和全局ban记录 - group_user = None - user = None - - try: - # 并行查询用户和群组的 ban 记录 - tasks = [] - if user_id and group_id: - tasks.append(ban_dao.safe_get_or_none(user_id=user_id, group_id=group_id)) - if user_id: - tasks.append( - ban_dao.safe_get_or_none(user_id=user_id, group_id__isnull=True) - ) - - # 等待所有查询完成,添加超时控制(使用更短的超时时间,快速失败) - if tasks: - try: - # 使用更短的超时时间(1.5秒),避免在高并发下等待太久 - # 如果查询超时,视为未ban,允许继续执行 - ban_records = await asyncio.wait_for( - asyncio.gather(*tasks, return_exceptions=True), - timeout=min(DB_TIMEOUT_SECONDS, 1.5), - ) - # 处理可能的异常 - valid_records = [] - for record in ban_records: - if isinstance(record, Exception): - logger.warning( - f"查询ban记录时出现异常: {record}", - LOGGER_COMMAND, - ) - continue - valid_records.append(record) - - if len(tasks) == 2: - group_user = valid_records[0] if len(valid_records) > 0 else None - user = valid_records[1] if len(valid_records) > 1 else None - elif user_id and group_id: - group_user = valid_records[0] if valid_records else None - else: - user = valid_records[0] if valid_records else None - except asyncio.TimeoutError: - logger.warning( - f"查询ban记录超时(视为未ban): " - f"user_id={user_id}, group_id={group_id}", - LOGGER_COMMAND, - ) - return 0 - - # 检查记录并计算ban时间 - results = [] - if group_user: - results.append(group_user) - if user: - results.append(user) - - # 如果没有找到记录,返回0 - if not results: - return 0 - - logger.debug(f"查询到的ban记录: {results}", LOGGER_COMMAND) - # 检查所有记录,找出最严格的ban(时间最长的) - max_ban_time: int = 0 - for result in results: - if result.duration > 0 or result.duration == -1: - # 直接计算ban时间,避免再次查询数据库 - ban_time = await calculate_ban_time(result) - if ban_time == -1 or ban_time > max_ban_time: - max_ban_time = ban_time - - return max_ban_time - finally: - # 记录执行时间 - elapsed = time.time() - start_time - if elapsed > WARNING_THRESHOLD: # 记录耗时超过500ms的检查 - logger.warning( - f"is_ban 耗时: {elapsed:.3f}s", - LOGGER_COMMAND, - session=user_id, - group_id=group_id, - ) - - def check_plugin_type(matcher: Matcher) -> bool: """判断插件类型是否是隐藏插件 @@ -190,45 +92,22 @@ def format_time(time_val: float) -> str: return time_str -async def group_handle(group_id: str) -> None: - """群组ban检查 - - 参数: - group_id: 群组id - - 异常: - SkipPluginException: 群组处于黑名单 - """ - start_time = time.time() - try: - if await is_ban(None, group_id): - raise SkipPluginException("群组处于黑名单中...") - finally: - # 记录执行时间 - elapsed = time.time() - start_time - if elapsed > WARNING_THRESHOLD: # 记录耗时超过500ms的检查 - logger.warning( - f"group_handle 耗时: {elapsed:.3f}s", - LOGGER_COMMAND, - group_id=group_id, - ) - - -async def user_handle(plugin: PluginInfo, entity: EntityIDs, session: Uninfo) -> None: +async def user_handle( + plugin: PluginInfo, entity: EntityIDs, session: Uninfo, time_val: int +) -> None: """用户ban检查 参数: module: 插件模块名 entity: 实体ID信息 session: Uninfo - + time_val: 剩余ban时间 异常: SkipPluginException: 用户处于黑名单 """ start_time = time.time() try: ban_result = Config.get_config("hook", "BAN_RESULT") - time_val = await is_ban(entity.user_id, entity.group_id) if not time_val: return time_str = format_time(time_val) @@ -284,24 +163,37 @@ async def auth_ban( entity = get_entity_ids(session) if entity.user_id in bot.config.superusers: return - if entity.group_id: - try: - await asyncio.wait_for( - group_handle(entity.group_id), timeout=DB_TIMEOUT_SECONDS - ) - except asyncio.TimeoutError: - logger.error(f"群组ban检查超时: {entity.group_id}", LOGGER_COMMAND) - # 超时时不阻塞,继续执行 - if entity.user_id: - try: - await asyncio.wait_for( - user_handle(plugin, entity, session), - timeout=DB_TIMEOUT_SECONDS, + cache_key = f"{entity.user_id}_{entity.group_id}" + + results = await CacheRoot.get(CacheType.BAN, cache_key) + if not results: + results = await BanConsole.is_ban(entity.user_id, entity.group_id) + await CacheRoot.set( + CacheType.BAN, + cache_key, + results or DataAccess._NULL_RESULT, + ) + else: + tmp_results: list[BanConsole] = [] + for r in results: + tmp_results.append(CacheRoot._deserialize_value(r, BanConsole)) + results = tmp_results + + for result in results: + if not result.user_id and result.group_id: + logger.debug( + f"群组{result.group_id}被ban: {result}", + target=f"{result.group_id}:{entity.user_id}", ) - except asyncio.TimeoutError: - logger.error(f"用户ban检查超时: {entity.user_id}", LOGGER_COMMAND) - # 超时时不阻塞,继续执行 + raise SkipPluginException(f"群组: {result.group_id} 处于黑名单中...") + if result.user_id: + logger.debug( + f"用户{result.user_id}被ban: {result}", + target=f"{result.group_id}:{entity.user_id}", + ) + await user_handle(plugin, entity, session, result.duration) + finally: # 记录总执行时间 elapsed = time.time() - start_time diff --git a/zhenxun/models/ban_console.py b/zhenxun/models/ban_console.py index 90127f45..55b55cb9 100644 --- a/zhenxun/models/ban_console.py +++ b/zhenxun/models/ban_console.py @@ -4,6 +4,7 @@ from typing import ClassVar from typing_extensions import Self from tortoise import fields +from tortoise.expressions import Q from zhenxun.services.data_access import DataAccess from zhenxun.services.db_context import Model @@ -61,16 +62,14 @@ class BanConsole(Model): try: dao = DataAccess(cls) if user_id: - result = ( - await dao.safe_get_or_none(user_id=user_id, group_id=group_id) - if group_id - else await dao.safe_get_or_none( - user_id=user_id, group_id__isnull=True - ) - ) + if group_id: + q = Q(user_id=user_id) & Q(group_id=group_id) + else: + q = Q(user_id=user_id) & Q(group_id__isnull=True) else: - result = await dao.safe_get_or_none(user_id="", group_id=group_id) + q = Q(user_id="") & Q(group_id=group_id) + result = await dao.safe_get_or_none(True, q) future.set_result(result) return result except Exception as e: @@ -128,21 +127,57 @@ class BanConsole(Model): return 0 @classmethod - async def is_ban(cls, user_id: str | None, group_id: str | None = None) -> bool: + async def is_ban( + cls, user_id: str | None, group_id: str | None = None + ) -> list[Self]: """判断用户是否被ban 参数: user_id: 用户id + group_id: 群组id 返回: - bool: 是否被ban + bool: list[Self] | None """ logger.debug("检测是否被ban", target=f"{group_id}:{user_id}") - if await cls.check_ban_time(user_id, group_id): - return True - else: - await cls.unban(user_id, group_id) - return False + + q_conditions = [] + + if user_id and group_id: + q_conditions.append(Q(user_id=user_id, group_id=group_id)) + if user_id: + q_conditions.append(Q(user_id=user_id, group_id__isnull=True)) + if group_id: + q_conditions.append(Q(group_id=group_id, user_id="")) + + if not q_conditions: + return [] + + q = q_conditions[0] + for condition in q_conditions[1:]: + q |= condition + + users = await cls.filter(q).all() + if not users: + return [] + + results = [] + for user in users: + # 永久封禁视为一直处于封禁中 + if user.duration == -1: + results.append(user) + continue + + _time = time.time() - (user.ban_time + user.duration) + # 还在封禁期内 + if _time < 0: + results.append(user) + continue + + # 已过期,删除记录并标记为不满足「全部仍在封禁」条件 + await user.delete() + + return results @classmethod async def ban(