mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-05 03:39:59 +08:00
refactor: enhance ban handling with caching and optimize user/group ban checks
This commit is contained in:
@@ -9,10 +9,11 @@ from nonebot_plugin_uninfo import Uninfo
|
|||||||
from zhenxun.configs.config import Config
|
from zhenxun.configs.config import Config
|
||||||
from zhenxun.models.ban_console import BanConsole
|
from zhenxun.models.ban_console import BanConsole
|
||||||
from zhenxun.models.plugin_info import PluginInfo
|
from zhenxun.models.plugin_info import PluginInfo
|
||||||
|
from zhenxun.services.cache import CacheRoot
|
||||||
from zhenxun.services.data_access import DataAccess
|
from zhenxun.services.data_access import DataAccess
|
||||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||||
from zhenxun.services.log import logger
|
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 zhenxun.utils.utils import EntityIDs, get_entity_ids
|
||||||
|
|
||||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||||
@@ -49,105 +50,6 @@ async def calculate_ban_time(ban_record: BanConsole | None) -> int:
|
|||||||
return 0
|
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:
|
def check_plugin_type(matcher: Matcher) -> bool:
|
||||||
"""判断插件类型是否是隐藏插件
|
"""判断插件类型是否是隐藏插件
|
||||||
|
|
||||||
@@ -190,45 +92,22 @@ def format_time(time_val: float) -> str:
|
|||||||
return time_str
|
return time_str
|
||||||
|
|
||||||
|
|
||||||
async def group_handle(group_id: str) -> None:
|
async def user_handle(
|
||||||
"""群组ban检查
|
plugin: PluginInfo, entity: EntityIDs, session: Uninfo, time_val: int
|
||||||
|
) -> None:
|
||||||
参数:
|
|
||||||
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:
|
|
||||||
"""用户ban检查
|
"""用户ban检查
|
||||||
|
|
||||||
参数:
|
参数:
|
||||||
module: 插件模块名
|
module: 插件模块名
|
||||||
entity: 实体ID信息
|
entity: 实体ID信息
|
||||||
session: Uninfo
|
session: Uninfo
|
||||||
|
time_val: 剩余ban时间
|
||||||
异常:
|
异常:
|
||||||
SkipPluginException: 用户处于黑名单
|
SkipPluginException: 用户处于黑名单
|
||||||
"""
|
"""
|
||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
try:
|
try:
|
||||||
ban_result = Config.get_config("hook", "BAN_RESULT")
|
ban_result = Config.get_config("hook", "BAN_RESULT")
|
||||||
time_val = await is_ban(entity.user_id, entity.group_id)
|
|
||||||
if not time_val:
|
if not time_val:
|
||||||
return
|
return
|
||||||
time_str = format_time(time_val)
|
time_str = format_time(time_val)
|
||||||
@@ -284,24 +163,37 @@ async def auth_ban(
|
|||||||
entity = get_entity_ids(session)
|
entity = get_entity_ids(session)
|
||||||
if entity.user_id in bot.config.superusers:
|
if entity.user_id in bot.config.superusers:
|
||||||
return
|
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:
|
cache_key = f"{entity.user_id}_{entity.group_id}"
|
||||||
try:
|
|
||||||
await asyncio.wait_for(
|
results = await CacheRoot.get(CacheType.BAN, cache_key)
|
||||||
user_handle(plugin, entity, session),
|
if not results:
|
||||||
timeout=DB_TIMEOUT_SECONDS,
|
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:
|
raise SkipPluginException(f"群组: {result.group_id} 处于黑名单中...")
|
||||||
logger.error(f"用户ban检查超时: {entity.user_id}", LOGGER_COMMAND)
|
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:
|
finally:
|
||||||
# 记录总执行时间
|
# 记录总执行时间
|
||||||
elapsed = time.time() - start_time
|
elapsed = time.time() - start_time
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ from typing import ClassVar
|
|||||||
from typing_extensions import Self
|
from typing_extensions import Self
|
||||||
|
|
||||||
from tortoise import fields
|
from tortoise import fields
|
||||||
|
from tortoise.expressions import Q
|
||||||
|
|
||||||
from zhenxun.services.data_access import DataAccess
|
from zhenxun.services.data_access import DataAccess
|
||||||
from zhenxun.services.db_context import Model
|
from zhenxun.services.db_context import Model
|
||||||
@@ -61,16 +62,14 @@ class BanConsole(Model):
|
|||||||
try:
|
try:
|
||||||
dao = DataAccess(cls)
|
dao = DataAccess(cls)
|
||||||
if user_id:
|
if user_id:
|
||||||
result = (
|
if group_id:
|
||||||
await dao.safe_get_or_none(user_id=user_id, group_id=group_id)
|
q = Q(user_id=user_id) & Q(group_id=group_id)
|
||||||
if group_id
|
else:
|
||||||
else await dao.safe_get_or_none(
|
q = Q(user_id=user_id) & Q(group_id__isnull=True)
|
||||||
user_id=user_id, group_id__isnull=True
|
|
||||||
)
|
|
||||||
)
|
|
||||||
else:
|
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)
|
future.set_result(result)
|
||||||
return result
|
return result
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -128,21 +127,57 @@ class BanConsole(Model):
|
|||||||
return 0
|
return 0
|
||||||
|
|
||||||
@classmethod
|
@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
|
"""判断用户是否被ban
|
||||||
|
|
||||||
参数:
|
参数:
|
||||||
user_id: 用户id
|
user_id: 用户id
|
||||||
|
group_id: 群组id
|
||||||
|
|
||||||
返回:
|
返回:
|
||||||
bool: 是否被ban
|
bool: list[Self] | None
|
||||||
"""
|
"""
|
||||||
logger.debug("检测是否被ban", target=f"{group_id}:{user_id}")
|
logger.debug("检测是否被ban", target=f"{group_id}:{user_id}")
|
||||||
if await cls.check_ban_time(user_id, group_id):
|
|
||||||
return True
|
q_conditions = []
|
||||||
else:
|
|
||||||
await cls.unban(user_id, group_id)
|
if user_id and group_id:
|
||||||
return False
|
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
|
@classmethod
|
||||||
async def ban(
|
async def ban(
|
||||||
|
|||||||
Reference in New Issue
Block a user