diff --git a/zhenxun/builtin_plugins/admin/group_member_update/__init__.py b/zhenxun/builtin_plugins/admin/group_member_update/__init__.py index d8fb1070..c19becd3 100644 --- a/zhenxun/builtin_plugins/admin/group_member_update/__init__.py +++ b/zhenxun/builtin_plugins/admin/group_member_update/__init__.py @@ -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() diff --git a/zhenxun/builtin_plugins/admin/group_member_update/_data_source.py b/zhenxun/builtin_plugins/admin/group_member_update/_data_source.py index 39bcee29..eeac663a 100644 --- a/zhenxun/builtin_plugins/admin/group_member_update/_data_source.py +++ b/zhenxun/builtin_plugins/admin/group_member_update/_data_source.py @@ -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", diff --git a/zhenxun/builtin_plugins/admin/plugin_switch/_data_source.py b/zhenxun/builtin_plugins/admin/plugin_switch/_data_source.py index 3a599dd1..2ea48c68 100644 --- a/zhenxun/builtin_plugins/admin/plugin_switch/_data_source.py +++ b/zhenxun/builtin_plugins/admin/plugin_switch/_data_source.py @@ -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 diff --git a/zhenxun/builtin_plugins/chat_history/chat_message.py b/zhenxun/builtin_plugins/chat_history/chat_message.py index 36ea4930..0c0f7ff1 100644 --- a/zhenxun/builtin_plugins/chat_history/chat_message.py +++ b/zhenxun/builtin_plugins/chat_history/chat_message.py @@ -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)} 条", "定时任务") diff --git a/zhenxun/builtin_plugins/help_help.py b/zhenxun/builtin_plugins/help_help.py deleted file mode 100644 index 2c792c71..00000000 --- a/zhenxun/builtin_plugins/help_help.py +++ /dev/null @@ -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() diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_admin.py b/zhenxun/builtin_plugins/hooks/auth/auth_admin.py index 19059f98..4483303f 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_admin.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_admin.py @@ -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( diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_ban.py b/zhenxun/builtin_plugins/hooks/auth/auth_ban.py index 7eea7f57..b877da83 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_ban.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_ban.py @@ -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: diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_bot.py b/zhenxun/builtin_plugins/hooks/auth/auth_bot.py index ab902991..e1e5ed86 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_bot.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_bot.py @@ -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不存在或休眠中阻断权限检测...") diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_cost.py b/zhenxun/builtin_plugins/hooks/auth/auth_cost.py index 53da21a9..314bf254 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_cost.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_cost.py @@ -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}) 金币限制...") diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_group.py b/zhenxun/builtin_plugins/hooks/auth/auth_group.py index 20114bef..a6ccf471 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_group.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_group.py @@ -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("群组信息不存在...") diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_limit.py b/zhenxun/builtin_plugins/hooks/auth/auth_limit.py index 80650472..0a915765 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_limit.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_limit.py @@ -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}) 正在限制中..." ) diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_plugin.py b/zhenxun/builtin_plugins/hooks/auth/auth_plugin.py index ddab3161..0bc00dbb 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_plugin.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_plugin.py @@ -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: # 记录总执行时间 diff --git a/zhenxun/builtin_plugins/hooks/auth/utils.py b/zhenxun/builtin_plugins/hooks/auth/utils.py index d2f1b551..158098dd 100644 --- a/zhenxun/builtin_plugins/hooks/auth/utils.py +++ b/zhenxun/builtin_plugins/hooks/auth/utils.py @@ -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: diff --git a/zhenxun/builtin_plugins/hooks/auth_checker.py b/zhenxun/builtin_plugins/hooks/auth_checker.py index ae757cec..9f77bf27 100644 --- a/zhenxun/builtin_plugins/hooks/auth_checker.py +++ b/zhenxun/builtin_plugins/hooks/auth_checker.py @@ -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() diff --git a/zhenxun/builtin_plugins/hooks/auth_hook.py b/zhenxun/builtin_plugins/hooks/auth_hook.py index 34ea8018..99cb088d 100644 --- a/zhenxun/builtin_plugins/hooks/auth_hook.py +++ b/zhenxun/builtin_plugins/hooks/auth_hook.py @@ -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 diff --git a/zhenxun/builtin_plugins/platform/qq/group_handle/data_source.py b/zhenxun/builtin_plugins/platform/qq/group_handle/data_source.py index e92a41e3..5a353ade 100644 --- a/zhenxun/builtin_plugins/platform/qq/group_handle/data_source.py +++ b/zhenxun/builtin_plugins/platform/qq/group_handle/data_source.py @@ -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} 送走了." diff --git a/zhenxun/builtin_plugins/record_request.py b/zhenxun/builtin_plugins/record_request.py index 88f5539f..c4000535 100644 --- a/zhenxun/builtin_plugins/record_request.py +++ b/zhenxun/builtin_plugins/record_request.py @@ -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"] ) diff --git a/zhenxun/builtin_plugins/scheduler/auto_update_group.py b/zhenxun/builtin_plugins/scheduler/auto_update_group.py index 8f4b2578..bae378ad 100644 --- a/zhenxun/builtin_plugins/scheduler/auto_update_group.py +++ b/zhenxun/builtin_plugins/scheduler/auto_update_group.py @@ -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() diff --git a/zhenxun/builtin_plugins/scheduler/chat_check.py b/zhenxun/builtin_plugins/scheduler/chat_check.py index d7559665..67b80e61 100644 --- a/zhenxun/builtin_plugins/scheduler/chat_check.py +++ b/zhenxun/builtin_plugins/scheduler/chat_check.py @@ -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 diff --git a/zhenxun/builtin_plugins/sign_in/__init__.py b/zhenxun/builtin_plugins/sign_in/__init__.py index 3dd1863c..589b2312 100644 --- a/zhenxun/builtin_plugins/sign_in/__init__.py +++ b/zhenxun/builtin_plugins/sign_in/__init__.py @@ -174,8 +174,9 @@ async def _( @scheduler.scheduled_job( - "interval", - hours=1, + "cron", + hour=4, + minute=10, ) async def _(): try: diff --git a/zhenxun/builtin_plugins/statistics/statistics_hook.py b/zhenxun/builtin_plugins/statistics/statistics_hook.py index 3ac15e2a..b71102b2 100644 --- a/zhenxun/builtin_plugins/statistics/statistics_hook.py +++ b/zhenxun/builtin_plugins/statistics/statistics_hook.py @@ -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: diff --git a/zhenxun/builtin_plugins/superuser/bot_manage/__init__.py b/zhenxun/builtin_plugins/superuser/bot_manage/__init__.py index f7923b53..484803a5 100644 --- a/zhenxun/builtin_plugins/superuser/bot_manage/__init__.py +++ b/zhenxun/builtin_plugins/superuser/bot_manage/__init__.py @@ -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( diff --git a/zhenxun/builtin_plugins/superuser/group_manage.py b/zhenxun/builtin_plugins/superuser/group_manage.py index b2f77f47..7d42282c 100644 --- a/zhenxun/builtin_plugins/superuser/group_manage.py +++ b/zhenxun/builtin_plugins/superuser/group_manage.py @@ -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 "添加" diff --git a/zhenxun/builtin_plugins/web_ui/api/tabs/manage/__init__.py b/zhenxun/builtin_plugins/web_ui/api/tabs/manage/__init__.py index dfa9dd31..89d8a6ce 100644 --- a/zhenxun/builtin_plugins/web_ui/api/tabs/manage/__init__.py +++ b/zhenxun/builtin_plugins/web_ui/api/tabs/manage/__init__.py @@ -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: diff --git a/zhenxun/builtin_plugins/web_ui/api/tabs/manage/data_source.py b/zhenxun/builtin_plugins/web_ui/api/tabs/manage/data_source.py index 0b068e17..5cd2f82e 100644 --- a/zhenxun/builtin_plugins/web_ui/api/tabs/manage/data_source.py +++ b/zhenxun/builtin_plugins/web_ui/api/tabs/manage/data_source.py @@ -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) diff --git a/zhenxun/models/ban_console.py b/zhenxun/models/ban_console.py index 4fec9608..612d6660 100644 --- a/zhenxun/models/ban_console.py +++ b/zhenxun/models/ban_console.py @@ -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( diff --git a/zhenxun/models/bot_console.py b/zhenxun/models/bot_console.py index 01a93535..f57a326d 100644 --- a/zhenxun/models/bot_console.py +++ b/zhenxun/models/bot_console.py @@ -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 [ diff --git a/zhenxun/models/group_console.py b/zhenxun/models/group_console.py index e73c4cde..5d596e7d 100644 --- a/zhenxun/models/group_console.py +++ b/zhenxun/models/group_console.py @@ -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( diff --git a/zhenxun/models/level_user.py b/zhenxun/models/level_user.py index 644c38d3..764a1f31 100644 --- a/zhenxun/models/level_user.py +++ b/zhenxun/models/level_user.py @@ -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 [ diff --git a/zhenxun/models/plugin_info.py b/zhenxun/models/plugin_info.py index 533fe0e8..c25f1de8 100644 --- a/zhenxun/models/plugin_info.py +++ b/zhenxun/models/plugin_info.py @@ -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);", ] diff --git a/zhenxun/models/plugin_limit.py b/zhenxun/models/plugin_limit.py index bee905ff..3d36db7e 100644 --- a/zhenxun/models/plugin_limit.py +++ b/zhenxun/models/plugin_limit.py @@ -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) diff --git a/zhenxun/models/user_console.py b/zhenxun/models/user_console.py index 27cd582f..f9dbcb11 100644 --- a/zhenxun/models/user_console.py +++ b/zhenxun/models/user_console.py @@ -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("未找到商品或道具数量不足...") diff --git a/zhenxun/services/avatar_service.py b/zhenxun/services/avatar_service.py index 2a051a39..3ced2bfe 100644 --- a/zhenxun/services/avatar_service.py +++ b/zhenxun/services/avatar_service.py @@ -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() diff --git a/zhenxun/services/cache/__init__.py b/zhenxun/services/cache/__init__.py index 9e222a44..79088dd6 100644 --- a/zhenxun/services/cache/__init__.py +++ b/zhenxun/services/cache/__init__.py @@ -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: diff --git a/zhenxun/services/cache/cache_containers.py b/zhenxun/services/cache/cache_containers.py index aad8878f..e6829007 100644 --- a/zhenxun/services/cache/cache_containers.py +++ b/zhenxun/services/cache/cache_containers.py @@ -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: diff --git a/zhenxun/services/cache/runtime_cache.py b/zhenxun/services/cache/runtime_cache.py new file mode 100644 index 00000000..f3bec2bc --- /dev/null +++ b/zhenxun/services/cache/runtime_cache.py @@ -0,0 +1,1728 @@ +from __future__ import annotations + +import asyncio +from dataclasses import dataclass, field +import json +import os +import time +from typing import TYPE_CHECKING, Any, ClassVar +import uuid + +from zhenxun.configs.config import Config +from zhenxun.services.cache.config import CacheMode +from zhenxun.services.log import logger +from zhenxun.utils.enum import LimitCheckType, LimitWatchType, PluginLimitType +from zhenxun.utils.manager.priority_manager import PriorityLifecycle + +if TYPE_CHECKING: + from zhenxun.models.plugin_info import PluginInfo + +LOG_COMMAND = "RuntimeCache" + +Config.add_plugin_config( + "hook", + "PLUGININFO_MEM_REFRESH_INTERVAL", + 300, + help="plugin info memory cache refresh seconds", +) +Config.add_plugin_config( + "hook", + "BAN_MEM_REFRESH_INTERVAL", + 60, + help="ban memory cache full refresh seconds", +) +Config.add_plugin_config( + "hook", + "BAN_MEM_CLEAN_INTERVAL", + 60, + help="ban memory cache cleanup seconds", +) +Config.add_plugin_config( + "hook", + "BAN_MEM_CLEANUP_DB", + True, + help="delete expired ban records from database", +) +Config.add_plugin_config( + "hook", + "BAN_MEM_NEGATIVE_TTL", + 5, + help="ban memory negative cache ttl seconds", +) +Config.add_plugin_config( + "hook", + "BOT_MEM_REFRESH_INTERVAL", + 60, + help="bot memory cache refresh seconds", +) +Config.add_plugin_config( + "hook", + "BOT_MEM_NEGATIVE_TTL", + 60, + help="bot memory negative cache ttl seconds", +) +Config.add_plugin_config( + "hook", + "GROUP_MEM_REFRESH_INTERVAL", + 60, + help="group memory cache refresh seconds", +) +Config.add_plugin_config( + "hook", + "GROUP_MEM_NEGATIVE_TTL", + 60, + help="group memory negative cache ttl seconds", +) +Config.add_plugin_config( + "hook", + "LEVEL_MEM_REFRESH_INTERVAL", + 120, + help="level memory cache refresh seconds", +) +Config.add_plugin_config( + "hook", + "LEVEL_MEM_NEGATIVE_TTL", + 60, + help="level memory negative cache ttl seconds", +) +Config.add_plugin_config( + "hook", + "LIMIT_MEM_REFRESH_INTERVAL", + 60, + help="plugin limit memory cache refresh seconds", +) +Config.add_plugin_config( + "hook", + "LIMIT_MEM_NEGATIVE_TTL", + 30, + help="plugin limit negative cache ttl seconds", +) +Config.add_plugin_config( + "hook", + "RUNTIME_CACHE_SYNC_ENABLED", + True, + help="enable redis pubsub runtime cache sync", +) +Config.add_plugin_config( + "hook", + "RUNTIME_CACHE_SYNC_CHANNEL", + "ZHENXUN_RUNTIME_CACHE_SYNC", + help="redis pubsub channel for runtime cache sync", +) + + +def _coerce_int(value, default: int) -> int: + try: + value_int = int(value) + except (TypeError, ValueError): + return default + return value_int if value_int >= 0 else default + + +INSTANCE_ID = uuid.uuid4().hex +_CACHE_READY_EVENT = asyncio.Event() + + +def _env_get(name: str, default: str | None = None) -> str | None: + value = os.getenv(name) + if value is None: + value = os.getenv(name.lower()) + return value if value is not None else default + + +def _redis_enabled() -> bool: + mode = (_env_get("CACHE_MODE") or "").upper() + if mode != CacheMode.REDIS: + return False + return bool(_env_get("REDIS_HOST")) + + +def is_cache_ready() -> bool: + return _CACHE_READY_EVENT.is_set() + + +async def wait_cache_ready(timeout: float | None = None) -> bool: + try: + await asyncio.wait_for(_CACHE_READY_EVENT.wait(), timeout=timeout) + return True + except asyncio.TimeoutError: + return False + + +def _parse_block_modules(value: str) -> frozenset[str]: + if not value: + return frozenset() + items = [] + for part in value.split("<"): + part = part.strip() + if not part: + continue + part = part.strip(",").strip() + if part: + items.append(part) + return frozenset(items) + + +@dataclass(frozen=True) +class BanEntry: + user_id: str | None + group_id: str | None + ban_level: int + ban_time: int + duration: int + expire_at: float | None + + def remaining(self, now: float | None = None) -> int: + if self.duration == -1: + return -1 + now_ts = time.time() if now is None else now + left = int(self.ban_time + self.duration - now_ts) + return left if left > 0 else 0 + + def to_payload(self) -> dict[str, Any]: + return { + "user_id": self.user_id, + "group_id": self.group_id, + "ban_level": self.ban_level, + "ban_time": self.ban_time, + "duration": self.duration, + "expire_at": self.expire_at, + } + + @classmethod + def from_payload(cls, payload: dict[str, Any]) -> "BanEntry": + user_id = payload.get("user_id") + group_id = payload.get("group_id") + ban_time = int(payload.get("ban_time", 0) or 0) + duration = int(payload.get("duration", 0) or 0) + expire_at = payload.get("expire_at") + if expire_at is None and duration != -1: + expire_at = float(ban_time + duration) + return cls( + user_id=str(user_id) if user_id else None, + group_id=str(group_id) if group_id else None, + ban_level=int(payload.get("ban_level", 0) or 0), + ban_time=ban_time, + duration=duration, + expire_at=expire_at, + ) + + +@dataclass(frozen=True) +class BotSnapshot: + bot_id: str + status: bool + platform: str | None + block_plugins: str + block_tasks: str + available_plugins: str + available_tasks: str + + @classmethod + def from_model(cls, model) -> "BotSnapshot": + return cls( + bot_id=str(model.bot_id), + status=bool(model.status), + platform=getattr(model, "platform", None), + block_plugins=getattr(model, "block_plugins", "") or "", + block_tasks=getattr(model, "block_tasks", "") or "", + available_plugins=getattr(model, "available_plugins", "") or "", + available_tasks=getattr(model, "available_tasks", "") or "", + ) + + def to_payload(self) -> dict[str, Any]: + return { + "bot_id": self.bot_id, + "status": self.status, + "platform": self.platform, + "block_plugins": self.block_plugins, + "block_tasks": self.block_tasks, + "available_plugins": self.available_plugins, + "available_tasks": self.available_tasks, + } + + @classmethod + def from_payload(cls, payload: dict[str, Any]) -> "BotSnapshot": + return cls( + bot_id=str(payload.get("bot_id", "")), + status=bool(payload.get("status", True)), + platform=payload.get("platform"), + block_plugins=payload.get("block_plugins", "") or "", + block_tasks=payload.get("block_tasks", "") or "", + available_plugins=payload.get("available_plugins", "") or "", + available_tasks=payload.get("available_tasks", "") or "", + ) + + +@dataclass(frozen=True) +class GroupSnapshot: + group_id: str + channel_id: str | None + group_name: str + max_member_count: int + member_count: int + status: bool + level: int + is_super: bool + group_flag: int + block_plugin: str + superuser_block_plugin: str + block_task: str + superuser_block_task: str + platform: str | None + block_plugin_set: frozenset[str] = field(default_factory=frozenset) + superuser_block_plugin_set: frozenset[str] = field(default_factory=frozenset) + block_task_set: frozenset[str] = field(default_factory=frozenset) + superuser_block_task_set: frozenset[str] = field(default_factory=frozenset) + + @classmethod + def from_model(cls, model) -> "GroupSnapshot": + block_plugin = getattr(model, "block_plugin", "") or "" + superuser_block_plugin = getattr(model, "superuser_block_plugin", "") or "" + block_task = getattr(model, "block_task", "") or "" + superuser_block_task = getattr(model, "superuser_block_task", "") or "" + return cls( + group_id=str(model.group_id), + channel_id=getattr(model, "channel_id", None), + group_name=getattr(model, "group_name", "") or "", + max_member_count=int(getattr(model, "max_member_count", 0) or 0), + member_count=int(getattr(model, "member_count", 0) or 0), + status=bool(getattr(model, "status", True)), + level=int(getattr(model, "level", 0) or 0), + is_super=bool(getattr(model, "is_super", False)), + group_flag=int(getattr(model, "group_flag", 0) or 0), + block_plugin=block_plugin, + superuser_block_plugin=superuser_block_plugin, + block_task=block_task, + superuser_block_task=superuser_block_task, + platform=getattr(model, "platform", None), + block_plugin_set=_parse_block_modules(block_plugin), + superuser_block_plugin_set=_parse_block_modules(superuser_block_plugin), + block_task_set=_parse_block_modules(block_task), + superuser_block_task_set=_parse_block_modules(superuser_block_task), + ) + + def to_payload(self) -> dict[str, Any]: + return { + "group_id": self.group_id, + "channel_id": self.channel_id, + "group_name": self.group_name, + "max_member_count": self.max_member_count, + "member_count": self.member_count, + "status": self.status, + "level": self.level, + "is_super": self.is_super, + "group_flag": self.group_flag, + "block_plugin": self.block_plugin, + "superuser_block_plugin": self.superuser_block_plugin, + "block_task": self.block_task, + "superuser_block_task": self.superuser_block_task, + "platform": self.platform, + } + + @classmethod + def from_payload(cls, payload: dict[str, Any]) -> "GroupSnapshot": + block_plugin = payload.get("block_plugin", "") or "" + superuser_block_plugin = payload.get("superuser_block_plugin", "") or "" + block_task = payload.get("block_task", "") or "" + superuser_block_task = payload.get("superuser_block_task", "") or "" + return cls( + group_id=str(payload.get("group_id", "")), + channel_id=payload.get("channel_id"), + group_name=payload.get("group_name", "") or "", + max_member_count=int(payload.get("max_member_count", 0) or 0), + member_count=int(payload.get("member_count", 0) or 0), + status=bool(payload.get("status", True)), + level=int(payload.get("level", 0) or 0), + is_super=bool(payload.get("is_super", False)), + group_flag=int(payload.get("group_flag", 0) or 0), + block_plugin=block_plugin, + superuser_block_plugin=superuser_block_plugin, + block_task=block_task, + superuser_block_task=superuser_block_task, + platform=payload.get("platform"), + block_plugin_set=_parse_block_modules(block_plugin), + superuser_block_plugin_set=_parse_block_modules(superuser_block_plugin), + block_task_set=_parse_block_modules(block_task), + superuser_block_task_set=_parse_block_modules(superuser_block_task), + ) + + +@dataclass(frozen=True) +class LevelUserSnapshot: + user_id: str + group_id: str | None + user_level: int + group_flag: int + + @classmethod + def from_model(cls, model) -> "LevelUserSnapshot": + return cls( + user_id=str(model.user_id), + group_id=getattr(model, "group_id", None), + user_level=int(getattr(model, "user_level", 0) or 0), + group_flag=int(getattr(model, "group_flag", 0) or 0), + ) + + def to_payload(self) -> dict[str, Any]: + return { + "user_id": self.user_id, + "group_id": self.group_id, + "user_level": self.user_level, + "group_flag": self.group_flag, + } + + @classmethod + def from_payload(cls, payload: dict[str, Any]) -> "LevelUserSnapshot": + return cls( + user_id=str(payload.get("user_id", "")), + group_id=payload.get("group_id"), + user_level=int(payload.get("user_level", 0) or 0), + group_flag=int(payload.get("group_flag", 0) or 0), + ) + + +@dataclass(frozen=True) +class PluginLimitSnapshot: + id: int + module: str + module_path: str + limit_type: PluginLimitType + watch_type: LimitWatchType + check_type: LimitCheckType + status: bool + result: str | None + cd: int | None + max_count: int | None + + @classmethod + def from_model(cls, model) -> "PluginLimitSnapshot": + return cls( + id=int(model.id), + module=str(model.module), + module_path=str(model.module_path), + limit_type=model.limit_type, + watch_type=model.watch_type, + check_type=model.check_type, + status=bool(model.status), + result=getattr(model, "result", None), + cd=getattr(model, "cd", None), + max_count=getattr(model, "max_count", None), + ) + + def to_payload(self) -> dict[str, Any]: + return { + "id": self.id, + "module": self.module, + "module_path": self.module_path, + "limit_type": self.limit_type.value, + "watch_type": self.watch_type.value, + "check_type": self.check_type.value, + "status": self.status, + "result": self.result, + "cd": self.cd, + "max_count": self.max_count, + } + + @classmethod + def from_payload(cls, payload: dict[str, Any]) -> "PluginLimitSnapshot": + return cls( + id=int(payload.get("id", 0) or 0), + module=str(payload.get("module", "")), + module_path=str(payload.get("module_path", "")), + limit_type=PluginLimitType(payload.get("limit_type", PluginLimitType.CD)), + watch_type=LimitWatchType(payload.get("watch_type", LimitWatchType.USER)), + check_type=LimitCheckType(payload.get("check_type", LimitCheckType.ALL)), + status=bool(payload.get("status", True)), + result=payload.get("result"), + cd=payload.get("cd"), + max_count=payload.get("max_count"), + ) + + +class RuntimeCacheSync: + _redis: ClassVar[Any | None] = None + _pubsub: ClassVar[Any | None] = None + _task: ClassVar[asyncio.Task | None] = None + _publish_tasks: ClassVar[set[asyncio.Task]] = set() + _ready: ClassVar[bool] = False + _channel: ClassVar[str] = "" + + @classmethod + def _sync_enabled(cls) -> bool: + enabled = bool(Config.get_config("hook", "RUNTIME_CACHE_SYNC_ENABLED", True)) + return enabled and _redis_enabled() + + @classmethod + async def start(cls) -> None: + if cls._ready: + return + if not cls._sync_enabled(): + return + try: + import redis.asyncio as redis_async + except ImportError: + logger.warning( + "redis not installed, runtime cache sync disabled", LOG_COMMAND + ) + return + + host = _env_get("REDIS_HOST") + if not host: + return + port = _coerce_int(_env_get("REDIS_PORT"), 6379) + password = _env_get("REDIS_PASSWORD") + cls._channel = str( + Config.get_config( + "hook", "RUNTIME_CACHE_SYNC_CHANNEL", "ZHENXUN_RUNTIME_CACHE_SYNC" + ) + ) + try: + cls._redis = redis_async.Redis( + host=host, + port=port, + password=password, + decode_responses=True, + ) + cls._pubsub = cls._redis.pubsub() + if cls._pubsub is None: + return + await cls._pubsub.subscribe(cls._channel) + cls._task = asyncio.create_task(cls._listen_loop()) + cls._ready = True + logger.info("runtime cache sync enabled", LOG_COMMAND) + except Exception as exc: + logger.error("runtime cache sync init failed", LOG_COMMAND, e=exc) + await cls.stop() + + @classmethod + async def stop(cls) -> None: + if cls._task and not cls._task.done(): + cls._task.cancel() + cls._task = None + try: + if cls._pubsub is not None: + await cls._pubsub.close() + except Exception: + pass + cls._pubsub = None + try: + if cls._redis is not None: + await cls._redis.close() + except Exception: + pass + cls._redis = None + cls._ready = False + + @classmethod + def publish_event(cls, cache_type: str, action: str, data: dict[str, Any]) -> None: + if not cls._ready: + return + payload = { + "source": INSTANCE_ID, + "type": cache_type, + "action": action, + "data": data, + } + task = asyncio.create_task(cls._publish(payload)) + cls._publish_tasks.add(task) + task.add_done_callback(cls._publish_tasks.discard) + + @classmethod + async def _publish(cls, payload: dict[str, Any]) -> None: + if not cls._ready or cls._redis is None: + return + try: + await cls._redis.publish(cls._channel, json.dumps(payload)) + except Exception as exc: + logger.error("runtime cache sync publish failed", LOG_COMMAND, e=exc) + + @classmethod + async def _listen_loop(cls) -> None: + if cls._pubsub is None: + return + try: + while True: + message = await cls._pubsub.get_message( + ignore_subscribe_messages=True, timeout=1.0 + ) + if not message: + await asyncio.sleep(0.05) + continue + await cls._handle_message(message.get("data")) + except asyncio.CancelledError: + return + except Exception as exc: + logger.error("runtime cache sync listener failed", LOG_COMMAND, e=exc) + + @classmethod + async def _handle_message(cls, raw: Any) -> None: + if raw is None: + return + if isinstance(raw, bytes | bytearray): + try: + raw = raw.decode() + except Exception: + return + if not raw: + return + try: + payload = json.loads(raw) + except Exception: + return + if payload.get("source") == INSTANCE_ID: + return + cache_type = payload.get("type") + action = payload.get("action") + data = payload.get("data") or {} + if cache_type == "bot": + await BotMemoryCache.apply_sync_event(action, data) + elif cache_type == "group": + await GroupMemoryCache.apply_sync_event(action, data) + elif cache_type == "ban": + await BanMemoryCache.apply_sync_event(action, data) + elif cache_type == "level": + await LevelUserMemoryCache.apply_sync_event(action, data) + elif cache_type == "plugin_limit": + await PluginLimitMemoryCache.apply_sync_event(action, data) + + +class PluginInfoMemoryCache: + _lock: ClassVar[asyncio.Lock] = asyncio.Lock() + _by_module: ClassVar[dict[str, "PluginInfo"]] = {} + _by_module_path: ClassVar[dict[str, "PluginInfo"]] = {} + _loaded: ClassVar[bool] = False + _refresh_task: ClassVar[asyncio.Task | None] = None + _last_refresh: ClassVar[float] = 0.0 + + @classmethod + async def refresh(cls) -> None: + from zhenxun.models.plugin_info import PluginInfo + + async with cls._lock: + plugins = await PluginInfo.all() + by_module: dict[str, "PluginInfo"] = {} + by_module_path: dict[str, "PluginInfo"] = {} + for plugin in plugins: + if plugin.module: + by_module[plugin.module] = plugin + if plugin.module_path: + by_module_path[plugin.module_path] = plugin + cls._by_module = by_module + cls._by_module_path = by_module_path + cls._loaded = True + cls._last_refresh = time.time() + logger.debug( + f"plugin cache refreshed: {len(by_module)} entries", LOG_COMMAND + ) + + @classmethod + async def ensure_loaded(cls) -> None: + if cls._loaded: + return + await cls.refresh() + + @classmethod + async def get_by_module(cls, module: str) -> "PluginInfo | None": + if not cls._loaded: + await cls.ensure_loaded() + return cls._by_module.get(module) + + @classmethod + def get_by_module_path(cls, module_path: str) -> "PluginInfo | None": + return cls._by_module_path.get(module_path) + + @classmethod + def set_plugin(cls, plugin) -> None: + if not plugin: + return + if plugin.module: + cls._by_module[plugin.module] = plugin + if getattr(plugin, "module_path", None): + cls._by_module_path[plugin.module_path] = plugin + + @classmethod + def remove_by_module(cls, module: str) -> None: + cls._by_module.pop(module, None) + + @classmethod + async def _refresh_loop(cls, interval: int) -> None: + while True: + await asyncio.sleep(interval) + try: + await cls.refresh() + except Exception as exc: + logger.error("plugin cache refresh failed", LOG_COMMAND, e=exc) + + @classmethod + def start_refresh_task(cls) -> None: + interval = _coerce_int( + Config.get_config("hook", "PLUGININFO_MEM_REFRESH_INTERVAL", 300), + 300, + ) + if interval <= 0: + return + if cls._refresh_task and not cls._refresh_task.done(): + return + cls._refresh_task = asyncio.create_task(cls._refresh_loop(interval)) + + @classmethod + def stop_tasks(cls) -> None: + if cls._refresh_task and not cls._refresh_task.done(): + cls._refresh_task.cancel() + cls._refresh_task = None + + +class BotMemoryCache: + _lock: ClassVar[asyncio.Lock] = asyncio.Lock() + _by_id: ClassVar[dict[str, BotSnapshot]] = {} + _negative: ClassVar[dict[str, float]] = {} + _loaded: ClassVar[bool] = False + _refresh_task: ClassVar[asyncio.Task | None] = None + + @classmethod + def _normalize(cls, bot_id: str | None) -> str | None: + if bot_id is None: + return None + bot_id = bot_id.strip() + return bot_id if bot_id else None + + @classmethod + def _negative_ttl(cls) -> int: + return _coerce_int(Config.get_config("hook", "BOT_MEM_NEGATIVE_TTL", 60), 60) + + @classmethod + def _is_negative(cls, bot_id: str) -> bool: + expire_at = cls._negative.get(bot_id) + if not expire_at: + return False + if expire_at <= time.time(): + cls._negative.pop(bot_id, None) + return False + return True + + @classmethod + def _mark_negative(cls, bot_id: str) -> None: + ttl = cls._negative_ttl() + if ttl <= 0: + return + cls._negative[bot_id] = time.time() + ttl + + @classmethod + async def refresh(cls) -> None: + from zhenxun.models.bot_console import BotConsole + + async with cls._lock: + records = await BotConsole.all() + cls._by_id = {str(r.bot_id): BotSnapshot.from_model(r) for r in records} + cls._negative = {} + cls._loaded = True + logger.debug(f"bot cache refreshed: {len(cls._by_id)} entries", LOG_COMMAND) + + @classmethod + async def ensure_loaded(cls) -> None: + if cls._loaded: + return + await cls.refresh() + + @classmethod + async def get(cls, bot_id: str | None) -> BotSnapshot | None: + bot_id = cls._normalize(bot_id) + if not bot_id: + return None + if not cls._loaded: + await cls.ensure_loaded() + entry = cls._by_id.get(bot_id) + if entry: + return entry + if cls._is_negative(bot_id): + return None + cls._mark_negative(bot_id) + return None + + @classmethod + async def update_status(cls, bot_id: str | None, status: bool) -> None: + bot_id = cls._normalize(bot_id) + if not bot_id: + return + async with cls._lock: + entry = cls._by_id.get(bot_id) + if not entry: + return + updated = BotSnapshot( + bot_id=entry.bot_id, + status=bool(status), + platform=entry.platform, + block_plugins=entry.block_plugins, + block_tasks=entry.block_tasks, + available_plugins=entry.available_plugins, + available_tasks=entry.available_tasks, + ) + cls._by_id[bot_id] = updated + RuntimeCacheSync.publish_event("bot", "upsert", updated.to_payload()) + + @classmethod + async def upsert_from_model(cls, record) -> None: + entry = BotSnapshot.from_model(record) + async with cls._lock: + cls._by_id[entry.bot_id] = entry + cls._negative.pop(entry.bot_id, None) + RuntimeCacheSync.publish_event("bot", "upsert", entry.to_payload()) + + @classmethod + async def upsert_from_payload(cls, payload: dict[str, Any]) -> None: + entry = BotSnapshot.from_payload(payload) + if not entry.bot_id: + return + async with cls._lock: + cls._by_id[entry.bot_id] = entry + cls._negative.pop(entry.bot_id, None) + + @classmethod + async def remove(cls, bot_id: str | None) -> None: + bot_id = cls._normalize(bot_id) + if not bot_id: + return + async with cls._lock: + cls._by_id.pop(bot_id, None) + RuntimeCacheSync.publish_event("bot", "delete", {"bot_id": bot_id}) + + @classmethod + async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None: + if action == "upsert": + await cls.upsert_from_payload(data) + elif action == "delete": + await cls.remove(data.get("bot_id")) + elif action == "refresh": + await cls.refresh() + + @classmethod + async def _refresh_loop(cls, interval: int) -> None: + while True: + await asyncio.sleep(interval) + try: + await cls.refresh() + except Exception as exc: + logger.error("bot cache refresh failed", LOG_COMMAND, e=exc) + + @classmethod + def start_tasks(cls) -> None: + interval = _coerce_int( + Config.get_config("hook", "BOT_MEM_REFRESH_INTERVAL", 60), 60 + ) + if interval <= 0: + return + if cls._refresh_task and not cls._refresh_task.done(): + return + cls._refresh_task = asyncio.create_task(cls._refresh_loop(interval)) + + @classmethod + def stop_tasks(cls) -> None: + if cls._refresh_task and not cls._refresh_task.done(): + cls._refresh_task.cancel() + cls._refresh_task = None + + +class GroupMemoryCache: + _lock: ClassVar[asyncio.Lock] = asyncio.Lock() + _by_key: ClassVar[dict[tuple[str, str], GroupSnapshot]] = {} + _negative: ClassVar[dict[tuple[str, str], float]] = {} + _loaded: ClassVar[bool] = False + _refresh_task: ClassVar[asyncio.Task | None] = None + + @classmethod + def _normalize(cls, value: str | None) -> str | None: + if value is None: + return None + value = value.strip() + return value if value else None + + @classmethod + def _key( + cls, group_id: str | None, channel_id: str | None + ) -> tuple[str, str] | None: + group_id = cls._normalize(group_id) + if not group_id: + return None + channel_id = cls._normalize(channel_id) or "" + return (group_id, channel_id) + + @classmethod + def _negative_ttl(cls) -> int: + return _coerce_int(Config.get_config("hook", "GROUP_MEM_NEGATIVE_TTL", 60), 60) + + @classmethod + def _is_negative(cls, key: tuple[str, str]) -> bool: + expire_at = cls._negative.get(key) + if not expire_at: + return False + if expire_at <= time.time(): + cls._negative.pop(key, None) + return False + return True + + @classmethod + def _mark_negative(cls, key: tuple[str, str]) -> None: + ttl = cls._negative_ttl() + if ttl <= 0: + return + cls._negative[key] = time.time() + ttl + + @classmethod + async def refresh(cls) -> None: + from zhenxun.models.group_console import GroupConsole + + async with cls._lock: + records = await GroupConsole.all() + by_key: dict[tuple[str, str], GroupSnapshot] = {} + for record in records: + entry = GroupSnapshot.from_model(record) + key = cls._key(entry.group_id, entry.channel_id) + if key: + by_key[key] = entry + cls._by_key = by_key + cls._negative = {} + cls._loaded = True + logger.debug(f"group cache refreshed: {len(by_key)} entries", LOG_COMMAND) + + @classmethod + async def ensure_loaded(cls) -> None: + if cls._loaded: + return + await cls.refresh() + + @classmethod + def is_loaded(cls) -> bool: + return cls._loaded + + @classmethod + async def get( + cls, group_id: str | None, channel_id: str | None = None + ) -> GroupSnapshot | None: + key = cls._key(group_id, channel_id) + if not key: + return None + if not cls._loaded: + await cls.ensure_loaded() + entry = cls._by_key.get(key) + if entry: + return entry + if cls._is_negative(key): + return None + cls._mark_negative(key) + return None + + @classmethod + def get_if_ready( + cls, group_id: str | None, channel_id: str | None = None + ) -> GroupSnapshot | None: + key = cls._key(group_id, channel_id) + if not key: + return None + if not cls._loaded: + return None + entry = cls._by_key.get(key) + if entry: + return entry + if cls._is_negative(key): + return None + cls._mark_negative(key) + return None + + @classmethod + async def upsert_from_model(cls, record) -> None: + entry = GroupSnapshot.from_model(record) + key = cls._key(entry.group_id, entry.channel_id) + if not key: + return + async with cls._lock: + cls._by_key[key] = entry + cls._negative.pop(key, None) + RuntimeCacheSync.publish_event("group", "upsert", entry.to_payload()) + + @classmethod + async def upsert_from_payload(cls, payload: dict[str, Any]) -> None: + entry = GroupSnapshot.from_payload(payload) + key = cls._key(entry.group_id, entry.channel_id) + if not key: + return + async with cls._lock: + cls._by_key[key] = entry + cls._negative.pop(key, None) + + @classmethod + async def remove(cls, group_id: str | None, channel_id: str | None = None) -> None: + key = cls._key(group_id, channel_id) + if not key: + return + async with cls._lock: + cls._by_key.pop(key, None) + RuntimeCacheSync.publish_event( + "group", "delete", {"group_id": key[0], "channel_id": key[1] or None} + ) + + @classmethod + async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None: + if action == "upsert": + await cls.upsert_from_payload(data) + elif action == "delete": + await cls.remove(data.get("group_id"), data.get("channel_id")) + elif action == "refresh": + await cls.refresh() + + @classmethod + async def _refresh_loop(cls, interval: int) -> None: + while True: + await asyncio.sleep(interval) + try: + await cls.refresh() + except Exception as exc: + logger.error("group cache refresh failed", LOG_COMMAND, e=exc) + + @classmethod + def start_tasks(cls) -> None: + interval = _coerce_int( + Config.get_config("hook", "GROUP_MEM_REFRESH_INTERVAL", 60), 60 + ) + if interval <= 0: + return + if cls._refresh_task and not cls._refresh_task.done(): + return + cls._refresh_task = asyncio.create_task(cls._refresh_loop(interval)) + + @classmethod + def stop_tasks(cls) -> None: + if cls._refresh_task and not cls._refresh_task.done(): + cls._refresh_task.cancel() + cls._refresh_task = None + + +class LevelUserMemoryCache: + _lock: ClassVar[asyncio.Lock] = asyncio.Lock() + _by_key: ClassVar[dict[tuple[str, str], LevelUserSnapshot]] = {} + _negative: ClassVar[dict[tuple[str, str], float]] = {} + _loaded: ClassVar[bool] = False + _refresh_task: ClassVar[asyncio.Task | None] = None + _last_refresh: ClassVar[float] = 0.0 + + @classmethod + def _normalize(cls, value: str | None) -> str | None: + if value is None: + return None + value = value.strip() + return value if value else "" + + @classmethod + def _key(cls, user_id: str | None, group_id: str | None) -> tuple[str, str] | None: + user_id = cls._normalize(user_id) + if not user_id: + return None + group_id = cls._normalize(group_id) or "" + return (user_id, group_id) + + @classmethod + def _negative_ttl(cls) -> int: + return _coerce_int(Config.get_config("hook", "LEVEL_MEM_NEGATIVE_TTL", 60), 60) + + @classmethod + def _is_negative(cls, key: tuple[str, str]) -> bool: + expire_at = cls._negative.get(key) + if not expire_at: + return False + if expire_at <= time.time(): + cls._negative.pop(key, None) + return False + return True + + @classmethod + def _mark_negative(cls, key: tuple[str, str]) -> None: + ttl = cls._negative_ttl() + if ttl <= 0: + return + cls._negative[key] = time.time() + ttl + + @classmethod + async def refresh(cls) -> None: + from zhenxun.models.level_user import LevelUser + + async with cls._lock: + records = await LevelUser.all() + by_key: dict[tuple[str, str], LevelUserSnapshot] = {} + for record in records: + entry = LevelUserSnapshot.from_model(record) + key = cls._key(entry.user_id, entry.group_id) + if key: + by_key[key] = entry + cls._by_key = by_key + cls._negative = {} + cls._loaded = True + cls._last_refresh = time.time() + logger.debug(f"level cache refreshed: {len(by_key)} entries", LOG_COMMAND) + + @classmethod + async def ensure_loaded(cls) -> None: + if cls._loaded: + return + await cls.refresh() + + @classmethod + async def ensure_fresh(cls) -> None: + interval = _coerce_int( + Config.get_config("hook", "LEVEL_MEM_REFRESH_INTERVAL", 120), 120 + ) + if not cls._loaded: + await cls.refresh() + return + if interval <= 0: + return + if time.time() - cls._last_refresh > interval: + await cls.refresh() + + @classmethod + async def get( + cls, user_id: str | None, group_id: str | None + ) -> LevelUserSnapshot | None: + key = cls._key(user_id, group_id) + if not key: + return None + if not cls._loaded: + await cls.ensure_loaded() + entry = cls._by_key.get(key) + if entry: + return entry + if cls._is_negative(key): + return None + cls._mark_negative(key) + return None + + @classmethod + async def get_levels( + cls, user_id: str | None, group_id: str | None + ) -> tuple[LevelUserSnapshot | None, LevelUserSnapshot | None]: + if not cls._loaded: + await cls.ensure_loaded() + global_user = None + group_user = None + global_key = cls._key(user_id, "") + if global_key: + global_user = cls._by_key.get(global_key) + if group_id: + group_key = cls._key(user_id, group_id) + if group_key: + group_user = cls._by_key.get(group_key) + return global_user, group_user + + @classmethod + async def upsert_from_model(cls, record) -> None: + entry = LevelUserSnapshot.from_model(record) + key = cls._key(entry.user_id, entry.group_id) + if not key: + return + async with cls._lock: + cls._by_key[key] = entry + cls._negative.pop(key, None) + RuntimeCacheSync.publish_event("level", "upsert", entry.to_payload()) + + @classmethod + async def upsert_from_payload(cls, payload: dict[str, Any]) -> None: + entry = LevelUserSnapshot.from_payload(payload) + key = cls._key(entry.user_id, entry.group_id) + if not key: + return + async with cls._lock: + cls._by_key[key] = entry + cls._negative.pop(key, None) + + @classmethod + async def remove(cls, user_id: str | None, group_id: str | None) -> None: + key = cls._key(user_id, group_id) + if not key: + return + async with cls._lock: + cls._by_key.pop(key, None) + RuntimeCacheSync.publish_event( + "level", "delete", {"user_id": key[0], "group_id": key[1] or None} + ) + + @classmethod + async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None: + if action == "upsert": + await cls.upsert_from_payload(data) + elif action == "delete": + await cls.remove(data.get("user_id"), data.get("group_id")) + elif action == "refresh": + await cls.refresh() + + @classmethod + async def _refresh_loop(cls, interval: int) -> None: + while True: + await asyncio.sleep(interval) + try: + await cls.refresh() + except Exception as exc: + logger.error("level cache refresh failed", LOG_COMMAND, e=exc) + + @classmethod + def start_tasks(cls) -> None: + interval = _coerce_int( + Config.get_config("hook", "LEVEL_MEM_REFRESH_INTERVAL", 120), 120 + ) + if interval <= 0: + return + if cls._refresh_task and not cls._refresh_task.done(): + return + cls._refresh_task = asyncio.create_task(cls._refresh_loop(interval)) + + @classmethod + def stop_tasks(cls) -> None: + if cls._refresh_task and not cls._refresh_task.done(): + cls._refresh_task.cancel() + cls._refresh_task = None + + +class PluginLimitMemoryCache: + _lock: ClassVar[asyncio.Lock] = asyncio.Lock() + _by_id: ClassVar[dict[int, PluginLimitSnapshot]] = {} + _by_module: ClassVar[dict[str, list[PluginLimitSnapshot]]] = {} + _negative: ClassVar[dict[str, float]] = {} + _loaded: ClassVar[bool] = False + _refresh_task: ClassVar[asyncio.Task | None] = None + + @classmethod + def _normalize(cls, value: str | None) -> str | None: + if value is None: + return None + value = value.strip() + return value if value else None + + @classmethod + def _negative_ttl(cls) -> int: + return _coerce_int(Config.get_config("hook", "LIMIT_MEM_NEGATIVE_TTL", 30), 30) + + @classmethod + def _is_negative(cls, module: str) -> bool: + expire_at = cls._negative.get(module) + if not expire_at: + return False + if expire_at <= time.time(): + cls._negative.pop(module, None) + return False + return True + + @classmethod + def _mark_negative(cls, module: str) -> None: + ttl = cls._negative_ttl() + if ttl <= 0: + return + cls._negative[module] = time.time() + ttl + + @classmethod + async def refresh(cls) -> None: + from zhenxun.models.plugin_limit import PluginLimit + + async with cls._lock: + records = await PluginLimit.filter(status=True).all() + by_id: dict[int, PluginLimitSnapshot] = {} + by_module: dict[str, list[PluginLimitSnapshot]] = {} + for record in records: + entry = PluginLimitSnapshot.from_model(record) + by_id[entry.id] = entry + by_module.setdefault(entry.module, []).append(entry) + cls._by_id = by_id + cls._by_module = by_module + cls._negative = {} + cls._loaded = True + logger.debug( + f"plugin limit cache refreshed: {len(by_id)} entries", + LOG_COMMAND, + ) + + @classmethod + async def ensure_loaded(cls) -> None: + if cls._loaded: + return + await cls.refresh() + + @classmethod + async def get_limits(cls, module: str) -> list[PluginLimitSnapshot]: + normalized = cls._normalize(module) + if not normalized: + return [] + module = normalized + if not cls._loaded: + await cls.ensure_loaded() + limits = cls._by_module.get(module) + if limits is not None: + return limits + if cls._is_negative(module): + return [] + cls._mark_negative(module) + return [] + + @classmethod + def get_all_limits(cls) -> list[PluginLimitSnapshot]: + return list(cls._by_id.values()) + + @classmethod + async def upsert_from_model(cls, record) -> None: + entry = PluginLimitSnapshot.from_model(record) + await cls._upsert_entry(entry) + RuntimeCacheSync.publish_event("plugin_limit", "upsert", entry.to_payload()) + + @classmethod + async def upsert_from_payload(cls, payload: dict[str, Any]) -> None: + try: + entry = PluginLimitSnapshot.from_payload(payload) + except Exception: + return + await cls._upsert_entry(entry) + + @classmethod + async def _upsert_entry(cls, entry: PluginLimitSnapshot) -> None: + async with cls._lock: + prev = cls._by_id.get(entry.id) + if prev and prev.module != entry.module: + cls._by_module[prev.module] = [ + item + for item in cls._by_module.get(prev.module, []) + if item.id != prev.id + ] + if not entry.status: + cls._by_id.pop(entry.id, None) + cls._by_module[entry.module] = [ + item + for item in cls._by_module.get(entry.module, []) + if item.id != entry.id + ] + return + cls._by_id[entry.id] = entry + module_limits = [ + item + for item in cls._by_module.get(entry.module, []) + if item.id != entry.id + ] + module_limits.append(entry) + cls._by_module[entry.module] = module_limits + cls._negative.pop(entry.module, None) + + @classmethod + async def remove_by_id(cls, limit_id: int | None) -> None: + if not limit_id: + return + async with cls._lock: + entry = cls._by_id.pop(limit_id, None) + if entry: + cls._by_module[entry.module] = [ + item + for item in cls._by_module.get(entry.module, []) + if item.id != entry.id + ] + RuntimeCacheSync.publish_event("plugin_limit", "delete", {"id": int(limit_id)}) + + @classmethod + async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None: + if action == "upsert": + await cls.upsert_from_payload(data) + elif action == "delete": + await cls.remove_by_id(data.get("id")) + elif action == "refresh": + await cls.refresh() + + @classmethod + async def _refresh_loop(cls, interval: int) -> None: + while True: + await asyncio.sleep(interval) + try: + await cls.refresh() + except Exception as exc: + logger.error("plugin limit cache refresh failed", LOG_COMMAND, e=exc) + + @classmethod + def start_tasks(cls) -> None: + interval = _coerce_int( + Config.get_config("hook", "LIMIT_MEM_REFRESH_INTERVAL", 60), 60 + ) + if interval <= 0: + return + if cls._refresh_task and not cls._refresh_task.done(): + return + cls._refresh_task = asyncio.create_task(cls._refresh_loop(interval)) + + @classmethod + def stop_tasks(cls) -> None: + if cls._refresh_task and not cls._refresh_task.done(): + cls._refresh_task.cancel() + cls._refresh_task = None + + +class BanMemoryCache: + _lock: ClassVar[asyncio.Lock] = asyncio.Lock() + _by_user: ClassVar[dict[str, BanEntry]] = {} + _by_group: ClassVar[dict[str, BanEntry]] = {} + _by_user_group: ClassVar[dict[tuple[str, str], BanEntry]] = {} + _negative: ClassVar[dict[tuple[str | None, str | None], float]] = {} + _loaded: ClassVar[bool] = False + _refresh_task: ClassVar[asyncio.Task | None] = None + _cleanup_task: ClassVar[asyncio.Task | None] = None + _remove_tasks: ClassVar[set[asyncio.Task]] = set() + + @classmethod + def _normalize_id(cls, value: str | None) -> str | None: + if value is None: + return None + value = value.strip() + return value if value else None + + @classmethod + def _neg_ttl(cls) -> int: + return _coerce_int(Config.get_config("hook", "BAN_MEM_NEGATIVE_TTL", 5), 5) + + @classmethod + def _neg_key( + cls, user_id: str | None, group_id: str | None + ) -> tuple[str | None, str | None]: + return (cls._normalize_id(user_id), cls._normalize_id(group_id)) + + @classmethod + def _is_negative(cls, key: tuple[str | None, str | None]) -> bool: + expire_at = cls._negative.get(key) + if not expire_at: + return False + if expire_at <= time.time(): + cls._negative.pop(key, None) + return False + return True + + @classmethod + def _mark_negative(cls, key: tuple[str | None, str | None]) -> None: + ttl = cls._neg_ttl() + if ttl <= 0: + return + cls._negative[key] = time.time() + ttl + + @classmethod + def _build_entry(cls, record) -> BanEntry | None: + user_id = cls._normalize_id(record.user_id) + group_id = cls._normalize_id(record.group_id) + duration = int(record.duration) + if duration == -1: + expire_at = None + else: + expire_at = float(record.ban_time + duration) + return BanEntry( + user_id=user_id, + group_id=group_id, + ban_level=int(record.ban_level), + ban_time=int(record.ban_time), + duration=duration, + expire_at=expire_at, + ) + + @classmethod + async def refresh(cls) -> None: + from zhenxun.models.ban_console import BanConsole + + async with cls._lock: + now_ts = time.time() + records = await BanConsole.all() + by_user: dict[str, BanEntry] = {} + by_group: dict[str, BanEntry] = {} + by_user_group: dict[tuple[str, str], BanEntry] = {} + for record in records: + entry = cls._build_entry(record) + if not entry: + continue + if entry.expire_at is not None and entry.expire_at <= now_ts: + continue + if entry.user_id and entry.group_id: + by_user_group[(entry.user_id, entry.group_id)] = entry + elif entry.user_id: + by_user[entry.user_id] = entry + elif entry.group_id: + by_group[entry.group_id] = entry + cls._by_user = by_user + cls._by_group = by_group + cls._by_user_group = by_user_group + cls._negative = {} + cls._loaded = True + logger.debug( + "ban cache refreshed: " + f"user={len(by_user)} group={len(by_group)} " + f"user_group={len(by_user_group)}", + LOG_COMMAND, + ) + + @classmethod + async def ensure_loaded(cls) -> None: + if cls._loaded: + return + await cls.refresh() + + @classmethod + def is_loaded(cls) -> bool: + return cls._loaded + + @classmethod + async def upsert_from_model(cls, record) -> None: + entry = cls._build_entry(record) + if not entry: + return + async with cls._lock: + if entry.user_id and entry.group_id: + cls._by_user_group[(entry.user_id, entry.group_id)] = entry + elif entry.user_id: + cls._by_user[entry.user_id] = entry + elif entry.group_id: + cls._by_group[entry.group_id] = entry + cls._negative = {} + RuntimeCacheSync.publish_event("ban", "upsert", entry.to_payload()) + + @classmethod + async def remove(cls, user_id: str | None, group_id: str | None) -> None: + await cls._remove_local(user_id, group_id) + RuntimeCacheSync.publish_event( + "ban", "delete", {"user_id": user_id, "group_id": group_id} + ) + + @classmethod + async def _remove_local(cls, user_id: str | None, group_id: str | None) -> None: + user_id = cls._normalize_id(user_id) + group_id = cls._normalize_id(group_id) + async with cls._lock: + if user_id and group_id: + cls._by_user_group.pop((user_id, group_id), None) + elif user_id: + cls._by_user.pop(user_id, None) + elif group_id: + cls._by_group.pop(group_id, None) + cls._negative = {} + + @classmethod + def _get_entry(cls, user_id: str | None, group_id: str | None) -> BanEntry | None: + user_id = cls._normalize_id(user_id) + group_id = cls._normalize_id(group_id) + if user_id and group_id: + entry = cls._by_user_group.get((user_id, group_id)) + if entry: + return entry + entry = cls._by_user.get(user_id) + if entry: + return entry + return None + if user_id: + return cls._by_user.get(user_id) + if group_id: + return cls._by_group.get(group_id) + return None + + @classmethod + def is_banned(cls, user_id: str | None, group_id: str | None) -> bool: + if not cls._loaded: + return False + neg_key = cls._neg_key(user_id, group_id) + if cls._is_negative(neg_key): + return False + entry = cls._get_entry(user_id, group_id) + if not entry: + cls._mark_negative(neg_key) + return False + remaining = entry.remaining() + if remaining == 0 and entry.duration != -1: + task = asyncio.create_task(cls.remove(entry.user_id, entry.group_id)) + cls._remove_tasks.add(task) + task.add_done_callback(cls._remove_tasks.discard) + return False + return True + + @classmethod + def remaining_time(cls, user_id: str | None, group_id: str | None) -> int: + if not cls._loaded: + return 0 + neg_key = cls._neg_key(user_id, group_id) + if cls._is_negative(neg_key): + return 0 + entry = cls._get_entry(user_id, group_id) + if not entry: + cls._mark_negative(neg_key) + return 0 + remaining = entry.remaining() + if remaining == 0 and entry.duration != -1: + task = asyncio.create_task(cls.remove(entry.user_id, entry.group_id)) + cls._remove_tasks.add(task) + task.add_done_callback(cls._remove_tasks.discard) + return 0 + return remaining + + @classmethod + def check_ban_level( + cls, user_id: str | None, group_id: str | None, level: int + ) -> bool: + if not cls._loaded: + return False + neg_key = cls._neg_key(user_id, group_id) + if cls._is_negative(neg_key): + return False + entry = cls._get_entry(user_id, group_id) + if not entry: + cls._mark_negative(neg_key) + return False + remaining = entry.remaining() + if remaining == 0 and entry.duration != -1: + task = asyncio.create_task(cls.remove(entry.user_id, entry.group_id)) + cls._remove_tasks.add(task) + task.add_done_callback(cls._remove_tasks.discard) + return False + return entry.ban_level <= level + + @classmethod + async def cleanup_expired(cls, delete_db: bool = True) -> None: + now_ts = time.time() + expired: list[BanEntry] = [] + async with cls._lock: + for entry in list(cls._by_user.values()): + if entry.expire_at is not None and entry.expire_at <= now_ts: + expired.append(entry) + for entry in list(cls._by_group.values()): + if entry.expire_at is not None and entry.expire_at <= now_ts: + expired.append(entry) + for entry in list(cls._by_user_group.values()): + if entry.expire_at is not None and entry.expire_at <= now_ts: + expired.append(entry) + for entry in expired: + if entry.user_id and entry.group_id: + cls._by_user_group.pop((entry.user_id, entry.group_id), None) + elif entry.user_id: + cls._by_user.pop(entry.user_id, None) + elif entry.group_id: + cls._by_group.pop(entry.group_id, None) + if expired: + cls._negative = {} + if not delete_db or not expired: + return + from tortoise.expressions import Q + + from zhenxun.models.ban_console import BanConsole + + for entry in expired: + query = BanConsole.filter() + if entry.user_id: + query = query.filter(user_id=entry.user_id) + else: + query = query.filter(Q(user_id__isnull=True) | Q(user_id="")) + if entry.group_id: + query = query.filter(group_id=entry.group_id) + else: + query = query.filter(Q(group_id__isnull=True) | Q(group_id="")) + await query.delete() + + @classmethod + async def _refresh_loop(cls, interval: int) -> None: + while True: + await asyncio.sleep(interval) + try: + await cls.refresh() + except Exception as exc: + logger.error("ban cache refresh failed", LOG_COMMAND, e=exc) + + @classmethod + async def _cleanup_loop(cls, interval: int, delete_db: bool) -> None: + while True: + await asyncio.sleep(interval) + try: + await cls.cleanup_expired(delete_db=delete_db) + except Exception as exc: + logger.error("ban cache cleanup failed", LOG_COMMAND, e=exc) + + @classmethod + def start_tasks(cls) -> None: + refresh_interval = _coerce_int( + Config.get_config("hook", "BAN_MEM_REFRESH_INTERVAL", 60), 60 + ) + clean_interval = _coerce_int( + Config.get_config("hook", "BAN_MEM_CLEAN_INTERVAL", 60), 60 + ) + cleanup_db = bool(Config.get_config("hook", "BAN_MEM_CLEANUP_DB", True)) + + if refresh_interval > 0 and (not cls._refresh_task or cls._refresh_task.done()): + cls._refresh_task = asyncio.create_task(cls._refresh_loop(refresh_interval)) + if clean_interval > 0 and (not cls._cleanup_task or cls._cleanup_task.done()): + cls._cleanup_task = asyncio.create_task( + cls._cleanup_loop(clean_interval, cleanup_db) + ) + + @classmethod + def stop_tasks(cls) -> None: + if cls._refresh_task and not cls._refresh_task.done(): + cls._refresh_task.cancel() + if cls._cleanup_task and not cls._cleanup_task.done(): + cls._cleanup_task.cancel() + cls._refresh_task = None + cls._cleanup_task = None + + @classmethod + async def upsert_from_payload(cls, payload: dict[str, Any]) -> None: + entry = BanEntry.from_payload(payload) + async with cls._lock: + if entry.user_id and entry.group_id: + cls._by_user_group[(entry.user_id, entry.group_id)] = entry + elif entry.user_id: + cls._by_user[entry.user_id] = entry + elif entry.group_id: + cls._by_group[entry.group_id] = entry + cls._negative = {} + + @classmethod + async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None: + if action == "upsert": + await cls.upsert_from_payload(data) + elif action == "delete": + await cls._remove_local(data.get("user_id"), data.get("group_id")) + elif action == "refresh": + await cls.refresh() + + +@PriorityLifecycle.on_startup(priority=6) +async def _init_runtime_cache(): + await RuntimeCacheSync.start() + try: + await PluginInfoMemoryCache.refresh() + except Exception as exc: + logger.error("plugin cache init failed", LOG_COMMAND, e=exc) + try: + await BotMemoryCache.refresh() + except Exception as exc: + logger.error("bot cache init failed", LOG_COMMAND, e=exc) + try: + await GroupMemoryCache.refresh() + except Exception as exc: + logger.error("group cache init failed", LOG_COMMAND, e=exc) + try: + await LevelUserMemoryCache.refresh() + except Exception as exc: + logger.error("level cache init failed", LOG_COMMAND, e=exc) + try: + await PluginLimitMemoryCache.refresh() + except Exception as exc: + logger.error("plugin limit cache init failed", LOG_COMMAND, e=exc) + try: + await BanMemoryCache.refresh() + except Exception as exc: + logger.error("ban cache init failed", LOG_COMMAND, e=exc) + PluginInfoMemoryCache.start_refresh_task() + BotMemoryCache.start_tasks() + GroupMemoryCache.start_tasks() + LevelUserMemoryCache.start_tasks() + PluginLimitMemoryCache.start_tasks() + BanMemoryCache.start_tasks() + _CACHE_READY_EVENT.set() + + +@PriorityLifecycle.on_shutdown(priority=6) +async def _stop_runtime_cache(): + PluginInfoMemoryCache.stop_tasks() + BotMemoryCache.stop_tasks() + GroupMemoryCache.stop_tasks() + LevelUserMemoryCache.stop_tasks() + PluginLimitMemoryCache.stop_tasks() + BanMemoryCache.stop_tasks() + await RuntimeCacheSync.stop() diff --git a/zhenxun/services/db_context/base_model.py b/zhenxun/services/db_context/base_model.py index ff642258..07c63ced 100644 --- a/zhenxun/services/db_context/base_model.py +++ b/zhenxun/services/db_context/base_model.py @@ -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 diff --git a/zhenxun/services/llm/tools.py b/zhenxun/services/llm/tools.py index bbf1b9ed..a0d79ee6 100644 --- a/zhenxun/services/llm/tools.py +++ b/zhenxun/services/llm/tools.py @@ -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: diff --git a/zhenxun/services/message_load.py b/zhenxun/services/message_load.py new file mode 100644 index 00000000..6823ed5b --- /dev/null +++ b/zhenxun/services/message_load.py @@ -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() diff --git a/zhenxun/services/scheduler/engine.py b/zhenxun/services/scheduler/engine.py index 7f711b9e..c5652514 100644 --- a/zhenxun/services/scheduler/engine.py +++ b/zhenxun/services/scheduler/engine.py @@ -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) diff --git a/zhenxun/services/send_queue.py b/zhenxun/services/send_queue.py new file mode 100644 index 00000000..3f9c7458 --- /dev/null +++ b/zhenxun/services/send_queue.py @@ -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)))