mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-02 10:10:01 +08:00
🎉🍾♻️refactor(core): 重构核心鉴权与缓存机制,引入异步发送队列以提升性能 🚀 (#2088)
* 添加测试插件 * ✨ 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>
This commit is contained in:
co-authored by
ATTomatoo
pre-commit-ci[bot]
parent
c9f0a8b9d9
commit
837330e30a
@@ -1,5 +1,6 @@
|
||||
import asyncio
|
||||
import random
|
||||
import time
|
||||
|
||||
import nonebot
|
||||
from nonebot import on_notice
|
||||
@@ -10,10 +11,12 @@ from nonebot.plugin import PluginMetadata
|
||||
from nonebot_plugin_alconna import Alconna, Arparma, on_alconna
|
||||
from nonebot_plugin_apscheduler import scheduler
|
||||
from nonebot_plugin_session import EventSession
|
||||
from nonebot_plugin_uninfo import Scene, SceneType, get_interface
|
||||
|
||||
from zhenxun.configs.config import BotConfig
|
||||
from zhenxun.configs.utils import PluginExtraData
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import should_pause_tasks
|
||||
from zhenxun.services.tags import tag_manager
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
@@ -38,6 +41,11 @@ __plugin_meta__ = PluginMetadata(
|
||||
).to_dict(),
|
||||
)
|
||||
|
||||
_FULL_REFRESH_INTERVAL_SECONDS = 24 * 60 * 60
|
||||
|
||||
_GROUP_LAST_UPDATE: dict[tuple[str, str], float] = {}
|
||||
_UPDATE_SEMAPHORE = asyncio.Semaphore(1)
|
||||
|
||||
|
||||
_matcher = on_alconna(
|
||||
Alconna("更新群组成员信息"),
|
||||
@@ -58,6 +66,34 @@ _update_all_matcher = on_alconna(
|
||||
)
|
||||
|
||||
|
||||
def _group_key(bot_id: str, group_id: str) -> tuple[str, str]:
|
||||
return bot_id, group_id
|
||||
|
||||
|
||||
async def _build_scene_map(bot: Bot) -> dict[str, Scene]:
|
||||
if not (interface := get_interface(bot)):
|
||||
return {}
|
||||
scenes = await interface.get_scenes(SceneType.GROUP)
|
||||
return {scene.id: scene for scene in scenes if scene.is_group}
|
||||
|
||||
|
||||
async def _run_update(
|
||||
bot: Bot,
|
||||
group_id: str,
|
||||
*,
|
||||
scene_map: dict[str, Scene] | None = None,
|
||||
platform: str | None = None,
|
||||
force: bool = False,
|
||||
) -> str | None:
|
||||
key = _group_key(bot.self_id, group_id)
|
||||
async with _UPDATE_SEMAPHORE:
|
||||
result = await MemberUpdateManage.update_group_member(
|
||||
bot, group_id, scene_map=scene_map, platform=platform
|
||||
)
|
||||
_GROUP_LAST_UPDATE[key] = time.time()
|
||||
return result
|
||||
|
||||
|
||||
async def _update_all_groups_task(bot: Bot, session: EventSession):
|
||||
"""
|
||||
在后台执行所有群组的更新任务,并向超级用户发送最终报告。
|
||||
@@ -69,21 +105,29 @@ async def _update_all_groups_task(bot: Bot, session: EventSession):
|
||||
|
||||
logger.info(f"Bot {bot_id}: 开始执行所有群组信息更新任务...", "更新所有群组")
|
||||
try:
|
||||
group_list, _ = await PlatformUtils.get_group_list(bot)
|
||||
total_count = len(group_list)
|
||||
for i, group in enumerate(group_list):
|
||||
scene_map = await _build_scene_map(bot)
|
||||
platform = PlatformUtils.get_platform(bot)
|
||||
group_ids = list(scene_map.keys())
|
||||
total_count = len(group_ids)
|
||||
for i, group_id in enumerate(group_ids):
|
||||
try:
|
||||
logger.debug(
|
||||
f"Bot {bot_id}: 正在更新第 {i + 1}/{total_count} 个群组: "
|
||||
f"{group.group_id}",
|
||||
f"{group_id}",
|
||||
"更新所有群组",
|
||||
)
|
||||
await MemberUpdateManage.update_group_member(bot, group.group_id)
|
||||
await _run_update(
|
||||
bot,
|
||||
group_id,
|
||||
scene_map=scene_map,
|
||||
platform=platform,
|
||||
force=True,
|
||||
)
|
||||
success_count += 1
|
||||
except Exception as e:
|
||||
fail_count += 1
|
||||
logger.error(
|
||||
f"Bot {bot_id}: 更新群组 {group.group_id} 信息失败",
|
||||
f"Bot {bot_id}: 更新群组 {group_id} 信息失败",
|
||||
"更新所有群组",
|
||||
e=e,
|
||||
)
|
||||
@@ -118,18 +162,19 @@ async def _(bot: Bot, session: EventSession):
|
||||
|
||||
@_matcher.handle()
|
||||
async def _(bot: Bot, session: EventSession, arparma: Arparma):
|
||||
if gid := session.id3 or session.id2:
|
||||
logger.info("更新群组成员信息", arparma.header_result, session=session)
|
||||
result = await MemberUpdateManage.update_group_member(bot, gid)
|
||||
await MessageUtils.build_message(result).finish(reply_to=True)
|
||||
await tag_manager._invalidate_cache()
|
||||
await MessageUtils.build_message("群组id为空...").send()
|
||||
if not (gid := session.id3 or session.id2):
|
||||
await MessageUtils.build_message("群组id为空...").send()
|
||||
return
|
||||
logger.info("更新群组成员信息", arparma.header_result, session=session)
|
||||
result = await _run_update(bot, gid, force=True)
|
||||
await MessageUtils.build_message(result or "更新已完成").finish(reply_to=True)
|
||||
await tag_manager._invalidate_cache()
|
||||
|
||||
|
||||
@_notice.handle()
|
||||
async def _(bot: Bot, event: GroupIncreaseNoticeEvent):
|
||||
if str(event.user_id) == bot.self_id:
|
||||
await MemberUpdateManage.update_group_member(bot, str(event.group_id))
|
||||
await _run_update(bot, str(event.group_id), force=True)
|
||||
logger.info(
|
||||
f"{BotConfig.self_nickname}加入群聊更新群组信息",
|
||||
"更新群组成员列表",
|
||||
@@ -140,29 +185,50 @@ async def _(bot: Bot, event: GroupIncreaseNoticeEvent):
|
||||
|
||||
|
||||
@scheduler.scheduled_job(
|
||||
"interval",
|
||||
minutes=5,
|
||||
"cron",
|
||||
hour=3,
|
||||
minute=0,
|
||||
max_instances=1,
|
||||
coalesce=True,
|
||||
)
|
||||
async def _():
|
||||
for bot in nonebot.get_bots().values():
|
||||
if PlatformUtils.get_platform(bot) == "qq":
|
||||
try:
|
||||
group_list, _ = await PlatformUtils.get_group_list(bot)
|
||||
if group_list:
|
||||
for group in group_list:
|
||||
try:
|
||||
await MemberUpdateManage.update_group_member(
|
||||
bot, group.group_id
|
||||
)
|
||||
logger.debug("自动更新群组成员信息成功...")
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Bot: {bot.self_id} 自动更新群组成员信息失败",
|
||||
target=group.group_id,
|
||||
e=e,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Bot: {bot.self_id} 自动更新群组信息", e=e)
|
||||
logger.debug(f"自动 Bot: {bot.self_id} 更新群组成员信息成功...")
|
||||
|
||||
await tag_manager._invalidate_cache()
|
||||
async def _nightly_full_refresh():
|
||||
if should_pause_tasks():
|
||||
return
|
||||
now = time.time()
|
||||
bots = nonebot.get_bots()
|
||||
if not bots:
|
||||
return
|
||||
updated = 0
|
||||
for bot in bots.values():
|
||||
platform = PlatformUtils.get_platform(bot)
|
||||
if platform != "qq":
|
||||
continue
|
||||
try:
|
||||
scene_map = await _build_scene_map(bot)
|
||||
if not scene_map:
|
||||
continue
|
||||
for group_id in scene_map:
|
||||
key = _group_key(bot.self_id, group_id)
|
||||
last_update = _GROUP_LAST_UPDATE.get(key, 0)
|
||||
if now - last_update < _FULL_REFRESH_INTERVAL_SECONDS:
|
||||
continue
|
||||
try:
|
||||
result = await _run_update(
|
||||
bot,
|
||||
group_id,
|
||||
scene_map=scene_map,
|
||||
platform=platform,
|
||||
force=True,
|
||||
)
|
||||
if result is not None:
|
||||
updated += 1
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Bot: {bot.self_id} 夜间更新群组成员信息失败",
|
||||
target=group_id,
|
||||
e=e,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Bot: {bot.self_id} 夜间更新群组信息", e=e)
|
||||
if updated:
|
||||
await tag_manager._invalidate_cache()
|
||||
|
||||
@@ -3,7 +3,7 @@ import re
|
||||
|
||||
import nonebot
|
||||
from nonebot.adapters import Bot
|
||||
from nonebot_plugin_uninfo import Member, SceneType, get_interface
|
||||
from nonebot_plugin_uninfo import Member, Scene, SceneType, get_interface
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
@@ -18,10 +18,13 @@ class MemberUpdateManage:
|
||||
async def __handle_user(
|
||||
cls,
|
||||
member: Member,
|
||||
db_user: list[GroupInfoUser],
|
||||
db_user_map: dict[str, list[GroupInfoUser]],
|
||||
group_id: str,
|
||||
data_list: tuple[list, list, list],
|
||||
data_list: tuple[list[GroupInfoUser], list[GroupInfoUser], list[int]],
|
||||
platform: str | None,
|
||||
*,
|
||||
default_auth: int | None,
|
||||
superusers: set[str],
|
||||
):
|
||||
"""单个成员操作
|
||||
|
||||
@@ -32,37 +35,32 @@ class MemberUpdateManage:
|
||||
data_list: 数据列表
|
||||
platform: 平台
|
||||
"""
|
||||
driver = nonebot.get_driver()
|
||||
default_auth = Config.get_config("admin_bot_manage", "ADMIN_DEFAULT_AUTH")
|
||||
nickname = re.sub(
|
||||
r"[\x00-\x09\x0b-\x1f\x7f-\x9f]", "", member.nick or member.user.name or ""
|
||||
)
|
||||
role = member.role
|
||||
db_user_uid = [u.user_id for u in db_user]
|
||||
uid2name = {u.user_id: u.user_name for u in db_user}
|
||||
if member.id in driver.config.superusers:
|
||||
await LevelUser.set_level(member.id, group_id, 9)
|
||||
member_id = str(member.id)
|
||||
if member_id in superusers:
|
||||
await LevelUser.set_level(member_id, group_id, 9)
|
||||
elif role and default_auth:
|
||||
if role.id != "MEMBER" and not await LevelUser.is_group_flag(
|
||||
member.id, group_id
|
||||
member_id, group_id
|
||||
):
|
||||
if role.id == "OWNER":
|
||||
await LevelUser.set_level(member.id, group_id, default_auth + 1)
|
||||
await LevelUser.set_level(member_id, group_id, default_auth + 1)
|
||||
elif role.id == "ADMINISTRATOR":
|
||||
await LevelUser.set_level(member.id, group_id, default_auth)
|
||||
if cnt := db_user_uid.count(member.id):
|
||||
users = [u for u in db_user if u.user_id == member.id]
|
||||
if cnt > 1:
|
||||
for u in users[1:]:
|
||||
data_list[2].append(u.id)
|
||||
if nickname != uid2name.get(member.id):
|
||||
await LevelUser.set_level(member_id, group_id, default_auth)
|
||||
if users := db_user_map.get(member_id):
|
||||
if len(users) > 1:
|
||||
data_list[2].extend(u.id for u in users[1:])
|
||||
if nickname != users[0].user_name:
|
||||
user = users[0]
|
||||
user.user_name = nickname
|
||||
data_list[1].append(user)
|
||||
else:
|
||||
data_list[0].append(
|
||||
GroupInfoUser(
|
||||
user_id=member.id,
|
||||
user_id=member_id,
|
||||
group_id=group_id,
|
||||
user_name=nickname,
|
||||
user_join_time=member.joined_at or datetime.now(),
|
||||
@@ -71,7 +69,14 @@ class MemberUpdateManage:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def update_group_member(cls, bot: Bot, group_id: str) -> str:
|
||||
async def update_group_member(
|
||||
cls,
|
||||
bot: Bot,
|
||||
group_id: str,
|
||||
*,
|
||||
scene_map: dict[str, Scene] | None = None,
|
||||
platform: str | None = None,
|
||||
) -> str:
|
||||
"""更新群组成员信息
|
||||
|
||||
参数:
|
||||
@@ -85,23 +90,26 @@ class MemberUpdateManage:
|
||||
logger.warning(f"bot: {bot.self_id},group_id为空,无法更新群成员信息...")
|
||||
return "群组id为空..."
|
||||
if interface := get_interface(bot):
|
||||
scenes = await interface.get_scenes()
|
||||
platform = PlatformUtils.get_platform(bot)
|
||||
group_list = [s for s in scenes if s.is_group and s.id == group_id]
|
||||
if not group_list:
|
||||
if scene_map is None:
|
||||
scenes = await interface.get_scenes(SceneType.GROUP)
|
||||
scene_map = {scene.id: scene for scene in scenes if scene.is_group}
|
||||
if platform is None:
|
||||
platform = PlatformUtils.get_platform(bot)
|
||||
group_scene = scene_map.get(group_id) if scene_map else None
|
||||
if not group_scene:
|
||||
logger.warning(
|
||||
f"bot: {bot.self_id},group_id: {group_id},群组不存在,"
|
||||
"无法更新群成员信息..."
|
||||
)
|
||||
return "更新群组失败,群组不存在..."
|
||||
members = await interface.get_members(SceneType.GROUP, group_list[0].id)
|
||||
members = await interface.get_members(SceneType.GROUP, group_scene.id)
|
||||
|
||||
try:
|
||||
group_console, _ = await GroupConsole.get_or_create(
|
||||
group_id=group_id, defaults={"platform": platform}
|
||||
)
|
||||
group_console.member_count = len(members)
|
||||
group_console.group_name = group_list[0].name or ""
|
||||
group_console.group_name = group_scene.name or ""
|
||||
await group_console.save(update_fields=["member_count", "group_name"])
|
||||
logger.debug(
|
||||
f"已更新群组 {group_id} 的成员总数为 {len(members)}",
|
||||
@@ -115,13 +123,31 @@ class MemberUpdateManage:
|
||||
)
|
||||
|
||||
db_user = await GroupInfoUser.filter(group_id=group_id).all()
|
||||
db_user_uid = [u.user_id for u in db_user]
|
||||
data_list = ([], [], [])
|
||||
exist_member_list = []
|
||||
db_user_map: dict[str, list[GroupInfoUser]] = {}
|
||||
for user in db_user:
|
||||
db_user_map.setdefault(user.user_id, []).append(user)
|
||||
db_user_ids = set(db_user_map)
|
||||
data_list: tuple[list[GroupInfoUser], list[GroupInfoUser], list[int]] = (
|
||||
[],
|
||||
[],
|
||||
[],
|
||||
)
|
||||
exist_member_ids: set[str] = set()
|
||||
driver = nonebot.get_driver()
|
||||
superusers = set(driver.config.superusers)
|
||||
default_auth = Config.get_config("admin_bot_manage", "ADMIN_DEFAULT_AUTH")
|
||||
for member in members:
|
||||
logger.debug(f"即将更新群组成员: {member}", "更新群组成员信息")
|
||||
await cls.__handle_user(member, db_user, group_id, data_list, platform)
|
||||
exist_member_list.append(member.id)
|
||||
member_id = str(member.id)
|
||||
await cls.__handle_user(
|
||||
member,
|
||||
db_user_map,
|
||||
group_id,
|
||||
data_list,
|
||||
platform,
|
||||
default_auth=default_auth,
|
||||
superusers=superusers,
|
||||
)
|
||||
exist_member_ids.add(member_id)
|
||||
if data_list[0]:
|
||||
try:
|
||||
await GroupInfoUser.bulk_create(
|
||||
@@ -145,14 +171,12 @@ class MemberUpdateManage:
|
||||
await GroupInfoUser.filter(id__in=data_list[2]).delete()
|
||||
logger.debug(f"删除重复数据 Ids: {data_list[2]}", "更新群组成员信息")
|
||||
|
||||
if delete_member_list := [
|
||||
uid for uid in db_user_uid if uid not in exist_member_list
|
||||
]:
|
||||
if delete_member_ids := db_user_ids - exist_member_ids:
|
||||
await GroupInfoUser.filter(
|
||||
user_id__in=delete_member_list, group_id=group_id
|
||||
user_id__in=list(delete_member_ids), group_id=group_id
|
||||
).delete()
|
||||
logger.info(
|
||||
f"删除已退群用户 {len(delete_member_list)} 条",
|
||||
f"删除已退群用户 {len(delete_member_ids)} 条",
|
||||
"更新群组成员信息",
|
||||
group_id=group_id,
|
||||
platform="qq",
|
||||
|
||||
@@ -100,7 +100,7 @@ async def build_task(group_id: str | None) -> BuildImage:
|
||||
column_name = ["ID", "模块", "名称", "群组状态", "全局状态", "运行时间"]
|
||||
group = None
|
||||
if group_id:
|
||||
group = await GroupConsole.get_group(group_id=group_id)
|
||||
group = await GroupConsole.get_group_db(group_id=group_id)
|
||||
if not group:
|
||||
raise GroupInfoNotFound()
|
||||
else:
|
||||
@@ -182,7 +182,7 @@ class PluginManager:
|
||||
)
|
||||
return f"成功将所有功能进群默认状态修改为: {'开启' if status else '关闭'}"
|
||||
if group_id:
|
||||
if group := await GroupConsole.get_group(group_id=group_id):
|
||||
if group := await GroupConsole.get_group_db(group_id=group_id):
|
||||
module_list = cast(
|
||||
list[str],
|
||||
await PluginInfo.filter(plugin_type=PluginType.NORMAL).values_list(
|
||||
@@ -214,7 +214,7 @@ class PluginManager:
|
||||
返回:
|
||||
bool: 是否醒来
|
||||
"""
|
||||
if c := await GroupConsole.get_group(group_id=group_id):
|
||||
if c := await GroupConsole.get_group_db(group_id=group_id):
|
||||
return c.status
|
||||
return False
|
||||
|
||||
|
||||
@@ -1,4 +1,8 @@
|
||||
from nonebot import on_message
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from nonebot import get_driver, on_message
|
||||
from nonebot.adapters import Event
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot_plugin_alconna import UniMsg
|
||||
from nonebot_plugin_apscheduler import scheduler
|
||||
@@ -8,6 +12,7 @@ from zhenxun.configs.config import Config
|
||||
from zhenxun.configs.utils import PluginExtraData, RegisterConfig
|
||||
from zhenxun.models.chat_history import ChatHistory
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import is_overloaded, should_pause_tasks
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.utils import get_entity_ids
|
||||
|
||||
@@ -33,28 +38,83 @@ __plugin_meta__ = PluginMetadata(
|
||||
)
|
||||
|
||||
|
||||
def rule(message: UniMsg) -> bool:
|
||||
return bool(Config.get_config("chat_history", "FLAG") and message)
|
||||
_COMMAND_STARTS = {str(item) for item in (get_driver().config.command_start or [])}
|
||||
_LAST_GROUP_SAVE: dict[str, float] = {}
|
||||
_LAST_USER_SAVE: dict[str, float] = {}
|
||||
_GROUP_MIN_INTERVAL = 0.5
|
||||
_USER_MIN_INTERVAL = 0.2
|
||||
|
||||
|
||||
def _is_command_like(text: str) -> bool:
|
||||
if not text:
|
||||
return False
|
||||
for start in _COMMAND_STARTS:
|
||||
if text.startswith(start):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
async def rule(event: Event, message: UniMsg, session: Uninfo) -> bool:
|
||||
if is_overloaded():
|
||||
return False
|
||||
if not Config.get_config("chat_history", "FLAG"):
|
||||
return False
|
||||
if not message:
|
||||
return False
|
||||
text = message.extract_plain_text().strip()
|
||||
if _is_command_like(text):
|
||||
return False
|
||||
entity = get_entity_ids(session)
|
||||
now = time.time()
|
||||
if entity.group_id:
|
||||
last_group = _LAST_GROUP_SAVE.get(entity.group_id, 0)
|
||||
if now - last_group < _GROUP_MIN_INTERVAL:
|
||||
return False
|
||||
if entity.user_id:
|
||||
last_user = _LAST_USER_SAVE.get(entity.user_id, 0)
|
||||
if now - last_user < _USER_MIN_INTERVAL:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
chat_history = on_message(rule=rule, priority=1, block=False)
|
||||
|
||||
TEMP_LIST = []
|
||||
_HISTORY_QUEUE: asyncio.Queue[ChatHistory] = asyncio.Queue(maxsize=5000)
|
||||
_DROP_COUNT = 0
|
||||
_LAST_DROP_LOG = 0.0
|
||||
_DROP_LOG_INTERVAL = 10.0
|
||||
|
||||
|
||||
@chat_history.handle()
|
||||
async def _(message: UniMsg, session: Uninfo):
|
||||
entity = get_entity_ids(session)
|
||||
TEMP_LIST.append(
|
||||
ChatHistory(
|
||||
user_id=entity.user_id,
|
||||
group_id=entity.group_id,
|
||||
text=str(message),
|
||||
plain_text=message.extract_plain_text(),
|
||||
bot_id=session.self_id,
|
||||
platform=session.platform,
|
||||
now = time.time()
|
||||
if entity.group_id:
|
||||
_LAST_GROUP_SAVE[entity.group_id] = now
|
||||
if entity.user_id:
|
||||
_LAST_USER_SAVE[entity.user_id] = now
|
||||
if is_overloaded():
|
||||
return
|
||||
try:
|
||||
_HISTORY_QUEUE.put_nowait(
|
||||
ChatHistory(
|
||||
user_id=entity.user_id,
|
||||
group_id=entity.group_id,
|
||||
text=str(message),
|
||||
plain_text=message.extract_plain_text(),
|
||||
bot_id=session.self_id,
|
||||
platform=session.platform,
|
||||
)
|
||||
)
|
||||
)
|
||||
except asyncio.QueueFull:
|
||||
global _DROP_COUNT, _LAST_DROP_LOG
|
||||
_DROP_COUNT += 1
|
||||
if now - _LAST_DROP_LOG > _DROP_LOG_INTERVAL:
|
||||
_LAST_DROP_LOG = now
|
||||
logger.debug(
|
||||
f"chat_history queue full, dropped {_DROP_COUNT} items",
|
||||
"chat_history",
|
||||
)
|
||||
|
||||
|
||||
@scheduler.scheduled_job(
|
||||
@@ -63,8 +123,14 @@ async def _(message: UniMsg, session: Uninfo):
|
||||
)
|
||||
async def _():
|
||||
try:
|
||||
message_list = TEMP_LIST.copy()
|
||||
TEMP_LIST.clear()
|
||||
if should_pause_tasks():
|
||||
return
|
||||
message_list: list[ChatHistory] = []
|
||||
while True:
|
||||
try:
|
||||
message_list.append(_HISTORY_QUEUE.get_nowait())
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
if message_list:
|
||||
await ChatHistory.bulk_create(message_list)
|
||||
logger.debug(f"批量添加聊天记录 {len(message_list)} 条", "定时任务")
|
||||
|
||||
@@ -1,82 +0,0 @@
|
||||
import os
|
||||
import random
|
||||
|
||||
from nonebot import on_message
|
||||
from nonebot.adapters import Event
|
||||
from nonebot.matcher import Matcher
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot_plugin_alconna import UniMsg
|
||||
from nonebot_plugin_session import EventSession
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.configs.path_config import IMAGE_PATH
|
||||
from zhenxun.configs.utils import PluginExtraData
|
||||
from zhenxun.models.ban_console import BanConsole
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
|
||||
__plugin_meta__ = PluginMetadata(
|
||||
name="笨蛋检测",
|
||||
description="功能名称当命令检测",
|
||||
usage="""当一些笨蛋直接输入功能名称时,提示笨蛋使用帮助指令查看功能帮助""".strip(),
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
version="0.1",
|
||||
plugin_type=PluginType.DEPENDANT,
|
||||
menu_type="其他",
|
||||
).to_dict(),
|
||||
)
|
||||
|
||||
|
||||
async def rule(event: Event, message: UniMsg, session: Uninfo) -> bool:
|
||||
group_id = session.group.id if session.group else None
|
||||
text = message.extract_plain_text().strip()
|
||||
if await BanConsole.is_ban(session.user.id, group_id):
|
||||
return False
|
||||
if group_id:
|
||||
if await BanConsole.is_ban(None, group_id):
|
||||
return False
|
||||
if g := await GroupConsole.get_group(group_id):
|
||||
if g.level < 0:
|
||||
return False
|
||||
return event.is_tome() and bool(text and len(text) < 20)
|
||||
|
||||
|
||||
_matcher = on_message(rule=rule, priority=996, block=False)
|
||||
|
||||
|
||||
_path = IMAGE_PATH / "_base" / "laugh"
|
||||
|
||||
|
||||
@_matcher.handle()
|
||||
async def _(matcher: Matcher, message: UniMsg, session: EventSession):
|
||||
text = message.extract_plain_text().strip()
|
||||
plugin = await PluginInfo.get_or_none(
|
||||
name=text,
|
||||
load_status=True,
|
||||
plugin_type=PluginType.NORMAL,
|
||||
block_type__isnull=True,
|
||||
status=True,
|
||||
)
|
||||
|
||||
if not plugin:
|
||||
return
|
||||
|
||||
image = None
|
||||
if _path.exists():
|
||||
if files := os.listdir(_path):
|
||||
image = _path / random.choice(files)
|
||||
message_list = []
|
||||
if image:
|
||||
message_list.append(image)
|
||||
message_list.append(
|
||||
"桀桀桀,预判到会有 '笨蛋' 把功能名称当命令用,特地前来嘲笑!"
|
||||
f"但还是好心来帮帮你啦!\n请at我发送 '帮助 {plugin.name}' 或者"
|
||||
f" '帮助 {plugin.id}' 来获取该功能帮助!"
|
||||
)
|
||||
logger.info("检测到功能名称当命令使用,已发送帮助信息", "功能帮助", session=session)
|
||||
await MessageUtils.build_message(message_list).send(reply_to=True)
|
||||
matcher.stop_propagation()
|
||||
@@ -1,4 +1,3 @@
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from nonebot_plugin_alconna import At
|
||||
@@ -6,8 +5,7 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.models.level_user import LevelUser
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.cache.runtime_cache import LevelUserMemoryCache, LevelUserSnapshot
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.utils import get_entity_ids
|
||||
|
||||
@@ -16,7 +14,14 @@ from .exception import SkipPluginException
|
||||
from .utils import send_message
|
||||
|
||||
|
||||
async def auth_admin(plugin: PluginInfo, session: Uninfo):
|
||||
async def auth_admin(
|
||||
plugin: PluginInfo,
|
||||
session: Uninfo,
|
||||
cached_levels: tuple[
|
||||
LevelUser | LevelUserSnapshot | None, LevelUser | LevelUserSnapshot | None
|
||||
]
|
||||
| None = None,
|
||||
):
|
||||
"""管理员命令 个人权限
|
||||
|
||||
参数:
|
||||
@@ -30,37 +35,17 @@ async def auth_admin(plugin: PluginInfo, session: Uninfo):
|
||||
|
||||
try:
|
||||
entity = get_entity_ids(session)
|
||||
level_dao = DataAccess(LevelUser)
|
||||
|
||||
# 并行查询用户权限数据
|
||||
global_user: LevelUser | None = None
|
||||
group_users: LevelUser | None = None
|
||||
global_user: LevelUser | LevelUserSnapshot | None = None
|
||||
group_users: LevelUser | LevelUserSnapshot | None = None
|
||||
|
||||
# 查询全局权限
|
||||
global_user_task = level_dao.safe_get_or_none(
|
||||
user_id=session.user.id, group_id__isnull=True
|
||||
)
|
||||
|
||||
# 如果在群组中,查询群组权限
|
||||
group_users_task = None
|
||||
if entity.group_id:
|
||||
group_users_task = level_dao.safe_get_or_none(
|
||||
user_id=session.user.id, group_id=entity.group_id
|
||||
if cached_levels is not None:
|
||||
global_user, group_users = cached_levels
|
||||
else:
|
||||
global_user, group_users = await LevelUserMemoryCache.get_levels(
|
||||
session.user.id, entity.group_id
|
||||
)
|
||||
|
||||
# 等待查询完成,添加超时控制
|
||||
try:
|
||||
results = await asyncio.wait_for(
|
||||
asyncio.gather(global_user_task, group_users_task or asyncio.sleep(0)),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
global_user = results[0]
|
||||
group_users = results[1] if group_users_task else None
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"查询用户权限超时: user_id={session.user.id}", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
return
|
||||
|
||||
user_level = global_user.user_level if global_user else 0
|
||||
if entity.group_id and group_users:
|
||||
user_level = max(user_level, group_users.user_level)
|
||||
@@ -73,6 +58,7 @@ async def auth_admin(plugin: PluginInfo, session: Uninfo):
|
||||
f"你的权限不足喔,该功能需要的权限等级: {plugin.admin_level}",
|
||||
],
|
||||
entity.user_id,
|
||||
background=True,
|
||||
)
|
||||
|
||||
raise SkipPluginException(
|
||||
@@ -83,6 +69,7 @@ async def auth_admin(plugin: PluginInfo, session: Uninfo):
|
||||
await send_message(
|
||||
session,
|
||||
f"你的权限不足喔,该功能需要的权限等级: {plugin.admin_level}",
|
||||
background=True,
|
||||
)
|
||||
|
||||
raise SkipPluginException(
|
||||
|
||||
@@ -9,7 +9,8 @@ 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.data_access import DataAccess
|
||||
from zhenxun.services.cache.cache_containers import CacheDict
|
||||
from zhenxun.services.cache.runtime_cache import BanMemoryCache
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import PluginType
|
||||
@@ -25,6 +26,76 @@ Config.add_plugin_config(
|
||||
"才不会给你发消息.",
|
||||
help="对被ban用户发送的消息",
|
||||
)
|
||||
Config.add_plugin_config(
|
||||
"hook",
|
||||
"BAN_CACHE_TTL",
|
||||
2,
|
||||
help="ban cache ttl seconds",
|
||||
)
|
||||
Config.add_plugin_config(
|
||||
"hook",
|
||||
"BAN_CACHE_TTL_POSITIVE",
|
||||
30,
|
||||
help="ban cache ttl seconds for banned users",
|
||||
)
|
||||
Config.add_plugin_config(
|
||||
"hook",
|
||||
"BAN_CACHE_TTL_NEGATIVE",
|
||||
5,
|
||||
help="ban cache ttl seconds for non-banned users",
|
||||
)
|
||||
|
||||
|
||||
def _coerce_ttl(value, default):
|
||||
try:
|
||||
value_int = int(value)
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
return value_int if value_int >= 0 else default
|
||||
|
||||
|
||||
_ban_cache_ttl_value = Config.get_config("hook", "BAN_CACHE_TTL", 2)
|
||||
try:
|
||||
_ban_cache_ttl_value = int(_ban_cache_ttl_value)
|
||||
except (TypeError, ValueError):
|
||||
_ban_cache_ttl_value = 2
|
||||
|
||||
_ban_cache_ttl_positive = _coerce_ttl(
|
||||
Config.get_config("hook", "BAN_CACHE_TTL_POSITIVE", _ban_cache_ttl_value),
|
||||
_ban_cache_ttl_value,
|
||||
)
|
||||
_ban_cache_ttl_negative = _coerce_ttl(
|
||||
Config.get_config("hook", "BAN_CACHE_TTL_NEGATIVE", _ban_cache_ttl_value),
|
||||
_ban_cache_ttl_value,
|
||||
)
|
||||
|
||||
BAN_CACHE = (
|
||||
CacheDict("AUTH_BAN_CACHE", expire=0)
|
||||
if max(_ban_cache_ttl_positive, _ban_cache_ttl_negative) > 0
|
||||
else None
|
||||
)
|
||||
|
||||
|
||||
def _ban_cache_key(user_id: str | None, group_id: str | None) -> str:
|
||||
return f"{user_id or ''}:{group_id or ''}"
|
||||
|
||||
|
||||
def _ban_cache_get(key: str) -> int | None:
|
||||
if not BAN_CACHE:
|
||||
return None
|
||||
try:
|
||||
return BAN_CACHE[key]
|
||||
except KeyError:
|
||||
return None
|
||||
|
||||
|
||||
def _ban_cache_set(key: str, value: int) -> None:
|
||||
if not BAN_CACHE:
|
||||
return
|
||||
ttl = _ban_cache_ttl_positive if value else _ban_cache_ttl_negative
|
||||
if ttl <= 0:
|
||||
return
|
||||
BAN_CACHE.set(key, value, expire=ttl)
|
||||
|
||||
|
||||
async def calculate_ban_time(ban_record: BanConsole | None) -> int:
|
||||
@@ -57,79 +128,13 @@ async def is_ban(user_id: str | None, group_id: str | None) -> int:
|
||||
group_id: 群组ID
|
||||
|
||||
返回:
|
||||
int: ban的剩余时间,0表示未被ban
|
||||
int: ban剩余时长,-1时为永久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:
|
||||
ban_records = await asyncio.wait_for(
|
||||
asyncio.gather(*tasks), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
if len(tasks) == 2:
|
||||
group_user, user = ban_records
|
||||
elif user_id and group_id:
|
||||
group_user = ban_records[0]
|
||||
else:
|
||||
user = ban_records[0]
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(
|
||||
f"查询ban记录超时: 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,
|
||||
)
|
||||
if not BanMemoryCache.is_loaded():
|
||||
return 0
|
||||
return BanMemoryCache.remaining_time(user_id, group_id)
|
||||
|
||||
|
||||
def check_plugin_type(matcher: Matcher) -> bool:
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from zhenxun.models.bot_console import BotConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.cache.runtime_cache import BotMemoryCache, BotSnapshot
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
|
||||
@@ -12,7 +10,12 @@ from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .exception import SkipPluginException
|
||||
|
||||
|
||||
async def auth_bot(plugin: PluginInfo, bot_id: str):
|
||||
async def auth_bot(
|
||||
plugin: PluginInfo,
|
||||
bot_id: str,
|
||||
bot_data: BotConsole | BotSnapshot | None = None,
|
||||
skip_fetch: bool = False,
|
||||
):
|
||||
"""bot层面的权限检查
|
||||
|
||||
参数:
|
||||
@@ -26,17 +29,9 @@ async def auth_bot(plugin: PluginInfo, bot_id: str):
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
# 从数据库或缓存中获取 bot 信息
|
||||
bot_dao = DataAccess(BotConsole)
|
||||
|
||||
try:
|
||||
bot: BotConsole | None = await asyncio.wait_for(
|
||||
bot_dao.safe_get_or_none(bot_id=bot_id), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"查询Bot信息超时: bot_id={bot_id}", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
return
|
||||
bot: BotConsole | BotSnapshot | None = bot_data
|
||||
if bot is None and not skip_fetch:
|
||||
bot = await BotMemoryCache.get(bot_id)
|
||||
|
||||
if not bot or not bot.status:
|
||||
raise SkipPluginException("Bot不存在或休眠中阻断权限检测...")
|
||||
|
||||
@@ -10,12 +10,16 @@ from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .exception import SkipPluginException
|
||||
from .utils import send_message
|
||||
|
||||
DEFAULT_GOLD = 100
|
||||
|
||||
async def auth_cost(user: UserConsole, plugin: PluginInfo, session: Uninfo) -> int:
|
||||
|
||||
async def auth_cost(
|
||||
user: UserConsole | None, plugin: PluginInfo, session: Uninfo
|
||||
) -> int:
|
||||
"""检测是否满足金币条件
|
||||
|
||||
参数:
|
||||
user: UserConsole
|
||||
user: UserConsole | None
|
||||
plugin: PluginInfo
|
||||
session: Uninfo
|
||||
|
||||
@@ -25,7 +29,8 @@ async def auth_cost(user: UserConsole, plugin: PluginInfo, session: Uninfo) -> i
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
if user.gold < plugin.cost_gold:
|
||||
user_gold = user.gold if user else DEFAULT_GOLD
|
||||
if user_gold < plugin.cost_gold:
|
||||
"""插件消耗金币不足"""
|
||||
await send_message(session, f"金币不足..该功能需要{plugin.cost_gold}金币..")
|
||||
raise SkipPluginException(f"{plugin.name}({plugin.module}) 金币限制...")
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
import time
|
||||
|
||||
from nonebot_plugin_alconna import UniMsg
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.cache.runtime_cache import GroupSnapshot
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
|
||||
@@ -12,8 +11,8 @@ from .exception import SkipPluginException
|
||||
|
||||
async def auth_group(
|
||||
plugin: PluginInfo,
|
||||
group: GroupConsole | None,
|
||||
message: UniMsg,
|
||||
group: GroupConsole | GroupSnapshot | None,
|
||||
text: str | None,
|
||||
group_id: str | None,
|
||||
):
|
||||
"""群黑名单检测 群总开关检测
|
||||
@@ -29,7 +28,7 @@ async def auth_group(
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
text = message.extract_plain_text()
|
||||
text = text or ""
|
||||
|
||||
if not group:
|
||||
raise SkipPluginException("群组信息不存在...")
|
||||
|
||||
@@ -1,13 +1,18 @@
|
||||
import asyncio
|
||||
import time
|
||||
from typing import ClassVar
|
||||
from typing import Any, ClassVar
|
||||
|
||||
import nonebot
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
from pydantic import BaseModel
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.plugin_limit import PluginLimit
|
||||
from zhenxun.services.cache.runtime_cache import (
|
||||
PluginLimitMemoryCache,
|
||||
PluginLimitSnapshot,
|
||||
)
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import LimitWatchType, PluginLimitType
|
||||
@@ -22,6 +27,16 @@ from .exception import SkipPluginException
|
||||
|
||||
driver = nonebot.get_driver()
|
||||
|
||||
Config.add_plugin_config(
|
||||
"hook",
|
||||
"AUTH_LIMIT_NOTICE_CD",
|
||||
2,
|
||||
help="auth limit notice cooldown seconds",
|
||||
)
|
||||
_LIMIT_NOTICE_CD = int(Config.get_config("hook", "AUTH_LIMIT_NOTICE_CD", 2) or 2)
|
||||
_LIMIT_NOTICE_LIMITER = FreqLimiter(_LIMIT_NOTICE_CD)
|
||||
_LIMIT_NOTICE_TASKS: set[asyncio.Task] = set()
|
||||
|
||||
|
||||
@PriorityLifecycle.on_startup(priority=5)
|
||||
async def _():
|
||||
@@ -30,13 +45,41 @@ async def _():
|
||||
|
||||
|
||||
class Limit(BaseModel):
|
||||
limit: PluginLimit
|
||||
limit: PluginLimit | PluginLimitSnapshot
|
||||
limiter: FreqLimiter | UserBlockLimiter | CountLimiter
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
|
||||
|
||||
def _limit_notice_key(
|
||||
limit: PluginLimit | PluginLimitSnapshot,
|
||||
user_id: str,
|
||||
group_id: str | None,
|
||||
channel_id: str | None,
|
||||
) -> str:
|
||||
key = user_id
|
||||
if group_id and limit.watch_type == LimitWatchType.GROUP:
|
||||
key = channel_id or group_id
|
||||
return f"{limit.module}:{limit.limit_type}:{key}"
|
||||
|
||||
|
||||
def _send_limit_notice(message: str, format_kwargs: dict[str, Any], key: str) -> None:
|
||||
if not _LIMIT_NOTICE_LIMITER.check(key):
|
||||
return
|
||||
_LIMIT_NOTICE_LIMITER.start_cd(key)
|
||||
|
||||
async def _send():
|
||||
try:
|
||||
await MessageUtils.build_message(message, format_args=format_kwargs).send()
|
||||
except Exception as exc:
|
||||
logger.error("limit notice send failed", LOGGER_COMMAND, e=exc)
|
||||
|
||||
task = asyncio.create_task(_send())
|
||||
_LIMIT_NOTICE_TASKS.add(task)
|
||||
task.add_done_callback(_LIMIT_NOTICE_TASKS.discard)
|
||||
|
||||
|
||||
class LimitManager:
|
||||
add_module: ClassVar[list] = []
|
||||
last_update_time: ClassVar[float] = 0
|
||||
@@ -48,8 +91,11 @@ class LimitManager:
|
||||
count_limit: ClassVar[dict[str, Limit]] = {}
|
||||
|
||||
# 模块限制缓存,避免频繁查询数据库
|
||||
module_limit_cache: ClassVar[dict[str, tuple[float, list[PluginLimit]]]] = {}
|
||||
module_limit_cache: ClassVar[
|
||||
dict[str, tuple[float, list[PluginLimitSnapshot], bool]]
|
||||
] = {}
|
||||
module_cache_ttl: ClassVar[float] = 60 # 模块缓存有效期(秒)
|
||||
module_cache_error_ttl: ClassVar[float] = 5 # 超时缓存有效期(秒)
|
||||
|
||||
@classmethod
|
||||
async def init_limit(cls):
|
||||
@@ -70,14 +116,8 @@ class LimitManager:
|
||||
cls.is_updating = True
|
||||
try:
|
||||
start_time = time.time()
|
||||
try:
|
||||
limit_list = await asyncio.wait_for(
|
||||
PluginLimit.filter(status=True).all(), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("查询限制信息超时", LOGGER_COMMAND)
|
||||
cls.is_updating = False
|
||||
return
|
||||
await PluginLimitMemoryCache.ensure_loaded()
|
||||
limit_list = PluginLimitMemoryCache.get_all_limits()
|
||||
|
||||
# 清空旧数据
|
||||
cls.add_module = []
|
||||
@@ -96,7 +136,7 @@ class LimitManager:
|
||||
cls.is_updating = False
|
||||
|
||||
@classmethod
|
||||
def add_limit(cls, limit: PluginLimit):
|
||||
def add_limit(cls, limit: PluginLimit | PluginLimitSnapshot):
|
||||
"""添加限制
|
||||
|
||||
参数:
|
||||
@@ -109,12 +149,16 @@ class LimitManager:
|
||||
limit=limit, limiter=UserBlockLimiter()
|
||||
)
|
||||
elif limit.limit_type == PluginLimitType.CD:
|
||||
cd_value = int(limit.cd or 0)
|
||||
cls.cd_limit[limit.module] = Limit(
|
||||
limit=limit, limiter=FreqLimiter(limit.cd)
|
||||
limit=limit, limiter=FreqLimiter(cd_value)
|
||||
)
|
||||
elif limit.limit_type == PluginLimitType.COUNT:
|
||||
max_count = int(limit.max_count or 0)
|
||||
if max_count <= 0:
|
||||
return
|
||||
cls.count_limit[limit.module] = Limit(
|
||||
limit=limit, limiter=CountLimiter(limit.max_count)
|
||||
limit=limit, limiter=CountLimiter(max_count)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -144,7 +188,7 @@ class LimitManager:
|
||||
limiter.set_false(key_type)
|
||||
|
||||
@classmethod
|
||||
async def get_module_limits(cls, module: str) -> list[PluginLimit]:
|
||||
async def get_module_limits(cls, module: str) -> list[PluginLimitSnapshot]:
|
||||
"""获取模块的限制信息,使用缓存减少数据库查询
|
||||
|
||||
参数:
|
||||
@@ -157,30 +201,20 @@ class LimitManager:
|
||||
|
||||
# 检查缓存
|
||||
if module in cls.module_limit_cache:
|
||||
cache_time, limits = cls.module_limit_cache[module]
|
||||
if current_time - cache_time < cls.module_cache_ttl:
|
||||
cache_time, limits, is_error = cls.module_limit_cache[module]
|
||||
ttl = cls.module_cache_error_ttl if is_error else cls.module_cache_ttl
|
||||
if current_time - cache_time < ttl:
|
||||
return limits
|
||||
|
||||
# 缓存不存在或已过期,从数据库查询
|
||||
# 缓存不存在或已过期,从内存缓存获取
|
||||
try:
|
||||
start_time = time.time()
|
||||
limits = await asyncio.wait_for(
|
||||
PluginLimit.filter(module=module, status=True).all(),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
elapsed = time.time() - start_time
|
||||
if elapsed > WARNING_THRESHOLD: # 记录耗时超过500ms的查询
|
||||
logger.warning(
|
||||
f"查询模块限制信息耗时: {elapsed:.3f}s, 模块: {module}",
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
|
||||
# 更新缓存
|
||||
cls.module_limit_cache[module] = (current_time, limits)
|
||||
await PluginLimitMemoryCache.ensure_loaded()
|
||||
limits = await PluginLimitMemoryCache.get_limits(module)
|
||||
cls.module_limit_cache[module] = (current_time, limits, False)
|
||||
return limits
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"查询模块限制信息超时: {module}", LOGGER_COMMAND)
|
||||
# 超时时返回空列表,避免阻塞
|
||||
except Exception as exc:
|
||||
logger.error(f"get module limits failed: {module}", LOGGER_COMMAND, e=exc)
|
||||
cls.module_limit_cache[module] = (current_time, [], True)
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
@@ -275,15 +309,8 @@ class LimitManager:
|
||||
left_time = limiter.left_time(key_type)
|
||||
cd_str = TimeUtils.format_duration(left_time)
|
||||
format_kwargs = {"cd": cd_str}
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
MessageUtils.build_message(
|
||||
limit.result, format_args=format_kwargs
|
||||
).send(),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"发送限制消息超时: {limit.module}", LOGGER_COMMAND)
|
||||
notice_key = _limit_notice_key(limit, user_id, group_id, channel_id)
|
||||
_send_limit_notice(limit.result, format_kwargs, notice_key)
|
||||
raise SkipPluginException(
|
||||
f"{limit.module}({limit.limit_type}) 正在限制中..."
|
||||
)
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from nonebot.adapters import Event
|
||||
@@ -6,9 +5,8 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.cache.runtime_cache import GroupSnapshot, _parse_block_modules
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
from zhenxun.utils.enum import BlockType
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
@@ -16,74 +14,89 @@ from .exception import IsSuperuserException, SkipPluginException
|
||||
from .utils import freq, is_poke, send_message
|
||||
|
||||
|
||||
def _get_group_block_sets(
|
||||
group: GroupConsole | GroupSnapshot,
|
||||
) -> tuple[frozenset[str], frozenset[str]]:
|
||||
block_set = getattr(group, "block_plugin_set", None)
|
||||
super_block_set = getattr(group, "superuser_block_plugin_set", None)
|
||||
if block_set is None:
|
||||
block_set = _parse_block_modules(getattr(group, "block_plugin", "") or "")
|
||||
setattr(group, "block_plugin_set", block_set)
|
||||
if super_block_set is None:
|
||||
super_block_set = _parse_block_modules(
|
||||
getattr(group, "superuser_block_plugin", "") or ""
|
||||
)
|
||||
setattr(group, "superuser_block_plugin_set", super_block_set)
|
||||
return block_set, super_block_set
|
||||
|
||||
|
||||
class GroupCheck:
|
||||
def __init__(
|
||||
self, plugin: PluginInfo, group: GroupConsole, session: Uninfo, is_poke: bool
|
||||
self,
|
||||
plugin: PluginInfo,
|
||||
group: GroupConsole | GroupSnapshot,
|
||||
session: Uninfo,
|
||||
is_poke: bool,
|
||||
skip_group_block: bool,
|
||||
) -> None:
|
||||
self.session = session
|
||||
self.is_poke = is_poke
|
||||
self.plugin = plugin
|
||||
self.group_data = group
|
||||
self.group_id = group.group_id
|
||||
self.skip_group_block = skip_group_block
|
||||
(
|
||||
self.block_plugin_set,
|
||||
self.superuser_block_plugin_set,
|
||||
) = _get_group_block_sets(group)
|
||||
|
||||
async def check(self):
|
||||
start_time = time.time()
|
||||
try:
|
||||
# 检查超级用户禁用
|
||||
if (
|
||||
self.group_data
|
||||
and CommonUtils.format(self.plugin.module)
|
||||
in self.group_data.superuser_block_plugin
|
||||
):
|
||||
if freq.is_send_limit_message(self.plugin, self.group_id, self.is_poke):
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
send_message(
|
||||
self.session,
|
||||
"超级管理员禁用了该群此功能...",
|
||||
self.group_id,
|
||||
),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
if not self.skip_group_block:
|
||||
# 检查超级用户禁用
|
||||
if (
|
||||
self.group_data
|
||||
and self.plugin.module in self.superuser_block_plugin_set
|
||||
):
|
||||
if freq.is_send_limit_message(
|
||||
self.plugin, self.group_id, self.is_poke
|
||||
):
|
||||
await send_message(
|
||||
self.session,
|
||||
"超级管理员禁用了该群此功能...",
|
||||
self.group_id,
|
||||
background=True,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"发送消息超时: {self.group_id}", LOGGER_COMMAND)
|
||||
raise SkipPluginException(
|
||||
f"{self.plugin.name}({self.plugin.module})"
|
||||
f" 超级管理员禁用了该群此功能..."
|
||||
)
|
||||
raise SkipPluginException(
|
||||
f"{self.plugin.name}({self.plugin.module})"
|
||||
f" 超级管理员禁用了该群此功能..."
|
||||
)
|
||||
|
||||
# 检查普通禁用
|
||||
if (
|
||||
self.group_data
|
||||
and CommonUtils.format(self.plugin.module)
|
||||
in self.group_data.block_plugin
|
||||
):
|
||||
if freq.is_send_limit_message(self.plugin, self.group_id, self.is_poke):
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
send_message(
|
||||
self.session, "该群未开启此功能...", self.group_id
|
||||
),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
# 检查普通禁用
|
||||
if self.group_data and self.plugin.module in self.block_plugin_set:
|
||||
if freq.is_send_limit_message(
|
||||
self.plugin, self.group_id, self.is_poke
|
||||
):
|
||||
await send_message(
|
||||
self.session,
|
||||
"该群未开启此功能...",
|
||||
self.group_id,
|
||||
background=True,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"发送消息超时: {self.group_id}", LOGGER_COMMAND)
|
||||
raise SkipPluginException(
|
||||
f"{self.plugin.name}({self.plugin.module}) 未开启此功能..."
|
||||
)
|
||||
raise SkipPluginException(
|
||||
f"{self.plugin.name}({self.plugin.module}) 未开启此功能..."
|
||||
)
|
||||
|
||||
# 检查全局禁用
|
||||
if self.plugin.block_type == BlockType.GROUP:
|
||||
if freq.is_send_limit_message(self.plugin, self.group_id, self.is_poke):
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
send_message(
|
||||
self.session, "该功能在群组中已被禁用...", self.group_id
|
||||
),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"发送消息超时: {self.group_id}", LOGGER_COMMAND)
|
||||
await send_message(
|
||||
self.session,
|
||||
"该功能在群组中已被禁用...",
|
||||
self.group_id,
|
||||
background=True,
|
||||
)
|
||||
raise SkipPluginException(
|
||||
f"{self.plugin.name}({self.plugin.module})该插件在群组中已被禁用..."
|
||||
)
|
||||
@@ -98,7 +111,9 @@ class GroupCheck:
|
||||
|
||||
|
||||
class PluginCheck:
|
||||
def __init__(self, group: GroupConsole | None, session: Uninfo, is_poke: bool):
|
||||
def __init__(
|
||||
self, group: GroupConsole | GroupSnapshot | None, session: Uninfo, is_poke: bool
|
||||
):
|
||||
self.session = session
|
||||
self.is_poke = is_poke
|
||||
self.group_data = group
|
||||
@@ -117,13 +132,11 @@ class PluginCheck:
|
||||
"""
|
||||
if plugin.block_type == BlockType.PRIVATE:
|
||||
if freq.is_send_limit_message(plugin, self.session.user.id, self.is_poke):
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
send_message(self.session, "该功能在私聊中已被禁用..."),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("发送消息超时", LOGGER_COMMAND)
|
||||
await send_message(
|
||||
self.session,
|
||||
"该功能在私聊中已被禁用...",
|
||||
background=True,
|
||||
)
|
||||
raise SkipPluginException(
|
||||
f"{plugin.name}({plugin.module}) 该插件在私聊中已被禁用..."
|
||||
)
|
||||
@@ -147,13 +160,12 @@ class PluginCheck:
|
||||
|
||||
sid = self.group_id or self.session.user.id
|
||||
if freq.is_send_limit_message(plugin, sid, self.is_poke):
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
send_message(self.session, "全局未开启此功能...", sid),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"发送消息超时: {sid}", LOGGER_COMMAND)
|
||||
await send_message(
|
||||
self.session,
|
||||
"全局未开启此功能...",
|
||||
sid,
|
||||
background=True,
|
||||
)
|
||||
raise SkipPluginException(
|
||||
f"{plugin.name}({plugin.module}) 全局未开启此功能..."
|
||||
)
|
||||
@@ -167,7 +179,12 @@ class PluginCheck:
|
||||
|
||||
|
||||
async def auth_plugin(
|
||||
plugin: PluginInfo, group: GroupConsole | None, session: Uninfo, event: Event
|
||||
plugin: PluginInfo,
|
||||
group: GroupConsole | GroupSnapshot | None,
|
||||
session: Uninfo,
|
||||
event: Event,
|
||||
*,
|
||||
skip_group_block: bool = False,
|
||||
):
|
||||
"""插件状态
|
||||
|
||||
@@ -181,19 +198,21 @@ async def auth_plugin(
|
||||
is_poke_event = is_poke(event)
|
||||
user_check = PluginCheck(group, session, is_poke_event)
|
||||
|
||||
tasks = []
|
||||
if group:
|
||||
tasks.append(GroupCheck(plugin, group, session, is_poke_event).check())
|
||||
block_set, super_block_set = _get_group_block_sets(group)
|
||||
if (
|
||||
plugin.status
|
||||
and plugin.block_type != BlockType.GROUP
|
||||
and not block_set
|
||||
and not super_block_set
|
||||
):
|
||||
return
|
||||
await GroupCheck(
|
||||
plugin, group, session, is_poke_event, skip_group_block
|
||||
).check()
|
||||
else:
|
||||
tasks.append(user_check.check_user(plugin))
|
||||
tasks.append(user_check.check_global(plugin))
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
asyncio.gather(*tasks), timeout=DB_TIMEOUT_SECONDS * 2
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("插件用户/群组/全局检查超时...", LOGGER_COMMAND)
|
||||
await user_check.check_user(plugin)
|
||||
await user_check.check_global(plugin)
|
||||
|
||||
finally:
|
||||
# 记录总执行时间
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
import contextlib
|
||||
|
||||
from nonebot.adapters import Event
|
||||
@@ -13,6 +14,7 @@ from zhenxun.utils.utils import FreqLimiter
|
||||
from .config import LOGGER_COMMAND
|
||||
|
||||
base_config = Config.get("hook")
|
||||
_SEND_TASKS: set[asyncio.Task] = set()
|
||||
|
||||
|
||||
def is_poke(event: Event) -> bool:
|
||||
@@ -32,7 +34,10 @@ def is_poke(event: Event) -> bool:
|
||||
|
||||
|
||||
async def send_message(
|
||||
session: Uninfo, message: list | str, check_tag: str | None = None
|
||||
session: Uninfo,
|
||||
message: list | str,
|
||||
check_tag: str | None = None,
|
||||
background: bool = False,
|
||||
):
|
||||
"""发送消息
|
||||
|
||||
@@ -41,19 +46,28 @@ async def send_message(
|
||||
message: 消息
|
||||
check_tag: cd flag
|
||||
"""
|
||||
try:
|
||||
if not check_tag:
|
||||
await MessageUtils.build_message(message).send(reply_to=True)
|
||||
elif freq._flmt.check(check_tag):
|
||||
freq._flmt.start_cd(check_tag)
|
||||
await MessageUtils.build_message(message).send(reply_to=True)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"发送消息失败",
|
||||
LOGGER_COMMAND,
|
||||
session=session,
|
||||
e=e,
|
||||
)
|
||||
|
||||
async def _send():
|
||||
try:
|
||||
if not check_tag:
|
||||
await MessageUtils.build_message(message).send(reply_to=True)
|
||||
elif freq._flmt.check(check_tag):
|
||||
freq._flmt.start_cd(check_tag)
|
||||
await MessageUtils.build_message(message).send(reply_to=True)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"发送消息失败",
|
||||
LOGGER_COMMAND,
|
||||
session=session,
|
||||
e=e,
|
||||
)
|
||||
|
||||
if background:
|
||||
task = asyncio.create_task(_send())
|
||||
_SEND_TASKS.add(task)
|
||||
task.add_done_callback(_SEND_TASKS.discard)
|
||||
return
|
||||
await _send()
|
||||
|
||||
|
||||
class FreqUtils:
|
||||
|
||||
@@ -1,25 +1,39 @@
|
||||
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 tortoise.exceptions import IntegrityError
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
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.utils.enum import GoldHandle, PluginType
|
||||
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
|
||||
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
|
||||
@@ -34,6 +48,54 @@ from .auth.exception import (
|
||||
)
|
||||
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
|
||||
# 熔断计数器
|
||||
@@ -51,13 +113,302 @@ CIRCUIT_RESET_TIME = 300 # 5分钟
|
||||
# 并发控制:限制同时进入 hooks 并行检查的协程数
|
||||
|
||||
# 默认为 6,可通过环境变量 AUTH_HOOKS_CONCURRENCY_LIMIT 调整
|
||||
HOOKS_CONCURRENCY_LIMIT = base_config.get("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):
|
||||
@@ -120,76 +471,73 @@ def check_circuit_breaker(name):
|
||||
return CIRCUIT_BREAKERS[name]["active"]
|
||||
|
||||
|
||||
async def get_plugin_and_user(
|
||||
module: str, user_id: str
|
||||
) -> tuple[PluginInfo, UserConsole]:
|
||||
"""获取用户数据和插件信息
|
||||
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
|
||||
|
||||
参数:
|
||||
module: 模块名
|
||||
user_id: 用户id
|
||||
|
||||
异常:
|
||||
PermissionExemption: 插件数据不存在
|
||||
PermissionExemption: 插件类型为HIDDEN
|
||||
PermissionExemption: 重复创建用户
|
||||
PermissionExemption: 用户数据不存在
|
||||
|
||||
返回:
|
||||
tuple[PluginInfo, UserConsole]: 插件信息,用户信息
|
||||
"""
|
||||
user_dao = DataAccess(UserConsole)
|
||||
plugin_dao = DataAccess(PluginInfo)
|
||||
|
||||
# 并行查询插件和用户数据
|
||||
plugin_task = plugin_dao.safe_get_or_none(module=module)
|
||||
user_task = user_dao.get_by_func_or_none(
|
||||
UserConsole.get_user, False, user_id=user_id
|
||||
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"
|
||||
)
|
||||
|
||||
try:
|
||||
plugin, user = await with_timeout(
|
||||
asyncio.gather(plugin_task, user_task), name="get_plugin_and_user"
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
# 如果并行查询超时,尝试串行查询
|
||||
logger.warning("并行查询超时,尝试串行查询", LOGGER_COMMAND)
|
||||
plugin = await with_timeout(
|
||||
plugin_dao.safe_get_or_none(module=module), name="get_plugin"
|
||||
)
|
||||
user = await with_timeout(
|
||||
user_dao.safe_get_or_none(user_id=user_id), name="get_user"
|
||||
)
|
||||
except IntegrityError:
|
||||
await asyncio.sleep(0.5)
|
||||
plugin_task = plugin_dao.safe_get_or_none(module=module)
|
||||
user_task = user_dao.get_by_func_or_none(
|
||||
UserConsole.get_user, False, user_id=user_id
|
||||
)
|
||||
plugin, user = await with_timeout(
|
||||
asyncio.gather(plugin_task, user_task), name="get_plugin_and_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"插件:{module} 数据不存在,已跳过权限检查...")
|
||||
raise PermissionExemption(f"plugin:{module} not found, skip permission check")
|
||||
if plugin.plugin_type == PluginType.HIDDEN:
|
||||
raise PermissionExemption(
|
||||
f"插件: {plugin.name}:{plugin.module} 为HIDDEN,已跳过权限检查..."
|
||||
)
|
||||
raise PermissionExemption(f"plugin {plugin.name}:{plugin.module} hidden, skip")
|
||||
|
||||
user = None
|
||||
try:
|
||||
user = await user_dao.get_by_func_or_none(
|
||||
UserConsole.get_user, False, user_id=user_id
|
||||
)
|
||||
except IntegrityError as e:
|
||||
raise PermissionExemption("重复创建用户,已跳过该次权限检查...") from e
|
||||
if not user:
|
||||
raise PermissionExemption("用户数据不存在,已跳过权限检查...")
|
||||
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, plugin: PluginInfo, session: Uninfo
|
||||
bot: Bot, user: UserConsole | None, plugin: PluginInfo, session: Uninfo
|
||||
) -> int:
|
||||
"""获取插件费用
|
||||
|
||||
@@ -278,7 +626,7 @@ async def _enter_hooks_section():
|
||||
await HOOKS_SEMAPHORE.acquire()
|
||||
async with HOOKS_ACTIVE_LOCK:
|
||||
HOOKS_ACTIVE_COUNT += 1
|
||||
logger.debug(f"当前并发权限检查数量: {HOOKS_ACTIVE_COUNT}", LOGGER_COMMAND)
|
||||
_debug_log(f"当前并发权限检查数量: {HOOKS_ACTIVE_COUNT}", LOGGER_COMMAND)
|
||||
|
||||
|
||||
async def _leave_hooks_section():
|
||||
@@ -292,7 +640,85 @@ async def _leave_hooks_section():
|
||||
HOOKS_ACTIVE_COUNT -= 1
|
||||
# 保证计数不为负
|
||||
HOOKS_ACTIVE_COUNT = max(HOOKS_ACTIVE_COUNT, 0)
|
||||
logger.debug(f"当前并发权限检查数量: {HOOKS_ACTIVE_COUNT}", LOGGER_COMMAND)
|
||||
_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(
|
||||
@@ -301,6 +727,8 @@ async def auth(
|
||||
bot: Bot,
|
||||
session: Uninfo,
|
||||
message: UniMsg,
|
||||
*,
|
||||
skip_ban: bool = False,
|
||||
):
|
||||
"""权限检查
|
||||
|
||||
@@ -316,6 +744,10 @@ async def auth(
|
||||
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 = {}
|
||||
@@ -328,11 +760,46 @@ async def auth(
|
||||
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), name="get_plugin_and_user"
|
||||
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:
|
||||
@@ -343,54 +810,200 @@ async def auth(
|
||||
)
|
||||
raise PermissionExemption("获取插件和用户数据超时,请稍后再试...")
|
||||
|
||||
# 进入 hooks 并行检查区域(会在高并发时排队)
|
||||
await _enter_hooks_section()
|
||||
entered_hooks = True
|
||||
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"
|
||||
|
||||
# 获取插件费用
|
||||
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
|
||||
)
|
||||
# 继续执行,不阻止权限检查
|
||||
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 = None
|
||||
if entity.group_id:
|
||||
group_dao = DataAccess(GroupConsole)
|
||||
group = await with_timeout(
|
||||
group_dao.safe_get_or_none(
|
||||
group_id=entity.group_id, channel_id__isnull=True
|
||||
),
|
||||
name="get_group",
|
||||
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 = [
|
||||
time_hook(auth_ban(matcher, bot, session, plugin), "auth_ban", hook_times),
|
||||
time_hook(auth_bot(plugin, bot.self_id), "auth_bot", hook_times),
|
||||
time_hook(
|
||||
auth_group(plugin, group, message, entity.group_id),
|
||||
"auth_group",
|
||||
hook_times,
|
||||
),
|
||||
time_hook(auth_admin(plugin, session), "auth_admin", hook_times),
|
||||
time_hook(
|
||||
auth_plugin(plugin, group, session, event), "auth_plugin", hook_times
|
||||
),
|
||||
time_hook(auth_limit(plugin, session), "auth_limit", hook_times),
|
||||
]
|
||||
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:
|
||||
@@ -408,15 +1021,19 @@ async def auth(
|
||||
# 不抛出异常,允许继续执行
|
||||
|
||||
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:
|
||||
@@ -428,6 +1045,8 @@ async def auth(
|
||||
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()
|
||||
|
||||
@@ -1,34 +1,159 @@
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from nonebot import get_driver
|
||||
from nonebot.adapters import Bot, Event
|
||||
from nonebot.exception import IgnoredException
|
||||
from nonebot.matcher import Matcher
|
||||
from nonebot.message import run_postprocessor, run_preprocessor
|
||||
from nonebot.message import event_preprocessor, run_postprocessor, run_preprocessor
|
||||
from nonebot_plugin_alconna import UniMsg
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.services.cache.runtime_cache import is_cache_ready
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import is_overloaded, signal_overload
|
||||
from zhenxun.utils.utils import get_entity_ids
|
||||
|
||||
from .auth.config import LOGGER_COMMAND
|
||||
from .auth_checker import LimitManager, auth
|
||||
from .auth.exception import SkipPluginException
|
||||
from .auth_checker import (
|
||||
LimitManager,
|
||||
_get_event_cache,
|
||||
auth,
|
||||
auth_ban_fast,
|
||||
auth_precheck,
|
||||
route_precheck,
|
||||
)
|
||||
|
||||
_SKIP_AUTH_PLUGINS = {"chat_history", "chat_message"}
|
||||
_BOT_CONNECT_TS: float | None = None
|
||||
_AUTH_QUEUE_MAXSIZE = 200
|
||||
_AUTH_QUEUE_HIGH_WATER = 160
|
||||
_AUTH_OVERLOAD_WINDOW = 5.0
|
||||
_AUTH_QUEUE: asyncio.Queue[tuple[Matcher, Event, Bot, Uninfo, UniMsg]] = asyncio.Queue(
|
||||
maxsize=_AUTH_QUEUE_MAXSIZE
|
||||
)
|
||||
_AUTH_QUEUE_STARTED = False
|
||||
_AUTH_WORKERS: list[asyncio.Task] = []
|
||||
_LAST_DROP_LOG = 0.0
|
||||
|
||||
driver = get_driver()
|
||||
|
||||
|
||||
@driver.on_bot_connect
|
||||
async def _mark_bot_connected(bot: Bot):
|
||||
del bot
|
||||
global _BOT_CONNECT_TS
|
||||
_BOT_CONNECT_TS = time.time()
|
||||
|
||||
|
||||
async def _auth_worker(worker_id: int) -> None:
|
||||
while True:
|
||||
matcher, event, bot, session, message = await _AUTH_QUEUE.get()
|
||||
try:
|
||||
await auth(
|
||||
matcher,
|
||||
event,
|
||||
bot,
|
||||
session,
|
||||
message,
|
||||
skip_ban=True,
|
||||
)
|
||||
except IgnoredException:
|
||||
pass
|
||||
except Exception as exc:
|
||||
if not is_overloaded():
|
||||
logger.error("async auth failed", LOGGER_COMMAND, e=exc)
|
||||
finally:
|
||||
_AUTH_QUEUE.task_done()
|
||||
|
||||
|
||||
@driver.on_startup
|
||||
async def _start_auth_queue():
|
||||
global _AUTH_QUEUE_STARTED
|
||||
if _AUTH_QUEUE_STARTED:
|
||||
return
|
||||
_AUTH_QUEUE_STARTED = True
|
||||
worker_count = max(1, min(6, _AUTH_QUEUE_MAXSIZE // 50))
|
||||
for idx in range(worker_count):
|
||||
_AUTH_WORKERS.append(asyncio.create_task(_auth_worker(idx)))
|
||||
|
||||
|
||||
def _skip_auth_for_plugin(matcher: Matcher) -> bool:
|
||||
if not matcher.plugin:
|
||||
return False
|
||||
name = (matcher.plugin.name or "").lower()
|
||||
if name in _SKIP_AUTH_PLUGINS:
|
||||
return True
|
||||
module_name = getattr(matcher.plugin, "module_name", "") or ""
|
||||
return "chat_history" in module_name
|
||||
|
||||
|
||||
@event_preprocessor
|
||||
async def _drop_message_before_cache_ready(event: Event):
|
||||
if event.get_type() != "message":
|
||||
return
|
||||
if not is_cache_ready():
|
||||
raise IgnoredException("cache not ready ignore")
|
||||
if _BOT_CONNECT_TS is not None:
|
||||
event_ts = getattr(event, "time", None)
|
||||
if event_ts is not None and event_ts < _BOT_CONNECT_TS:
|
||||
raise IgnoredException("drop backlog message")
|
||||
|
||||
|
||||
# # 权限检测
|
||||
@run_preprocessor
|
||||
async def _(matcher: Matcher, event: Event, bot: Bot, session: Uninfo, message: UniMsg):
|
||||
async def _auth_preprocessor(
|
||||
matcher: Matcher, event: Event, bot: Bot, session: Uninfo, message: UniMsg
|
||||
):
|
||||
if event.get_type() == "message" and not is_cache_ready():
|
||||
raise IgnoredException("cache not ready ignore")
|
||||
if _skip_auth_for_plugin(matcher):
|
||||
return
|
||||
start_time = time.time()
|
||||
await auth(
|
||||
matcher,
|
||||
event,
|
||||
bot,
|
||||
session,
|
||||
message,
|
||||
)
|
||||
logger.debug(f"权限检测耗时:{time.time() - start_time}秒", LOGGER_COMMAND)
|
||||
entity = get_entity_ids(session)
|
||||
event_cache = _get_event_cache(event, session, entity)
|
||||
if await route_precheck(matcher, event, session, message):
|
||||
return
|
||||
try:
|
||||
await auth_ban_fast(matcher, event, bot, session)
|
||||
except SkipPluginException as exc:
|
||||
logger.info(str(exc), LOGGER_COMMAND, session=session)
|
||||
raise IgnoredException("ban fast ignore") from exc
|
||||
try:
|
||||
await auth_precheck(matcher, event, bot, session, message)
|
||||
except SkipPluginException as exc:
|
||||
logger.info(str(exc), LOGGER_COMMAND, session=session)
|
||||
raise IgnoredException("precheck ignore") from exc
|
||||
|
||||
if event_cache is not None and event_cache.get("route_skip") is True:
|
||||
if not is_overloaded():
|
||||
logger.debug("route miss skip auth task", LOGGER_COMMAND)
|
||||
return
|
||||
|
||||
try:
|
||||
_AUTH_QUEUE.put_nowait((matcher, event, bot, session, message))
|
||||
except asyncio.QueueFull:
|
||||
signal_overload(_AUTH_OVERLOAD_WINDOW)
|
||||
now = time.monotonic()
|
||||
global _LAST_DROP_LOG
|
||||
if now - _LAST_DROP_LOG > 1.0:
|
||||
_LAST_DROP_LOG = now
|
||||
logger.warning("auth queue full, skip auth task", LOGGER_COMMAND)
|
||||
return
|
||||
if _AUTH_QUEUE.qsize() >= _AUTH_QUEUE_HIGH_WATER:
|
||||
signal_overload(_AUTH_OVERLOAD_WINDOW)
|
||||
now = time.monotonic()
|
||||
last_log = getattr(_auth_preprocessor, "_last_log", 0.0)
|
||||
if now - last_log > 1.0 and not is_overloaded():
|
||||
setattr(_auth_preprocessor, "_last_log", now)
|
||||
logger.debug(
|
||||
f"auth check cost: {time.time() - start_time:.3f}s",
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
|
||||
|
||||
# 解除命令block阻塞
|
||||
@run_postprocessor
|
||||
async def _(matcher: Matcher, session: Uninfo):
|
||||
async def _unblock_after_matcher(matcher: Matcher, session: Uninfo):
|
||||
user_id = session.user.id
|
||||
group_id = None
|
||||
channel_id = None
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
from datetime import datetime
|
||||
import os
|
||||
from pathlib import Path
|
||||
@@ -33,6 +34,57 @@ WELCOME_PATH = DATA_PATH / "welcome_message"
|
||||
|
||||
DEFAULT_IMAGE_PATH = IMAGE_PATH / "qxz"
|
||||
|
||||
_API_SEMAPHORE = asyncio.Semaphore(4)
|
||||
_API_TIMEOUT = 5.0
|
||||
_REFRESH_TASKS: set[asyncio.Task] = set()
|
||||
|
||||
|
||||
def _normalize_platform(platform: str | set[str] | None) -> str | None:
|
||||
if isinstance(platform, set):
|
||||
return next(iter(platform), None)
|
||||
return platform
|
||||
|
||||
|
||||
async def _safe_get_group_member_info(bot: Bot, group_id: str, user_id: str) -> dict:
|
||||
async with _API_SEMAPHORE:
|
||||
try:
|
||||
return await asyncio.wait_for(
|
||||
bot.get_group_member_info(
|
||||
group_id=int(group_id), user_id=int(user_id), no_cache=True
|
||||
),
|
||||
timeout=_API_TIMEOUT,
|
||||
)
|
||||
except (asyncio.TimeoutError, ActionFailed, Exception) as e:
|
||||
logger.warning("获取用户信息失败", e=e)
|
||||
return {"user_id": user_id, "group_id": group_id, "nickname": ""}
|
||||
|
||||
|
||||
async def _safe_get_group_info(bot: Bot, group_id: str) -> dict | None:
|
||||
async with _API_SEMAPHORE:
|
||||
try:
|
||||
return await asyncio.wait_for(
|
||||
bot.get_group_info(group_id=group_id),
|
||||
timeout=_API_TIMEOUT,
|
||||
)
|
||||
except (asyncio.TimeoutError, ActionFailed, Exception) as e:
|
||||
logger.warning("获取群信息失败", e=e)
|
||||
return None
|
||||
|
||||
|
||||
async def _refresh_member_info_async(
|
||||
bot: Bot, group_id: str, user_id: str, platform: str | None
|
||||
) -> None:
|
||||
user_info = await _safe_get_group_member_info(bot, group_id, user_id)
|
||||
await GroupInfoUser.update_or_create(
|
||||
user_id=str(user_info["user_id"]),
|
||||
group_id=str(user_info["group_id"]),
|
||||
defaults={
|
||||
"user_name": user_info.get("nickname") or "",
|
||||
"nickname": user_info.get("card") or user_info.get("nickname") or "",
|
||||
"platform": platform,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class GroupManager:
|
||||
_flmt = FreqLimiter(limit_cd)
|
||||
@@ -56,7 +108,14 @@ class GroupManager:
|
||||
if plugin_list := await PluginInfo.filter(default_status=False).all():
|
||||
for plugin in plugin_list:
|
||||
block_plugin += f"<{plugin.module},"
|
||||
group_info = await bot.get_group_info(group_id=group_id)
|
||||
group_info = await _safe_get_group_info(bot, group_id)
|
||||
if not group_info:
|
||||
logger.warning(
|
||||
"获取群信息失败,跳过群信息写入",
|
||||
"入群检测",
|
||||
group_id=group_id,
|
||||
)
|
||||
return
|
||||
await GroupConsole.update_or_create(
|
||||
group_id=group_info["group_id"],
|
||||
defaults={
|
||||
@@ -259,22 +318,27 @@ class GroupManager:
|
||||
else:
|
||||
group_id = session.group.id
|
||||
join_time = datetime.now()
|
||||
try:
|
||||
user_info = await bot.get_group_member_info(
|
||||
group_id=int(group_id), user_id=int(user_id), no_cache=True
|
||||
)
|
||||
except ActionFailed as e:
|
||||
logger.warning("获取用户信息识别...", e=e)
|
||||
user_info = {"user_id": user_id, "group_id": group_id, "nickname": ""}
|
||||
user_name = getattr(session.user, "name", None) or getattr(
|
||||
session.user, "nick", None
|
||||
)
|
||||
platform = PlatformUtils.get_platform(session)
|
||||
await GroupInfoUser.update_or_create(
|
||||
user_id=str(user_info["user_id"]),
|
||||
group_id=str(user_info["group_id"]),
|
||||
user_id=str(user_id),
|
||||
group_id=str(group_id),
|
||||
defaults={
|
||||
"user_name": user_info["nickname"],
|
||||
"user_name": user_name or "",
|
||||
"user_join_time": join_time,
|
||||
"platform": platform,
|
||||
},
|
||||
)
|
||||
logger.info(f"用户{user_info['user_id']} 所属{user_info['group_id']} 更新成功")
|
||||
task = asyncio.create_task(
|
||||
_refresh_member_info_async(
|
||||
bot, str(group_id), str(user_id), _normalize_platform(platform)
|
||||
)
|
||||
)
|
||||
_REFRESH_TASKS.add(task)
|
||||
task.add_done_callback(_REFRESH_TASKS.discard)
|
||||
logger.info(f"用户{user_id} 所属{group_id} 更新成功")
|
||||
if not await CommonUtils.task_is_block(
|
||||
session, "group_welcome"
|
||||
) and cls._flmt.check(group_id):
|
||||
@@ -295,7 +359,7 @@ class GroupManager:
|
||||
operator_name = user.user_name
|
||||
else:
|
||||
operator_name = "None"
|
||||
group = await GroupConsole.get_group(group_id)
|
||||
group = await GroupConsole.get_group_db(group_id)
|
||||
group_name = group.group_name if group else ""
|
||||
if group:
|
||||
await group.delete()
|
||||
@@ -342,10 +406,15 @@ class GroupManager:
|
||||
)
|
||||
if sub_type == "kick":
|
||||
if operator_id != "0":
|
||||
operator = await bot.get_group_member_info(
|
||||
user_id=int(operator_id), group_id=int(group_id)
|
||||
operator_user = await GroupInfoUser.get_or_none(
|
||||
user_id=operator_id, group_id=group_id
|
||||
)
|
||||
operator_name = operator["card"] or operator["nickname"]
|
||||
if operator_user:
|
||||
operator_name = (
|
||||
operator_user.nickname or operator_user.user_name or operator_id
|
||||
)
|
||||
else:
|
||||
operator_name = operator_id
|
||||
else:
|
||||
operator_name = ""
|
||||
return f"{user_name} 被 {operator_name} 送走了."
|
||||
|
||||
@@ -72,6 +72,18 @@ _t = on_message(priority=999, block=False, rule=lambda: False)
|
||||
|
||||
|
||||
cache = CacheRoot.cache_dict("REQUEST_CACHE", 60, str)
|
||||
_API_TIMEOUT = 5.0
|
||||
|
||||
|
||||
async def _safe_get_group_info(bot, group_id: str):
|
||||
try:
|
||||
return await asyncio.wait_for(
|
||||
bot.get_group_info(group_id=group_id),
|
||||
timeout=_API_TIMEOUT,
|
||||
)
|
||||
except (asyncio.TimeoutError, ActionFailed, Exception) as e:
|
||||
logger.warning("获取群信息失败", "群邀请", e=e)
|
||||
return None
|
||||
|
||||
|
||||
@friend_req.handle()
|
||||
@@ -162,17 +174,16 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio
|
||||
await bot.set_group_add_request(
|
||||
flag=event.flag, sub_type="invite", approve=True
|
||||
)
|
||||
if isinstance(bot, v11Bot):
|
||||
group_info = await bot.get_group_info(group_id=event.group_id)
|
||||
max_member_count = group_info["max_member_count"]
|
||||
member_count = group_info["member_count"]
|
||||
group_info = await _safe_get_group_info(bot, str(event.group_id))
|
||||
if isinstance(bot, v11Bot) and group_info:
|
||||
max_member_count = group_info.get("max_member_count", 0)
|
||||
member_count = group_info.get("member_count", 0)
|
||||
else:
|
||||
group_info = await bot.get_group_info(group_id=str(event.group_id))
|
||||
max_member_count = 0
|
||||
member_count = 0
|
||||
group.max_member_count = max_member_count
|
||||
group.member_count = member_count
|
||||
group.group_name = group_info["group_name"]
|
||||
group.group_name = group_info.get("group_name", "") if group_info else ""
|
||||
await group.save(
|
||||
update_fields=["group_name", "max_member_count", "member_count"]
|
||||
)
|
||||
|
||||
@@ -2,6 +2,7 @@ import nonebot
|
||||
from nonebot_plugin_apscheduler import scheduler
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import should_pause_tasks
|
||||
from zhenxun.services.tags import tag_manager
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
@@ -13,6 +14,8 @@ from zhenxun.utils.platform import PlatformUtils
|
||||
minute=1,
|
||||
)
|
||||
async def _():
|
||||
if should_pause_tasks():
|
||||
return
|
||||
bots = nonebot.get_bots()
|
||||
for bot in bots.values():
|
||||
try:
|
||||
@@ -29,6 +32,8 @@ async def _():
|
||||
minute=1,
|
||||
)
|
||||
async def _():
|
||||
if should_pause_tasks():
|
||||
return
|
||||
bots = nonebot.get_bots()
|
||||
for bot in bots.values():
|
||||
try:
|
||||
@@ -43,8 +48,8 @@ async def _():
|
||||
# 自动清理静态标签中的无效群组
|
||||
@scheduler.scheduled_job(
|
||||
"cron",
|
||||
hour=23,
|
||||
minute=30,
|
||||
hour=4,
|
||||
minute=50,
|
||||
)
|
||||
async def _prune_stale_tags():
|
||||
deleted_count = await tag_manager.prune_stale_group_links()
|
||||
|
||||
@@ -9,6 +9,7 @@ from zhenxun.models.chat_history import ChatHistory
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.task_info import TaskInfo
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import should_pause_tasks
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
Config.add_plugin_config(
|
||||
@@ -27,6 +28,8 @@ Config.add_plugin_config(
|
||||
minute=40,
|
||||
)
|
||||
async def _():
|
||||
if should_pause_tasks():
|
||||
return
|
||||
if not Config.get_config("chat_history", "FLAG"):
|
||||
logger.debug("未开启历史发言记录,过滤群组发言检测...")
|
||||
return
|
||||
|
||||
@@ -174,8 +174,9 @@ async def _(
|
||||
|
||||
|
||||
@scheduler.scheduled_job(
|
||||
"interval",
|
||||
hours=1,
|
||||
"cron",
|
||||
hour=4,
|
||||
minute=10,
|
||||
)
|
||||
async def _():
|
||||
try:
|
||||
|
||||
@@ -12,6 +12,7 @@ from zhenxun.configs.utils import PluginExtraData
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.statistics import Statistics
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import should_pause_tasks
|
||||
from zhenxun.utils.enum import PluginType
|
||||
|
||||
__plugin_meta__ = PluginMetadata(
|
||||
@@ -53,9 +54,11 @@ async def _(
|
||||
)
|
||||
|
||||
|
||||
@scheduler.scheduled_job("interval", minutes=1, max_instances=5)
|
||||
@scheduler.scheduled_job("interval", minutes=30, max_instances=1, coalesce=True)
|
||||
async def _():
|
||||
try:
|
||||
if should_pause_tasks():
|
||||
return
|
||||
call_list = TEMP_LIST.copy()
|
||||
TEMP_LIST.clear()
|
||||
if call_list:
|
||||
|
||||
@@ -3,6 +3,8 @@ from typing import cast
|
||||
import nonebot
|
||||
from nonebot.adapters import Bot
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from tortoise.exceptions import IntegrityError, TransactionManagementError
|
||||
from tortoise.transactions import in_transaction
|
||||
|
||||
from zhenxun.configs.utils import PluginExtraData
|
||||
from zhenxun.models.bot_console import BotConsole
|
||||
@@ -72,9 +74,16 @@ async def init_bot_console(bot: Bot):
|
||||
list[str], await TaskInfo.filter(status=True).values_list("module", flat=True)
|
||||
)
|
||||
platform = PlatformUtils.get_platform(bot)
|
||||
bot_data, created = await BotConsole.get_or_create(
|
||||
bot_id=bot.self_id, platform=platform
|
||||
)
|
||||
try:
|
||||
bot_data, created = await BotConsole.get_or_create(
|
||||
bot_id=bot.self_id, platform=platform
|
||||
)
|
||||
except (IntegrityError, TransactionManagementError):
|
||||
async with in_transaction() as connection:
|
||||
bot_data = (
|
||||
await BotConsole.filter(bot_id=bot.self_id).using_db(connection).get()
|
||||
)
|
||||
created = False
|
||||
|
||||
if not created:
|
||||
task_list = await _filter_blocked_items(
|
||||
|
||||
@@ -163,7 +163,7 @@ async def _(session: EventSession, arparma: Arparma, state: T_State, level: int)
|
||||
@_matcher.assign("super-handle", parameterless=[CheckGroupId()])
|
||||
async def _(session: EventSession, arparma: Arparma, state: T_State):
|
||||
gid = state["group_id"]
|
||||
group = await GroupConsole.get_group(group_id=gid)
|
||||
group = await GroupConsole.get_group_db(group_id=gid)
|
||||
if not group:
|
||||
await MessageUtils.build_message("群组信息不存在, 请更新群组信息...").finish()
|
||||
s = "删除" if arparma.find("delete") else "添加"
|
||||
|
||||
@@ -209,7 +209,7 @@ async def _(param: HandleRequest) -> Result:
|
||||
if not (req := await FgRequest.get_or_none(id=param.id)):
|
||||
return Result.warning_("未找到此Id请求...")
|
||||
if req.request_type == RequestType.GROUP:
|
||||
if group := await GroupConsole.get_group(group_id=req.group_id):
|
||||
if group := await GroupConsole.get_group_db(group_id=req.group_id):
|
||||
group.group_flag = 1
|
||||
await group.save(update_fields=["group_flag"])
|
||||
else:
|
||||
|
||||
@@ -33,7 +33,7 @@ class ApiDataSource:
|
||||
参数:
|
||||
group: UpdateGroup
|
||||
"""
|
||||
db_group = await GroupConsole.get_group(group.group_id) or GroupConsole(
|
||||
db_group = await GroupConsole.get_group_db(group.group_id) or GroupConsole(
|
||||
group_id=group.group_id
|
||||
)
|
||||
task_list = await TaskInfo.all().values_list("module", flat=True)
|
||||
@@ -250,7 +250,7 @@ class ApiDataSource:
|
||||
返回:
|
||||
GroupDetail | None: 群组详情数据
|
||||
"""
|
||||
group = await GroupConsole.get_group(group_id=group_id)
|
||||
group = await GroupConsole.get_group_db(group_id=group_id)
|
||||
if not group:
|
||||
return None
|
||||
like_plugin = await cls.__get_group_detail_like_plugin(group_id)
|
||||
|
||||
@@ -4,6 +4,7 @@ from typing_extensions import Self
|
||||
|
||||
from tortoise import fields
|
||||
|
||||
from zhenxun.services.cache.runtime_cache import BanMemoryCache
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.db_context import Model
|
||||
from zhenxun.services.log import logger
|
||||
@@ -42,6 +43,18 @@ class BanConsole(Model):
|
||||
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT]
|
||||
"""开启锁"""
|
||||
|
||||
@classmethod
|
||||
async def create(cls, *args, **kwargs) -> Self:
|
||||
result = await super().create(*args, **kwargs)
|
||||
await BanMemoryCache.upsert_from_model(result)
|
||||
return result
|
||||
|
||||
async def delete(self, *args, **kwargs):
|
||||
user_id = self.user_id
|
||||
group_id = self.group_id
|
||||
await super().delete(*args, **kwargs)
|
||||
await BanMemoryCache.remove(user_id, group_id)
|
||||
|
||||
@classmethod
|
||||
async def _get_data(cls, user_id: str | None, group_id: str | None) -> Self | None:
|
||||
"""获取数据
|
||||
@@ -82,14 +95,10 @@ class BanConsole(Model):
|
||||
返回:
|
||||
bool: 权限判断,能否unban
|
||||
"""
|
||||
user = await cls._get_data(user_id, group_id)
|
||||
if user:
|
||||
logger.debug(
|
||||
f"检测用户被ban等级,user_level: {user.ban_level},level: {level}",
|
||||
target=f"{group_id}:{user_id}",
|
||||
)
|
||||
return user.ban_level <= level
|
||||
return False
|
||||
logger.debug("检测用户被ban等级", target=f"{group_id}:{user_id}")
|
||||
if not BanMemoryCache.is_loaded():
|
||||
return False
|
||||
return BanMemoryCache.check_ban_level(user_id, group_id, level)
|
||||
|
||||
@classmethod
|
||||
async def check_ban_time(
|
||||
@@ -104,17 +113,9 @@ class BanConsole(Model):
|
||||
int: ban剩余时长,-1时为永久ban,0表示未被ban
|
||||
"""
|
||||
logger.debug("获取用户ban时长", target=f"{group_id}:{user_id}")
|
||||
user = await cls._get_data(user_id, group_id)
|
||||
if not user and user_id:
|
||||
user = await cls._get_data(user_id, None)
|
||||
if user:
|
||||
if user.duration == -1:
|
||||
return -1
|
||||
_time = time.time() - (user.ban_time + user.duration)
|
||||
if _time < 0:
|
||||
return int(abs(_time))
|
||||
await user.delete()
|
||||
return 0
|
||||
if not BanMemoryCache.is_loaded():
|
||||
return 0
|
||||
return BanMemoryCache.remaining_time(user_id, group_id)
|
||||
|
||||
@classmethod
|
||||
async def is_ban(cls, user_id: str | None, group_id: str | None = None) -> bool:
|
||||
@@ -127,11 +128,7 @@ class BanConsole(Model):
|
||||
bool: 是否被ban
|
||||
"""
|
||||
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
|
||||
return (await cls.check_ban_time(user_id, group_id)) != 0
|
||||
|
||||
@classmethod
|
||||
async def ban(
|
||||
|
||||
@@ -2,6 +2,7 @@ from typing import Literal, overload
|
||||
|
||||
from tortoise import fields
|
||||
|
||||
from zhenxun.services.cache.runtime_cache import BotMemoryCache
|
||||
from zhenxun.services.db_context import Model
|
||||
from zhenxun.utils.enum import CacheType
|
||||
|
||||
@@ -160,8 +161,10 @@ class BotConsole(Model):
|
||||
affected_rows = await cls.filter(bot_id=bot_id).update(status=status)
|
||||
if not affected_rows:
|
||||
raise ValueError(f"未找到 bot_id: {bot_id}")
|
||||
await BotMemoryCache.update_status(bot_id, status)
|
||||
else:
|
||||
await cls.all().update(status=status)
|
||||
await BotMemoryCache.refresh()
|
||||
|
||||
@overload
|
||||
@classmethod
|
||||
@@ -434,6 +437,27 @@ class BotConsole(Model):
|
||||
bot_data, _ = await cls.get_or_create(bot_id=bot_id)
|
||||
return cls.format(task_name) in bot_data.block_tasks
|
||||
|
||||
@classmethod
|
||||
async def create(cls, *args, **kwargs):
|
||||
result = await super().create(*args, **kwargs)
|
||||
await BotMemoryCache.upsert_from_model(result)
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
async def update_or_create(cls, *args, **kwargs):
|
||||
result = await super().update_or_create(*args, **kwargs)
|
||||
await BotMemoryCache.upsert_from_model(result[0])
|
||||
return result
|
||||
|
||||
async def save(self, *args, **kwargs):
|
||||
await super().save(*args, **kwargs)
|
||||
await BotMemoryCache.upsert_from_model(self)
|
||||
|
||||
async def delete(self, *args, **kwargs):
|
||||
bot_id = self.bot_id
|
||||
await super().delete(*args, **kwargs)
|
||||
await BotMemoryCache.remove(bot_id)
|
||||
|
||||
@classmethod
|
||||
async def _run_script(cls):
|
||||
return [
|
||||
|
||||
+78
-101
@@ -1,4 +1,4 @@
|
||||
from typing import Any, ClassVar, cast, overload
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, cast, overload
|
||||
from typing_extensions import Self
|
||||
|
||||
from tortoise import fields
|
||||
@@ -7,10 +7,14 @@ from tortoise.backends.base.client import BaseDBAsyncClient
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.task_info import TaskInfo
|
||||
from zhenxun.services.cache import CacheRoot
|
||||
from zhenxun.services.cache.runtime_cache import GroupMemoryCache
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.db_context import Model
|
||||
from zhenxun.utils.enum import CacheType, DbLockType, PluginType
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from zhenxun.services.cache.runtime_cache import GroupSnapshot
|
||||
|
||||
|
||||
def add_disable_marker(name: str) -> str:
|
||||
"""添加模块禁用标记符
|
||||
@@ -155,6 +159,7 @@ class GroupConsole(Model):
|
||||
|
||||
# 更新缓存
|
||||
await cls._update_cache(group)
|
||||
await GroupMemoryCache.upsert_from_model(group)
|
||||
|
||||
return group
|
||||
|
||||
@@ -210,6 +215,7 @@ class GroupConsole(Model):
|
||||
# 更新缓存
|
||||
if is_create:
|
||||
await cls._update_cache(group)
|
||||
await GroupMemoryCache.upsert_from_model(group)
|
||||
|
||||
return group, is_create
|
||||
|
||||
@@ -235,26 +241,37 @@ class GroupConsole(Model):
|
||||
|
||||
# 更新缓存
|
||||
await cls._update_cache(group)
|
||||
await GroupMemoryCache.upsert_from_model(group)
|
||||
|
||||
return group, is_create
|
||||
|
||||
async def save(self, *args, **kwargs):
|
||||
await super().save(*args, **kwargs)
|
||||
await GroupMemoryCache.upsert_from_model(self)
|
||||
|
||||
async def delete(self, *args, **kwargs):
|
||||
group_id = self.group_id
|
||||
channel_id = self.channel_id
|
||||
await super().delete(*args, **kwargs)
|
||||
await GroupMemoryCache.remove(group_id, channel_id)
|
||||
|
||||
@classmethod
|
||||
async def get_group(
|
||||
cls,
|
||||
group_id: str,
|
||||
channel_id: str | None = None,
|
||||
clean_duplicates: bool = True,
|
||||
) -> "GroupSnapshot | None":
|
||||
return GroupMemoryCache.get_if_ready(group_id, channel_id)
|
||||
|
||||
@classmethod
|
||||
async def get_group_db(
|
||||
cls,
|
||||
group_id: str,
|
||||
channel_id: str | None = None,
|
||||
clean_duplicates: bool = True,
|
||||
) -> Self | None:
|
||||
"""获取群组
|
||||
|
||||
参数:
|
||||
group_id: 群组id
|
||||
channel_id: 频道id
|
||||
clean_duplicates: 是否删除重复的记录,仅保留最新的
|
||||
|
||||
返回:
|
||||
Self: GroupConsole
|
||||
"""
|
||||
"""获取群组(数据库)"""
|
||||
dao = DataAccess(cls)
|
||||
if channel_id:
|
||||
return await dao.safe_get_or_none(
|
||||
@@ -270,49 +287,32 @@ class GroupConsole(Model):
|
||||
|
||||
@classmethod
|
||||
async def is_super_group(cls, group_id: str) -> bool:
|
||||
"""是否超级用户指定群
|
||||
|
||||
参数:
|
||||
group_id: 群组id
|
||||
|
||||
返回:
|
||||
bool: 是否超级用户指定群
|
||||
"""
|
||||
return group.is_super if (group := await cls.get_group(group_id)) else False
|
||||
group = GroupMemoryCache.get_if_ready(group_id, None)
|
||||
return bool(group and group.is_super)
|
||||
|
||||
@classmethod
|
||||
async def is_superuser_block_plugin(cls, group_id: str, module: str) -> bool:
|
||||
"""查看群组是否超级用户禁用功能
|
||||
|
||||
参数:
|
||||
group_id: 群组id
|
||||
module: 模块名称
|
||||
|
||||
返回:
|
||||
bool: 是否禁用被动
|
||||
"""
|
||||
return await cls.exists(
|
||||
group_id=group_id,
|
||||
superuser_block_plugin__contains=add_disable_marker(module),
|
||||
group = GroupMemoryCache.get_if_ready(group_id, None)
|
||||
if not group:
|
||||
return False
|
||||
return bool(
|
||||
group.superuser_block_plugin_set
|
||||
and module in group.superuser_block_plugin_set
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def is_block_plugin(cls, group_id: str, module: str) -> bool:
|
||||
"""查看群组是否禁用插件
|
||||
|
||||
参数:
|
||||
group_id: 群组id
|
||||
plugin: 插件名称
|
||||
|
||||
返回:
|
||||
bool: 是否禁用插件
|
||||
"""
|
||||
module = add_disable_marker(module)
|
||||
return await cls.exists(
|
||||
group_id=group_id, block_plugin__contains=module
|
||||
) or await cls.exists(
|
||||
group_id=group_id, superuser_block_plugin__contains=module
|
||||
)
|
||||
group = GroupMemoryCache.get_if_ready(group_id, None)
|
||||
if not group:
|
||||
return False
|
||||
if group.block_plugin_set and module in group.block_plugin_set:
|
||||
return True
|
||||
if (
|
||||
group.superuser_block_plugin_set
|
||||
and module in group.superuser_block_plugin_set
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
async def set_block_plugin(
|
||||
@@ -396,70 +396,47 @@ class GroupConsole(Model):
|
||||
async def is_normal_block_plugin(
|
||||
cls, group_id: str, module: str, channel_id: str | None = None
|
||||
) -> bool:
|
||||
"""查看群组是否禁用功能
|
||||
|
||||
参数:
|
||||
group_id: 群组id
|
||||
module: 模块名称
|
||||
channel_id: 频道id
|
||||
|
||||
返回:
|
||||
bool: 是否禁用被动
|
||||
"""
|
||||
return await cls.exists(
|
||||
group_id=group_id,
|
||||
channel_id=channel_id,
|
||||
block_plugin__contains=f"<{module},",
|
||||
)
|
||||
group = GroupMemoryCache.get_if_ready(group_id, channel_id)
|
||||
if not group:
|
||||
return False
|
||||
return bool(group.block_plugin_set and module in group.block_plugin_set)
|
||||
|
||||
@classmethod
|
||||
async def is_superuser_block_task(cls, group_id: str, task: str) -> bool:
|
||||
"""查看群组是否超级用户禁用被动
|
||||
|
||||
参数:
|
||||
group_id: 群组id
|
||||
task: 模块名称
|
||||
|
||||
返回:
|
||||
bool: 是否禁用被动
|
||||
"""
|
||||
return await cls.exists(
|
||||
group_id=group_id,
|
||||
superuser_block_task__contains=add_disable_marker(task),
|
||||
group = GroupMemoryCache.get_if_ready(group_id, None)
|
||||
if not group:
|
||||
return False
|
||||
return bool(
|
||||
group.superuser_block_task_set and task in group.superuser_block_task_set
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def is_block_task(
|
||||
cls, group_id: str, task: str, channel_id: str | None = None
|
||||
) -> bool:
|
||||
"""查看群组是否禁用被动
|
||||
|
||||
参数:
|
||||
group_id: 群组id
|
||||
task: 任务模块
|
||||
channel_id: 频道id
|
||||
|
||||
返回:
|
||||
bool: 是否禁用被动
|
||||
"""
|
||||
task = add_disable_marker(task)
|
||||
if not channel_id:
|
||||
return await cls.exists(
|
||||
group_id=group_id,
|
||||
channel_id__isnull=True,
|
||||
block_task__contains=task,
|
||||
) or await cls.exists(
|
||||
group_id=group_id,
|
||||
channel_id__isnull=True,
|
||||
superuser_block_task__contains=task,
|
||||
)
|
||||
return await cls.exists(
|
||||
group_id=group_id, channel_id=channel_id, block_task__contains=task
|
||||
) or await cls.exists(
|
||||
group_id=group_id,
|
||||
channel_id__isnull=True,
|
||||
superuser_block_task__contains=task,
|
||||
)
|
||||
group = GroupMemoryCache.get_if_ready(group_id, None)
|
||||
if not group:
|
||||
return False
|
||||
if group.block_task_set and task in group.block_task_set:
|
||||
return True
|
||||
if (
|
||||
group.superuser_block_task_set
|
||||
and task in group.superuser_block_task_set
|
||||
):
|
||||
return True
|
||||
return False
|
||||
group = GroupMemoryCache.get_if_ready(group_id, channel_id)
|
||||
if group and group.block_task_set and task in group.block_task_set:
|
||||
return True
|
||||
super_group = GroupMemoryCache.get_if_ready(group_id, None)
|
||||
if (
|
||||
super_group
|
||||
and super_group.superuser_block_task_set
|
||||
and task in super_group.superuser_block_task_set
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
async def set_block_task(
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from tortoise import fields
|
||||
|
||||
from zhenxun.services.cache.runtime_cache import LevelUserMemoryCache
|
||||
from zhenxun.services.db_context import Model
|
||||
from zhenxun.utils.enum import CacheType
|
||||
|
||||
@@ -124,6 +125,28 @@ class LevelUser(Model):
|
||||
return user.group_flag == 1
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
async def create(cls, *args, **kwargs):
|
||||
result = await super().create(*args, **kwargs)
|
||||
await LevelUserMemoryCache.upsert_from_model(result)
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
async def update_or_create(cls, *args, **kwargs):
|
||||
result = await super().update_or_create(*args, **kwargs)
|
||||
await LevelUserMemoryCache.upsert_from_model(result[0])
|
||||
return result
|
||||
|
||||
async def save(self, *args, **kwargs):
|
||||
await super().save(*args, **kwargs)
|
||||
await LevelUserMemoryCache.upsert_from_model(self)
|
||||
|
||||
async def delete(self, *args, **kwargs):
|
||||
user_id = self.user_id
|
||||
group_id = self.group_id
|
||||
await super().delete(*args, **kwargs)
|
||||
await LevelUserMemoryCache.remove(user_id, group_id)
|
||||
|
||||
@classmethod
|
||||
async def _run_script(cls):
|
||||
return [
|
||||
|
||||
@@ -109,4 +109,5 @@ class PluginInfo(Model):
|
||||
"ALTER TABLE plugin_info ADD COLUMN is_show boolean DEFAULT true;",
|
||||
"ALTER TABLE plugin_info ADD COLUMN ignore_prompt boolean DEFAULT false;",
|
||||
"ALTER TABLE plugin_info ADD COLUMN impression float DEFAULT 0;",
|
||||
"CREATE INDEX idx_plugin_info_module ON plugin_info(module);",
|
||||
]
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from tortoise import fields
|
||||
|
||||
from zhenxun.services.cache.runtime_cache import PluginLimitMemoryCache
|
||||
from zhenxun.services.db_context import Model
|
||||
from zhenxun.utils.enum import LimitCheckType, LimitWatchType, PluginLimitType
|
||||
|
||||
@@ -38,3 +39,24 @@ class PluginLimit(Model):
|
||||
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
|
||||
table = "plugin_limit"
|
||||
table_description = "插件限制"
|
||||
|
||||
@classmethod
|
||||
async def create(cls, *args, **kwargs):
|
||||
result = await super().create(*args, **kwargs)
|
||||
await PluginLimitMemoryCache.upsert_from_model(result)
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
async def update_or_create(cls, *args, **kwargs):
|
||||
result = await super().update_or_create(*args, **kwargs)
|
||||
await PluginLimitMemoryCache.upsert_from_model(result[0])
|
||||
return result
|
||||
|
||||
async def save(self, *args, **kwargs):
|
||||
await super().save(*args, **kwargs)
|
||||
await PluginLimitMemoryCache.upsert_from_model(self)
|
||||
|
||||
async def delete(self, *args, **kwargs):
|
||||
limit_id = self.id
|
||||
await super().delete(*args, **kwargs)
|
||||
await PluginLimitMemoryCache.remove_by_id(limit_id)
|
||||
|
||||
@@ -1,4 +1,8 @@
|
||||
import asyncio
|
||||
from typing import ClassVar
|
||||
|
||||
from tortoise import fields
|
||||
from tortoise.exceptions import IntegrityError
|
||||
|
||||
from zhenxun.models.goods_info import GoodsInfo
|
||||
from zhenxun.services.db_context import Model
|
||||
@@ -36,6 +40,29 @@ class UserConsole(Model):
|
||||
cache_key_field = "user_id"
|
||||
"""缓存键字段"""
|
||||
|
||||
_uid_counter: ClassVar[int | None] = None
|
||||
_uid_lock: ClassVar[asyncio.Lock] = asyncio.Lock()
|
||||
|
||||
@classmethod
|
||||
async def get_or_create_user(
|
||||
cls, user_id: str, platform: str | None = None
|
||||
) -> tuple["UserConsole", bool]:
|
||||
for attempt in range(2):
|
||||
try:
|
||||
return await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
except IntegrityError:
|
||||
async with cls._uid_lock:
|
||||
cls._uid_counter = None
|
||||
if attempt >= 1:
|
||||
raise
|
||||
return await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def get_user(cls, user_id: str, platform: str | None = None) -> "UserConsole":
|
||||
"""获取用户
|
||||
@@ -47,15 +74,8 @@ class UserConsole(Model):
|
||||
返回:
|
||||
UserConsole: UserConsole
|
||||
"""
|
||||
if not await cls.exists(user_id=user_id):
|
||||
await cls.create(
|
||||
user_id=user_id, platform=platform, uid=await cls.get_new_uid()
|
||||
)
|
||||
# user, _ = await UserConsole.get_or_create(
|
||||
# user_id=user_id,
|
||||
# defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
# )
|
||||
return await cls.get(user_id=user_id)
|
||||
user, _ = await cls.get_or_create_user(user_id=user_id, platform=platform)
|
||||
return user
|
||||
|
||||
@classmethod
|
||||
async def get_new_uid(cls) -> int:
|
||||
@@ -64,9 +84,12 @@ class UserConsole(Model):
|
||||
返回:
|
||||
int: 最新uid
|
||||
"""
|
||||
if user := await cls.annotate().order_by("-uid").first():
|
||||
return user.uid + 1
|
||||
return 1
|
||||
async with cls._uid_lock:
|
||||
if cls._uid_counter is None:
|
||||
user = await cls.annotate().order_by("-uid").first()
|
||||
cls._uid_counter = user.uid if user else 0
|
||||
cls._uid_counter += 1
|
||||
return cls._uid_counter
|
||||
|
||||
@classmethod
|
||||
async def add_gold(
|
||||
@@ -80,10 +103,7 @@ class UserConsole(Model):
|
||||
source: 来源
|
||||
platform: 平台.
|
||||
"""
|
||||
user, _ = await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
user, _ = await cls.get_or_create_user(user_id=user_id, platform=platform)
|
||||
user.gold += gold
|
||||
await user.save(update_fields=["gold"])
|
||||
await UserGoldLog.create(
|
||||
@@ -111,10 +131,7 @@ class UserConsole(Model):
|
||||
异常:
|
||||
InsufficientGold: 金币不足
|
||||
"""
|
||||
user, _ = await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
user, _ = await cls.get_or_create_user(user_id=user_id, platform=platform)
|
||||
if user.gold < gold:
|
||||
raise InsufficientGold()
|
||||
user.gold -= gold
|
||||
@@ -135,10 +152,7 @@ class UserConsole(Model):
|
||||
num: 道具数量.
|
||||
platform: 平台.
|
||||
"""
|
||||
user, _ = await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
user, _ = await cls.get_or_create_user(user_id=user_id, platform=platform)
|
||||
if goods_uuid not in user.props:
|
||||
user.props[goods_uuid] = 0
|
||||
user.props[goods_uuid] += num
|
||||
@@ -172,10 +186,7 @@ class UserConsole(Model):
|
||||
num: 道具数量.
|
||||
platform: 平台.
|
||||
"""
|
||||
user, _ = await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
user, _ = await cls.get_or_create_user(user_id=user_id, platform=platform)
|
||||
|
||||
if goods_uuid not in user.props or user.props[goods_uuid] < num:
|
||||
raise GoodsNotFound("未找到商品或道具数量不足...")
|
||||
|
||||
@@ -135,7 +135,9 @@ avatar_service = AvatarService()
|
||||
|
||||
|
||||
@scheduler.scheduled_job(
|
||||
"interval", hours=Config.get_config("avatar_cache", "CLEANUP_INTERVAL_HOURS", 24)
|
||||
"cron",
|
||||
hour=4,
|
||||
minute=30,
|
||||
)
|
||||
async def _run_avatar_cache_cleanup():
|
||||
await avatar_service._cleanup_cache()
|
||||
|
||||
Vendored
+6
-3
@@ -81,7 +81,8 @@ import asyncio
|
||||
from collections.abc import Callable
|
||||
from datetime import datetime
|
||||
from functools import wraps
|
||||
from typing import Any, ClassVar, Generic, TypeVar, get_type_hints
|
||||
from typing import Any, ClassVar, Generic, TypeVar, cast, get_type_hints
|
||||
from typing_extensions import Self
|
||||
|
||||
from aiocache import Cache as AioCache
|
||||
from aiocache import SimpleMemoryCache
|
||||
@@ -115,6 +116,8 @@ __all__ = [
|
||||
"CacheRoot",
|
||||
]
|
||||
|
||||
from . import runtime_cache as _runtime_cache # noqa: F401
|
||||
|
||||
T = TypeVar("T")
|
||||
U = TypeVar("U")
|
||||
|
||||
@@ -291,11 +294,11 @@ class CacheManager:
|
||||
_dict_caches: ClassVar[dict[str, "CacheDict"]] = {}
|
||||
_enabled = False # 缓存启用标记
|
||||
|
||||
def __new__(cls) -> "CacheManager":
|
||||
def __new__(cls) -> Self:
|
||||
"""单例模式"""
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
return cls._instance
|
||||
return cast(Self, cls._instance)
|
||||
|
||||
@property
|
||||
def enabled(self) -> bool:
|
||||
|
||||
+16
-12
@@ -47,7 +47,8 @@ class CacheDict(Generic[T]):
|
||||
T: 字典值
|
||||
"""
|
||||
if value := self._data.get(key):
|
||||
if self.expire_time(key):
|
||||
if value.expire_time > 0 and value.expire_time < time.time():
|
||||
del self._data[key]
|
||||
raise KeyError(f"键 {key} 已过期")
|
||||
return value.value
|
||||
raise KeyError(f"键 {key} 不存在")
|
||||
@@ -80,7 +81,13 @@ class CacheDict(Generic[T]):
|
||||
返回:
|
||||
bool: 是否存在
|
||||
"""
|
||||
return False if key not in self._data else bool(self.expire_time(key))
|
||||
data = self._data.get(key)
|
||||
if data is None:
|
||||
return False
|
||||
if data.expire_time > 0 and data.expire_time < time.time():
|
||||
del self._data[key]
|
||||
return False
|
||||
return True
|
||||
|
||||
def get(self, key: str, default: Any = None) -> T | None:
|
||||
"""获取字典项,如果不存在返回默认值
|
||||
@@ -93,11 +100,11 @@ class CacheDict(Generic[T]):
|
||||
Any: 字典值或默认值
|
||||
"""
|
||||
if value := self._data.get(key):
|
||||
if self.expire_time(key):
|
||||
if value.expire_time > 0 and value.expire_time < time.time():
|
||||
del self._data[key]
|
||||
return default
|
||||
if not value:
|
||||
return default
|
||||
return default if value.value is None else value.value
|
||||
return default if value.value is None else value.value
|
||||
return default
|
||||
|
||||
def set(self, key: str, value: Any, expire: int | None = None):
|
||||
"""设置字典项
|
||||
@@ -126,16 +133,13 @@ class CacheDict(Generic[T]):
|
||||
返回:
|
||||
Any: 字典值或默认值
|
||||
"""
|
||||
if key not in self._data:
|
||||
data = self._data.get(key)
|
||||
if data is None:
|
||||
return default
|
||||
|
||||
data = self._data.pop(key)
|
||||
|
||||
# 检查是否过期
|
||||
if data.expire_time > 0 and data.expire_time < time.time():
|
||||
del self._data[key]
|
||||
return default
|
||||
|
||||
del self._data[key]
|
||||
return data.value
|
||||
|
||||
def clear(self) -> None:
|
||||
|
||||
+1728
File diff suppressed because it is too large
Load Diff
@@ -5,7 +5,11 @@ from typing import Any, ClassVar
|
||||
from typing_extensions import Self
|
||||
|
||||
from tortoise.backends.base.client import BaseDBAsyncClient
|
||||
from tortoise.exceptions import IntegrityError, MultipleObjectsReturned
|
||||
from tortoise.exceptions import (
|
||||
IntegrityError,
|
||||
MultipleObjectsReturned,
|
||||
TransactionManagementError,
|
||||
)
|
||||
from tortoise.models import Model as TortoiseModel
|
||||
from tortoise.transactions import in_transaction
|
||||
|
||||
@@ -129,9 +133,23 @@ class Model(TortoiseModel):
|
||||
**kwargs: Any,
|
||||
) -> tuple[Self, bool]:
|
||||
"""获取或创建数据(无锁版本,依赖数据库约束)"""
|
||||
result = await super().get_or_create(
|
||||
defaults=defaults, using_db=using_db, **kwargs
|
||||
)
|
||||
try:
|
||||
result = await super().get_or_create(
|
||||
defaults=defaults, using_db=using_db, **kwargs
|
||||
)
|
||||
except IntegrityError:
|
||||
# 并发创建冲突时,回退为查询已存在记录
|
||||
try:
|
||||
if using_db is not None:
|
||||
obj = await cls.filter(**kwargs).using_db(using_db).get()
|
||||
result = (obj, False)
|
||||
else:
|
||||
raise TransactionManagementError("fallback to new transaction")
|
||||
except TransactionManagementError:
|
||||
async with in_transaction() as connection:
|
||||
obj = await cls.filter(**kwargs).using_db(connection).get()
|
||||
result = (obj, False)
|
||||
|
||||
if cache_type := cls.get_cache_type():
|
||||
await CacheRoot.invalidate_cache(cache_type, cls.get_cache_key(result[0]))
|
||||
return result
|
||||
|
||||
@@ -21,7 +21,7 @@ from typing import (
|
||||
get_origin,
|
||||
get_type_hints,
|
||||
)
|
||||
from typing_extensions import override
|
||||
from typing_extensions import Self, override
|
||||
|
||||
from httpx import NetworkError, TimeoutException
|
||||
|
||||
@@ -421,10 +421,10 @@ class ToolProviderManager:
|
||||
|
||||
_instance: "ToolProviderManager | None" = None
|
||||
|
||||
def __new__(cls) -> "ToolProviderManager":
|
||||
def __new__(cls) -> Self:
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
return cls._instance
|
||||
return cast(Self, cls._instance)
|
||||
|
||||
def __init__(self):
|
||||
if hasattr(self, "_initialized") and self._initialized:
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
_OVERLOAD_UNTIL = 0.0
|
||||
|
||||
|
||||
def signal_overload(duration: float = 5.0) -> None:
|
||||
"""Mark the system as overloaded for a short time window."""
|
||||
global _OVERLOAD_UNTIL
|
||||
if duration <= 0:
|
||||
return
|
||||
now = time.monotonic()
|
||||
until = now + duration
|
||||
if until > _OVERLOAD_UNTIL:
|
||||
_OVERLOAD_UNTIL = until
|
||||
|
||||
|
||||
def is_overloaded() -> bool:
|
||||
return time.monotonic() < _OVERLOAD_UNTIL
|
||||
|
||||
|
||||
def should_pause_tasks() -> bool:
|
||||
return is_overloaded()
|
||||
@@ -10,6 +10,7 @@ from collections.abc import Callable
|
||||
from datetime import datetime
|
||||
from functools import partial
|
||||
import random
|
||||
import time
|
||||
|
||||
import nonebot
|
||||
from nonebot.adapters import Bot
|
||||
@@ -23,6 +24,7 @@ from pydantic import BaseModel
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.scheduled_job import ScheduledJob
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import should_pause_tasks
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
from zhenxun.utils.decorator.retry import Retry
|
||||
from zhenxun.utils.pydantic_compat import parse_as
|
||||
@@ -32,6 +34,7 @@ from .types import ExecutionPolicy, ScheduleContext
|
||||
|
||||
JOB_PREFIX = "zhenxun_schedule_"
|
||||
SCHEDULE_CONCURRENCY_KEY = "all_groups_concurrency_limit"
|
||||
_LAST_PRESSURE_SKIP = 0.0
|
||||
|
||||
|
||||
class APSchedulerAdapter:
|
||||
@@ -362,6 +365,14 @@ async def _execute_job(
|
||||
logger.error("执行持久化任务时 schedule_id 不能为空。")
|
||||
return
|
||||
|
||||
global _LAST_PRESSURE_SKIP
|
||||
if should_pause_tasks():
|
||||
now = time.time()
|
||||
if now - _LAST_PRESSURE_SKIP > 30:
|
||||
_LAST_PRESSURE_SKIP = now
|
||||
logger.info("scheduler paused due to message pressure")
|
||||
return
|
||||
|
||||
scheduler_manager._running_tasks.add(schedule_id)
|
||||
try:
|
||||
schedule = await ScheduleRepository.get_by_id(schedule_id)
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
import asyncio
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import nonebot
|
||||
from nonebot.adapters import Bot
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
_SEND_APIS = {"send_msg", "send_like"}
|
||||
_QUEUE: asyncio.Queue[tuple[Bot, str, dict[str, Any], asyncio.Future]] = asyncio.Queue()
|
||||
_WORKERS = 3
|
||||
_MIN_INTERVAL = 0.05
|
||||
_SEND_LOCK = asyncio.Lock()
|
||||
_LAST_SEND_TS = 0.0
|
||||
_API_SEMAPHORE = asyncio.Semaphore(3)
|
||||
_ORIG_CALL_API = Bot.call_api
|
||||
_PATCHED = False
|
||||
_WORKER_TASKS: list[asyncio.Task] = []
|
||||
|
||||
|
||||
async def _rate_limit():
|
||||
global _LAST_SEND_TS
|
||||
async with _SEND_LOCK:
|
||||
now = time.monotonic()
|
||||
wait = _MIN_INTERVAL - (now - _LAST_SEND_TS)
|
||||
if wait > 0:
|
||||
await asyncio.sleep(wait)
|
||||
_LAST_SEND_TS = time.monotonic()
|
||||
|
||||
|
||||
async def _worker(worker_id: int):
|
||||
while True:
|
||||
bot, api, data, future = await _QUEUE.get()
|
||||
try:
|
||||
await _rate_limit()
|
||||
async with _API_SEMAPHORE:
|
||||
result = await _ORIG_CALL_API(bot, api, **data)
|
||||
if not future.done():
|
||||
future.set_result(result)
|
||||
except Exception as exc:
|
||||
if not future.done():
|
||||
future.set_exception(exc)
|
||||
logger.warning(
|
||||
f"send queue failed: {api}",
|
||||
"SendQueue",
|
||||
target=getattr(bot, "self_id", None),
|
||||
e=exc,
|
||||
)
|
||||
finally:
|
||||
_QUEUE.task_done()
|
||||
|
||||
|
||||
async def _queued_call_api(self: Bot, api: str, **data: Any):
|
||||
if api not in _SEND_APIS:
|
||||
return await _ORIG_CALL_API(self, api, **data)
|
||||
loop = asyncio.get_running_loop()
|
||||
future: asyncio.Future = loop.create_future()
|
||||
await _QUEUE.put((self, api, data, future))
|
||||
return await future
|
||||
|
||||
|
||||
def patch_send_queue() -> None:
|
||||
global _PATCHED
|
||||
if _PATCHED:
|
||||
return
|
||||
Bot.call_api = _queued_call_api # type: ignore[assignment]
|
||||
_PATCHED = True
|
||||
|
||||
|
||||
driver = nonebot.get_driver()
|
||||
|
||||
|
||||
@driver.on_startup
|
||||
async def _start_send_queue():
|
||||
patch_send_queue()
|
||||
for idx in range(_WORKERS):
|
||||
_WORKER_TASKS.append(asyncio.create_task(_worker(idx)))
|
||||
Reference in New Issue
Block a user