🎉🍾♻️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:
ManyManyTomato
2026-01-28 10:15:24 +08:00
committed by GitHub
co-authored by ATTomatoo pre-commit-ci[bot]
parent c9f0a8b9d9
commit 837330e30a
41 changed files with 3656 additions and 763 deletions
@@ -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)} 条", "定时任务")
-82
View File
@@ -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(
+76 -71
View File
@@ -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:
+10 -15
View File
@@ -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:
# 记录总执行时间
+28 -14
View File
@@ -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:
+721 -102
View File
@@ -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()
+139 -14
View File
@@ -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} 送走了."
+17 -6
View File
@@ -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
+3 -2
View File
@@ -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)
+21 -24
View File
@@ -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(
+24
View File
@@ -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
View File
@@ -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(
+23
View File
@@ -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 [
+1
View File
@@ -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);",
]
+22
View File
@@ -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)
+39 -28
View File
@@ -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("未找到商品或道具数量不足...")
+3 -1
View File
@@ -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()
+6 -3
View File
@@ -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
View File
@@ -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:
File diff suppressed because it is too large Load Diff
+22 -4
View File
@@ -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
+3 -3
View File
@@ -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:
+24
View File
@@ -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()
+11
View File
@@ -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)
+78
View File
@@ -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)))