From 8afc8f8673b0943bc91bbc69e3987ba8d66a640a Mon Sep 17 00:00:00 2001 From: Copaan <98086483+Copaan@users.noreply.github.com> Date: Sun, 7 Jun 2026 18:14:01 +0800 Subject: [PATCH] =?UTF-8?q?bugfix:=E6=9B=B4=E6=8D=A2=E6=95=B0=E6=8D=AE?= =?UTF-8?q?=E5=BA=93=E5=88=9D=E5=A7=8B=E5=8C=96=E8=B6=85=E6=97=B6=E8=B7=AF?= =?UTF-8?q?=E5=BE=84=E4=BB=A5=E4=BF=AE=E5=A4=8D=E8=BF=9E=E6=8E=A5=E8=B6=85?= =?UTF-8?q?=E6=97=B6=E9=97=AE=E9=A2=98=20(#2137)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * bugfix:更换数据库初始化超时路径以修复连接超时问题 * 移除部分观测链路 * 细节修改 * 权限检查细节修改2 * bugfix:修复金币懒加载造成插件金币消耗不了的问题 * bugfix:整理鉴权逻辑 * 完善缓存系统 * 优化官端使用 * bugfix:增加官机和频道的前导统一剥离,新增获取群列表自愈 * bugfix:修复导入问题 --- .env.example | 19 + .../admin/group_member_update/_data_source.py | 2 +- zhenxun/builtin_plugins/admin/group_update.py | 5 + .../admin/plugin_switch/data_source.py | 2 +- .../admin/plugin_switch/strategy.py | 2 + zhenxun/builtin_plugins/catchphrase.py | 3 + zhenxun/builtin_plugins/hooks/__init__.py | 10 - .../builtin_plugins/hooks/auth/auth_admin.py | 7 +- .../builtin_plugins/hooks/auth/auth_ban.py | 7 +- .../builtin_plugins/hooks/auth/auth_bot.py | 5 +- .../builtin_plugins/hooks/auth/auth_limit.py | 41 +- zhenxun/builtin_plugins/hooks/auth/context.py | 20 +- .../hooks/auth/data_provider.py | 151 ++++ .../builtin_plugins/hooks/auth_activation.py | 669 +++++++++++++----- zhenxun/builtin_plugins/hooks/auth_checker.py | 310 ++------ .../hooks/auth_event_selector.py | 85 ++- zhenxun/builtin_plugins/hooks/auth_hook.py | 5 - .../builtin_plugins/hooks/auth_pipeline.py | 42 +- zhenxun/builtin_plugins/hooks/auth_policy.py | 3 - zhenxun/builtin_plugins/hooks/auth_profile.py | 13 +- zhenxun/builtin_plugins/hooks/auth_route.py | 60 -- .../builtin_plugins/hooks/auth_snapshot.py | 154 +++- zhenxun/builtin_plugins/hooks/call_hook.py | 40 +- zhenxun/builtin_plugins/init/__init__.py | 9 +- zhenxun/builtin_plugins/init/init_task.py | 10 +- zhenxun/builtin_plugins/platform/__init__.py | 16 +- .../platform/qq/group_handle/__init__.py | 4 +- .../platform/qq/group_handle/data_source.py | 3 +- .../platform/qq_api/ug_watch.py | 35 +- zhenxun/builtin_plugins/record_request.py | 3 +- .../scheduler/auto_update_group.py | 4 + .../builtin_plugins/scheduler/chat_check.py | 4 +- .../builtin_plugins/superuser/group_manage.py | 6 +- .../superuser/update_fg_info.py | 10 + .../web_ui/api/tabs/manage/__init__.py | 9 +- zhenxun/builtin_plugins/withdraw.py | 2 +- zhenxun/models/_bot_message_buffer.py | 113 --- zhenxun/models/auth_decision_log.py | 49 -- zhenxun/models/ban_console.py | 8 +- zhenxun/models/bot_console.py | 8 +- zhenxun/models/bot_message_store.py | 55 -- zhenxun/models/fg_request.py | 6 +- zhenxun/models/group_console.py | 170 ++++- zhenxun/models/runtime_backpressure_log.py | 39 - zhenxun/models/user_console.py | 1 + zhenxun/services/auth_observability.py | 547 -------------- zhenxun/services/cache/runtime_cache.py | 450 ++++++++++-- zhenxun/services/data_access.py | 4 + zhenxun/services/db_context/config.py | 3 + zhenxun/services/db_context/schema_guard.py | 16 +- zhenxun/services/runtime_bootstrap.py | 3 - zhenxun/services/scheduler/engine.py | 54 +- zhenxun/services/uninfo_patch.py | 143 +++- zhenxun/utils/common_utils.py | 22 +- zhenxun/utils/enum.py | 5 - zhenxun/utils/platform.py | 87 ++- 56 files changed, 1899 insertions(+), 1654 deletions(-) create mode 100644 zhenxun/builtin_plugins/hooks/auth/data_provider.py delete mode 100644 zhenxun/builtin_plugins/hooks/auth_route.py delete mode 100644 zhenxun/models/_bot_message_buffer.py delete mode 100644 zhenxun/models/auth_decision_log.py delete mode 100644 zhenxun/models/bot_message_store.py delete mode 100644 zhenxun/models/runtime_backpressure_log.py delete mode 100644 zhenxun/services/auth_observability.py diff --git a/.env.example b/.env.example index f8a9c81e..fde84b30 100644 --- a/.env.example +++ b/.env.example @@ -66,6 +66,25 @@ PORT = 8080 # qq adapter load = True QQ_ADAPTER_LOAD=False +# QQ官方适配器配置,启用 QQ_ADAPTER_LOAD 后填写 +# QQ_BOTS=' +# [ +# { +# "id": "", +# "token": "", +# "secret": "", +# "use_websocket": true, +# "intent": { +# "guilds": true, +# "guild_members": true, +# "message_audit": true, +# "at_messages": true, +# "c2c_group_at_messages": false, +# "direct_message": false +# } +# } +# ] +# ' # kook adapter toekn # kaiheila_bots =[{"token": ""}] 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 bb1c12f5..1f9775dc 100644 --- a/zhenxun/builtin_plugins/admin/group_member_update/_data_source.py +++ b/zhenxun/builtin_plugins/admin/group_member_update/_data_source.py @@ -109,7 +109,7 @@ class MemberUpdateManage: members = await interface.get_members(SceneType.GROUP, group_scene.id) try: - group_console, _ = await GroupConsole.get_or_create( + group_console, _ = await GroupConsole.get_or_create_root_group( group_id=group_id, defaults={"platform": platform} ) group_console.member_count = len(members) diff --git a/zhenxun/builtin_plugins/admin/group_update.py b/zhenxun/builtin_plugins/admin/group_update.py index ab2170e0..1a7e3493 100644 --- a/zhenxun/builtin_plugins/admin/group_update.py +++ b/zhenxun/builtin_plugins/admin/group_update.py @@ -39,6 +39,11 @@ _matcher = on_alconna( async def _(bot: Bot, session: EventSession, arparma: Arparma): logger.info("更新群组信息", arparma.header_result, session=session) try: + if PlatformUtils.get_platform_scope(bot) != "qq_client": + await MessageUtils.build_message( + "当前平台不支持旧群组信息同步,仅 OneBot 协议端可用。" + ).send(reply_to=True) + return await PlatformUtils.update_group(bot) await MessageUtils.build_message("已经成功更新了群组信息!").send(reply_to=True) except Exception: diff --git a/zhenxun/builtin_plugins/admin/plugin_switch/data_source.py b/zhenxun/builtin_plugins/admin/plugin_switch/data_source.py index cabea9b2..061f16eb 100644 --- a/zhenxun/builtin_plugins/admin/plugin_switch/data_source.py +++ b/zhenxun/builtin_plugins/admin/plugin_switch/data_source.py @@ -87,7 +87,7 @@ class PluginManager: for gid in groups_to_open | groups_to_close: platform = bot.adapter.get_name() if bot else "qq" - await GroupConsole.get_or_create( + await GroupConsole.get_or_create_root_group( group_id=gid, defaults={"platform": platform} ) diff --git a/zhenxun/builtin_plugins/admin/plugin_switch/strategy.py b/zhenxun/builtin_plugins/admin/plugin_switch/strategy.py index fb1c484e..eac67819 100644 --- a/zhenxun/builtin_plugins/admin/plugin_switch/strategy.py +++ b/zhenxun/builtin_plugins/admin/plugin_switch/strategy.py @@ -173,10 +173,12 @@ class TaskStrategy(SwitchStrategy): async def set_all_default_status(self, status: bool) -> None: await TaskInfo.all().update(default_status=status) + # Bulk updates bypass model save hooks; keep TaskInfoMemoryCache in sync. await self.refresh_cache() async def set_all_global_status(self, status: bool) -> None: await TaskInfo.all().update(status=status) + # Bulk updates bypass model save hooks; keep TaskInfoMemoryCache in sync. await self.refresh_cache() async def refresh_cache(self) -> None: diff --git a/zhenxun/builtin_plugins/catchphrase.py b/zhenxun/builtin_plugins/catchphrase.py index 23736b5b..178e9899 100644 --- a/zhenxun/builtin_plugins/catchphrase.py +++ b/zhenxun/builtin_plugins/catchphrase.py @@ -4,6 +4,7 @@ from nonebot.adapters import Bot from zhenxun.configs.config import Config from zhenxun.services.log import logger +from zhenxun.utils.platform import PlatformUtils Config.add_plugin_config( "catchphrase", @@ -16,6 +17,8 @@ Config.add_plugin_config( @Bot.on_calling_api async def handle_api_call(bot: Bot, api: str, data: dict[str, Any]): + if PlatformUtils.get_platform_scope(bot) != "qq_client": + return if api == "send_msg": catchphrase = Config.get_config("catchphrase", "CATCHPHRASE") if catchphrase and (message := data.get("message")): diff --git a/zhenxun/builtin_plugins/hooks/__init__.py b/zhenxun/builtin_plugins/hooks/__init__.py index 2f8c79de..3ad29d71 100644 --- a/zhenxun/builtin_plugins/hooks/__init__.py +++ b/zhenxun/builtin_plugins/hooks/__init__.py @@ -49,14 +49,4 @@ Config.add_plugin_config( type=bool, ) -Config.add_plugin_config( - "hook", - "RECORD_BOT_SENT_MESSAGES", - True, - help="记录bot消息发送", - default_value=True, - type=bool, -) - - nonebot.load_plugins(str(Path(__file__).parent.resolve())) diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_admin.py b/zhenxun/builtin_plugins/hooks/auth/auth_admin.py index 9c307580..0aac05b1 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_admin.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_admin.py @@ -5,12 +5,12 @@ from nonebot_plugin_uninfo import Uninfo from zhenxun.models.level_user import LevelUser from zhenxun.models.plugin_info import PluginInfo -from zhenxun.services.cache.runtime_cache import LevelUserMemoryCache, LevelUserSnapshot from zhenxun.services.log import logger from zhenxun.utils.utils import EntityIDs, get_entity_ids from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .context import PermissionContext +from .data_provider import DEFAULT_PERMISSION_DATA_PROVIDER, LevelUserSnapshot from .exception import SkipPluginException @@ -50,7 +50,10 @@ async def auth_admin( if cached_levels is not None: global_user, group_users = cached_levels else: - global_user, group_users = await LevelUserMemoryCache.get_levels( + ( + global_user, + group_users, + ) = await DEFAULT_PERMISSION_DATA_PROVIDER.get_admin_levels( entity.user_id, entity.group_id ) diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_ban.py b/zhenxun/builtin_plugins/hooks/auth/auth_ban.py index bc1f660a..7da482e1 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_ban.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_ban.py @@ -7,7 +7,6 @@ from nonebot_plugin_uninfo import Uninfo from zhenxun.configs.config import Config from zhenxun.models.ban_console import BanConsole from zhenxun.models.plugin_info import PluginInfo -from zhenxun.services.cache.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 @@ -15,6 +14,7 @@ from zhenxun.utils.utils import EntityIDs, get_entity_ids from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .context import PermissionContext +from .data_provider import DEFAULT_PERMISSION_DATA_PROVIDER from .exception import SkipPluginException from .utils import freq @@ -60,9 +60,10 @@ async def is_ban(user_id: str | None, group_id: str | None) -> int: """ if not user_id and not group_id: return 0 - if not BanMemoryCache.is_loaded(): + provider = DEFAULT_PERMISSION_DATA_PROVIDER + if not provider.ban_cache_loaded(): return 0 - return BanMemoryCache.remaining_time(user_id, group_id) + return provider.get_ban_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 6a0b30b6..ae76e691 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_bot.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_bot.py @@ -2,12 +2,12 @@ import time from zhenxun.models.bot_console import BotConsole from zhenxun.models.plugin_info import PluginInfo -from zhenxun.services.cache.runtime_cache import BotMemoryCache, BotSnapshot from zhenxun.services.log import logger from zhenxun.utils.common_utils import CommonUtils from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .context import PermissionContext +from .data_provider import DEFAULT_PERMISSION_DATA_PROVIDER, BotSnapshot from .exception import SkipPluginException @@ -33,12 +33,13 @@ async def auth_bot( start_time = time.time() try: + provider = DEFAULT_PERMISSION_DATA_PROVIDER if context is not None: bot_id = context.event.bot_id bot_data = context.bot_data bot: BotConsole | BotSnapshot | None = bot_data if bot is None and not skip_fetch: - bot = await BotMemoryCache.get(bot_id) + bot = await provider.get_bot(bot_id) if bot is None: raise SkipPluginException("Bot不存在,阻断权限检测...") diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_limit.py b/zhenxun/builtin_plugins/hooks/auth/auth_limit.py index b18d15f5..3b570baf 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_limit.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_limit.py @@ -10,10 +10,6 @@ from pydantic import BaseModel 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 @@ -25,6 +21,10 @@ from zhenxun.utils.utils import EntityIDs, get_entity_ids from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .context import PermissionContext +from .data_provider import ( + DEFAULT_PERMISSION_DATA_PROVIDER, + PluginLimitSnapshot, +) from .exception import SkipPluginException driver = nonebot.get_driver() @@ -106,11 +106,10 @@ class LimitManager: block_limit: ClassVar[dict[str, Limit]] = {} count_limit: ClassVar[dict[str, Limit]] = {} - # 模块限制缓存,避免频繁查询数据库 - module_limit_cache: ClassVar[ - dict[str, tuple[float, list[PluginLimitSnapshot], bool]] + # 只缓存异常短路结果;正常 limit 列表统一从 PluginLimitMemoryCache 读取。 + module_limit_error_cache: ClassVar[ + dict[str, tuple[float, list[PluginLimitSnapshot]]] ] = {} - module_cache_ttl: ClassVar[float] = 60 # 模块缓存有效期(秒) module_cache_error_ttl: ClassVar[float] = 5 # 超时缓存有效期(秒) @classmethod @@ -132,15 +131,16 @@ class LimitManager: cls.is_updating = True try: start_time = time.time() - await PluginLimitMemoryCache.ensure_loaded() - limit_list = PluginLimitMemoryCache.get_all_limits() + provider = DEFAULT_PERMISSION_DATA_PROVIDER + await provider.ensure_module_limits_loaded() + limit_list = await provider.get_all_module_limits() # 清空旧数据 cls.add_module = [] cls.cd_limit = {} cls.block_limit = {} cls.count_limit = {} - cls.module_limit_cache.clear() + cls.module_limit_error_cache.clear() # 添加新数据 for limit in limit_list: cls.add_limit(limit) @@ -216,22 +216,21 @@ class LimitManager: """ current_time = time.time() - # 检查缓存 - if module in cls.module_limit_cache: - 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: + # 正常路径不再二次缓存列表,避免与 PluginLimitMemoryCache 形成双真源。 + if module in cls.module_limit_error_cache: + cache_time, limits = cls.module_limit_error_cache[module] + if current_time - cache_time < cls.module_cache_error_ttl: return limits + cls.module_limit_error_cache.pop(module, None) # 缓存不存在或已过期,从内存缓存获取 try: - await PluginLimitMemoryCache.ensure_loaded() - limits = await PluginLimitMemoryCache.get_limits(module) - cls.module_limit_cache[module] = (current_time, limits, False) - return limits + provider = DEFAULT_PERMISSION_DATA_PROVIDER + await provider.ensure_module_limits_loaded() + return await provider.get_module_limits(module) except Exception as exc: logger.error(f"get module limits failed: {module}", LOGGER_COMMAND, e=exc) - cls.module_limit_cache[module] = (current_time, [], True) + cls.module_limit_error_cache[module] = (current_time, []) return [] @classmethod diff --git a/zhenxun/builtin_plugins/hooks/auth/context.py b/zhenxun/builtin_plugins/hooks/auth/context.py index e491398a..34c37bb8 100644 --- a/zhenxun/builtin_plugins/hooks/auth/context.py +++ b/zhenxun/builtin_plugins/hooks/auth/context.py @@ -39,6 +39,7 @@ if TYPE_CHECKING: class EventContext: bot_id: str platform: str + platform_scope: str event_type: str message_id: str | int | None entity: EntityIDs @@ -170,6 +171,7 @@ def event_cache_key( *, bot_id: str, platform: str, + platform_scope: str | None = None, entity: EntityIDs, ) -> str: msg_id = _event_message_id(event) @@ -177,7 +179,11 @@ def event_cache_key( msg_id = id(event) group_id = entity.group_id or "" channel_id = entity.channel_id or "" - return f"{platform}:{bot_id}:{entity.user_id}:{group_id}:{channel_id}:{msg_id}" + scope = platform_scope or platform + return ( + f"{scope}:{platform}:{bot_id}:{entity.user_id}:" + f"{group_id}:{channel_id}:{msg_id}" + ) def get_event_cache( @@ -185,11 +191,18 @@ def get_event_cache( *, bot_id: str, platform: str, + platform_scope: str | None = None, entity: EntityIDs, ) -> dict[str, Any] | None: if not EVENT_CACHE: return None - key = event_cache_key(event, bot_id=bot_id, platform=platform, entity=entity) + key = event_cache_key( + event, + bot_id=bot_id, + platform=platform, + platform_scope=platform_scope, + entity=entity, + ) try: return EVENT_CACHE[key] except KeyError: @@ -255,6 +268,7 @@ def get_or_create_event_context( entity = resolve_entity_ids(event, session) platform = PlatformUtils.get_platform(session) + platform_scope = PlatformUtils.get_platform_scope(session) bot_id = str(bot.self_id) event_cache = state.get(STATE_EVENT_CACHE) if not isinstance(event_cache, dict): @@ -262,6 +276,7 @@ def get_or_create_event_context( event, bot_id=bot_id, platform=platform, + platform_scope=platform_scope, entity=entity, ) @@ -292,6 +307,7 @@ def get_or_create_event_context( context = EventContext( bot_id=bot_id, platform=platform, + platform_scope=platform_scope, event_type=event.get_type(), message_id=_event_message_id(event), entity=entity, diff --git a/zhenxun/builtin_plugins/hooks/auth/data_provider.py b/zhenxun/builtin_plugins/hooks/auth/data_provider.py new file mode 100644 index 00000000..1d3fee89 --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth/data_provider.py @@ -0,0 +1,151 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +from zhenxun.services.cache.runtime_cache import ( + BanMemoryCache, + BotMemoryCache, + BotSnapshot, + GroupMemoryCache, + GroupSnapshot, + LevelUserMemoryCache, + LevelUserSnapshot, + PluginLimitMemoryCache, + PluginLimitSnapshot, +) + +if TYPE_CHECKING: + from zhenxun.models.plugin_info import PluginInfo + + +AdminLevels = tuple[LevelUserSnapshot | None, LevelUserSnapshot | None] + + +class PermissionDataProvider: + """Auth data facade over runtime caches. + + Permission checks should read stable runtime snapshots through this provider + instead of reaching into individual cache classes from multiple auth modules. + The provider does not own policy semantics and does not query the database + directly. + """ + + @staticmethod + def plugin_cache_loaded() -> bool: + from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache + + return PluginInfoMemoryCache.is_loaded() + + @staticmethod + def get_plugin_if_ready(module: str) -> "PluginInfo | None": + from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache + + return PluginInfoMemoryCache.get_by_module_if_ready(module) + + @staticmethod + async def get_plugin(module: str) -> "PluginInfo | None": + from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache + + return await PluginInfoMemoryCache.get_by_module(module) + + @staticmethod + def module_limit_cache_loaded() -> bool: + return PluginLimitMemoryCache.is_loaded() + + @staticmethod + async def ensure_module_limits_loaded() -> None: + await PluginLimitMemoryCache.ensure_loaded() + + @staticmethod + def get_module_limits_if_ready( + module: str, + ) -> list[PluginLimitSnapshot] | None: + return PluginLimitMemoryCache.get_limits_if_ready(module) + + @staticmethod + async def get_module_limits(module: str) -> list[PluginLimitSnapshot]: + return await PluginLimitMemoryCache.get_limits(module) + + @staticmethod + async def get_all_module_limits() -> list[PluginLimitSnapshot]: + if not PluginLimitMemoryCache.is_loaded(): + await PluginLimitMemoryCache.ensure_loaded() + return PluginLimitMemoryCache.get_all_limits() + + @staticmethod + def bot_cache_loaded() -> bool: + return BotMemoryCache.is_loaded() + + @staticmethod + def get_bot_if_ready(bot_id: str | None) -> BotSnapshot | None: + return BotMemoryCache.get_if_ready(bot_id) + + @staticmethod + async def get_bot(bot_id: str | None) -> BotSnapshot | None: + return await BotMemoryCache.get(bot_id) + + @staticmethod + def group_cache_loaded() -> bool: + return GroupMemoryCache.is_loaded() + + @staticmethod + def get_group_if_ready( + group_id: str | None, + channel_id: str | None = None, + ) -> GroupSnapshot | None: + return GroupMemoryCache.get_if_ready(group_id, channel_id) + + @staticmethod + async def get_group( + group_id: str | None, + channel_id: str | None = None, + ) -> GroupSnapshot | None: + return await GroupMemoryCache.get(group_id, channel_id) + + @staticmethod + def admin_cache_loaded() -> bool: + return LevelUserMemoryCache.is_loaded() + + @staticmethod + def get_admin_levels_if_ready( + user_id: str | None, + group_id: str | None, + ) -> AdminLevels | None: + return LevelUserMemoryCache.get_levels_if_ready(user_id, group_id) + + @staticmethod + async def get_admin_levels( + user_id: str | None, + group_id: str | None, + ) -> AdminLevels: + return await LevelUserMemoryCache.get_levels(user_id, group_id) + + @staticmethod + def ban_cache_loaded() -> bool: + return BanMemoryCache.is_loaded() + + @staticmethod + async def ensure_ban_loaded() -> None: + await BanMemoryCache.ensure_loaded() + + @staticmethod + def is_banned(user_id: str | None, group_id: str | None) -> bool: + return BanMemoryCache.is_banned(user_id, group_id) + + @staticmethod + def get_ban_remaining_time(user_id: str | None, group_id: str | None) -> int: + return BanMemoryCache.remaining_time(user_id, group_id) + + +DEFAULT_PERMISSION_DATA_PROVIDER = PermissionDataProvider() + + +__all__ = [ + "DEFAULT_PERMISSION_DATA_PROVIDER", + "AdminLevels", + "BotSnapshot", + "GroupSnapshot", + "LevelUserSnapshot", + "PermissionDataProvider", + "PluginLimitSnapshot", +] diff --git a/zhenxun/builtin_plugins/hooks/auth_activation.py b/zhenxun/builtin_plugins/hooks/auth_activation.py index 0255ccb9..8b2ce9fc 100644 --- a/zhenxun/builtin_plugins/hooks/auth_activation.py +++ b/zhenxun/builtin_plugins/hooks/auth_activation.py @@ -8,6 +8,7 @@ import re from typing import Any, Literal import weakref +from loguru import logger from nonebot.matcher import Matcher ActivationDecision = Literal["match", "miss", "unknown"] @@ -23,6 +24,20 @@ ActivationLane = Literal[ "passive_render", ] +KNOWN_SAFE_RULE_NAMES = frozenset( + { + "CommandRule", + "ShellCommandRule", + "RegexRule", + "StartswithRule", + "EndswithRule", + "FullmatchRule", + "KeywordsRule", + "IsTypeRule", + "ToMeRule", + } +) + @dataclass(frozen=True, slots=True) class ActivationRuleDescriptor: @@ -48,9 +63,96 @@ class HandlerDescriptor: has_custom_rule: bool = False commands: tuple[str, ...] = () shortcuts: tuple[str, ...] | None = None + alconna: tuple[AlconnaDescriptor, ...] = () rules: tuple[ActivationRuleDescriptor, ...] = () +@dataclass(frozen=True, slots=True) +class AlconnaShortcutDescriptor: + pattern: str + fuzzy: bool = False + flags: int = 0 + + +@dataclass(frozen=True, slots=True) +class AlconnaDescriptor: + command: str = "" + aliases: tuple[str, ...] = () + prefixes: tuple[str, ...] = () + shortcuts: tuple[AlconnaShortcutDescriptor, ...] = () + compact: bool = False + skip_for_unmatch: bool = True + before_rule_count: int = 0 + after_rule_count: int = 0 + before_rule_known_safe: bool = True + after_rule_known_safe: bool = True + input_rewrite_extensions: tuple[str, ...] = () + + @property + def has_reply_merge_extension(self) -> bool: + return "ReplyMergeExtension" in self.input_rewrite_extensions + + @property + def regex_command(self) -> bool: + return self.command.startswith("re:") + + +class AlconnaActivationIndex: + """Safe, metadata-only Alconna prefilter. + + This index never executes Alconna.parse() or Matcher.check_rule(). It only + skips matchers when the command head/shortcut is known to miss. Anything + custom or ambiguous stays fail-open to preserve NoneBot compatibility. + """ + + def __init__(self) -> None: + self._safe = 0 + self._unknown = 0 + + @property + def safe_count(self) -> int: + return self._safe + + @property + def unknown_count(self) -> int: + return self._unknown + + def rebuild(self, descriptors: Iterable[HandlerDescriptor]) -> None: + safe = 0 + unknown = 0 + for descriptor in descriptors: + if descriptor.alconna: + if any(_alconna_can_prefilter(item) for item in descriptor.alconna): + safe += 1 + else: + unknown += 1 + self._safe = safe + self._unknown = unknown + if logger.level("DEBUG"): + logger.debug( + "alconna activation index rebuilt: safe={}, unknown={}", + safe, + unknown, + ) + + def select( + self, + descriptor: HandlerDescriptor, + context: ActivationContext, + texts: tuple[str, ...], + ) -> ActivationDecision: + if not descriptor.alconna: + return "unknown" + saw_unknown = False + for alconna in descriptor.alconna: + decision = matcher_alconna_head_matches(alconna, texts, context) + if decision == "match": + return "match" + if decision == "unknown": + saw_unknown = True + return "unknown" if saw_unknown else "miss" + + @dataclass(slots=True) class ActivationContext: event_type: str @@ -69,37 +171,25 @@ class ActivationContext: @dataclass(slots=True) class ActivationResult: selected: list[type[Matcher]] - fallback_required: bool = False - selected_by_lane: dict[str, int] = field(default_factory=dict) - skipped_by_lane: dict[str, int] = field(default_factory=dict) - selected_by_reason: dict[str, int] = field(default_factory=dict) - skipped_by_reason: dict[str, int] = field(default_factory=dict) deterministic_selected: set[type[Matcher]] = field(default_factory=set) total_descriptors: int = 0 candidate_count: int = 0 - def mark_selected(self, lane: str, reason: str = "selected") -> None: - self.selected_by_lane[lane] = self.selected_by_lane.get(lane, 0) + 1 - self.selected_by_reason[reason] = self.selected_by_reason.get(reason, 0) + 1 - - def mark_skipped(self, lane: str, reason: str = "skipped") -> None: - self.skipped_by_lane[lane] = self.skipped_by_lane.get(lane, 0) + 1 - self.skipped_by_reason[reason] = self.skipped_by_reason.get(reason, 0) + 1 - class HandlerActivationIndex: """In-memory matcher activation index. - The index is intentionally fail-open: only proven misses are rejected before - matcher task creation. Unknown custom rules, incomplete command metadata, - and Alconna shortcut misses stay selected so plugin compatibility wins over - dispatch aggressiveness. + The index is intentionally fail-open: only known NoneBot rule misses are + rejected before matcher task creation. Custom rules, incomplete command + metadata, and Alconna shortcut misses stay selected so plugin compatibility + wins over dispatch aggressiveness. """ def __init__(self) -> None: self._by_priority: dict[int, list[HandlerDescriptor]] = {} self._matcher_map: dict[type[Matcher], HandlerDescriptor] = {} self._source_keys: set[tuple[int, tuple[int, ...]]] = set() + self._alconna_index = AlconnaActivationIndex() self._compiled = False @property @@ -122,6 +212,7 @@ class HandlerActivationIndex: (int(priority), tuple(id(matcher) for matcher in priority_matchers)) ) self._source_keys = source_keys + self._alconna_index.rebuild(self._matcher_map.values()) self._compiled = True def ensure_fresh(self, matchers: dict[int, list[type[Matcher]]]) -> None: @@ -164,32 +255,16 @@ class HandlerActivationIndex: ] for descriptor in descriptors: decision = self._select_descriptor(descriptor, context) - lane = descriptor.lane - if decision == "fallback": - result.fallback_required = True - if not _consume_uncertain_budget(descriptor, budget): - result.mark_skipped(lane, "fallback_budget_exhausted") - continue - result.selected.append(descriptor.matcher) - result.mark_selected(lane, "fallback_budgeted") - continue if decision == "miss": - result.mark_skipped(lane, _miss_reason(descriptor, context)) continue if decision == "deterministic": result.selected.append(descriptor.matcher) result.deterministic_selected.add(descriptor.matcher) - result.mark_selected(lane, "deterministic") continue - if not _selection_is_guaranteed(descriptor, context): - if not _consume_uncertain_budget(descriptor, budget): - result.mark_skipped(lane, "unknown_budget_exhausted") + if _is_throttleable_broad_passive(descriptor, context): + if not _consume_broad_passive_budget(descriptor, budget): continue - selected_reason = "unknown_budgeted" - else: - selected_reason = "guaranteed" result.selected.append(descriptor.matcher) - result.mark_selected(lane, selected_reason) result.candidate_count = len(result.selected) return result @@ -197,7 +272,7 @@ class HandlerActivationIndex: self, descriptor: HandlerDescriptor, context: ActivationContext, - ) -> Literal["select", "miss", "fallback", "deterministic"]: + ) -> Literal["select", "miss", "deterministic"]: if descriptor.temp: return "select" matcher_type = descriptor.matcher_type @@ -224,7 +299,7 @@ class HandlerActivationIndex: self, descriptor: HandlerDescriptor, context: ActivationContext, - ) -> Literal["select", "miss", "fallback", "deterministic"]: + ) -> Literal["select", "miss", "deterministic"]: texts = text_match_candidates( context.plain_text, context.raw_text, @@ -232,8 +307,25 @@ class HandlerActivationIndex: ) if not texts: return "select" + rule_match = matcher_rule_matches_text( + descriptor.rules, + context.raw_text, + context.plain_text, + event=context.event, + to_me=context.to_me, + ) + if rule_match == "miss": + return "miss" command_matched = False - if descriptor.commands: + if descriptor.alconna: + alconna_match = self._alconna_index.select(descriptor, context, texts) + if alconna_match == "match": + command_matched = True + elif alconna_match == "miss": + return "miss" + else: + return "select" + elif descriptor.commands: if any( matcher_command_matches(text, command) for text in texts @@ -248,9 +340,7 @@ class HandlerActivationIndex: if shortcut_match == "match": command_matched = True else: - # Command extraction for Alconna/custom matchers is incomplete by - # design; a miss here is not proof that NoneBot will miss. - return "select" + return "miss" if not descriptor.has_custom_rule else "select" else: shortcut_match = matcher_alconna_shortcut_matches_any( descriptor.shortcuts, @@ -259,22 +349,19 @@ class HandlerActivationIndex: if shortcut_match == "match": command_matched = True - rule_match = matcher_rule_matches_text( - descriptor.rules, - context.raw_text, - context.plain_text, - event=context.event, - to_me=context.to_me, - ) - if rule_match == "match": + if ( + rule_match == "match" + and not descriptor.has_custom_rule + and descriptor.shortcuts is None + and not descriptor.alconna + ): command_matched = True - elif rule_match == "miss": - return "miss" - elif ( + if ( not (descriptor.commands or descriptor.shortcuts is not None) and not command_matched + and not descriptor.alconna ): - return "fallback" + return "select" if ( context.ai_route_modules @@ -284,7 +371,7 @@ class HandlerActivationIndex: return "select" if command_matched: - return "deterministic" + return "select" if descriptor.has_custom_rule else "deterministic" return "select" def _build_descriptor( @@ -298,7 +385,10 @@ class HandlerActivationIndex: if hasattr(matcher, "command"): command_like = True commands = extract_matcher_command_literals(matcher) or () + alconna_descriptors = extract_matcher_alconna_descriptors(matcher) shortcuts = extract_matcher_alconna_shortcuts(matcher) + if alconna_descriptors: + command_like = True if shortcuts is not None: command_like = True module = matcher_module_name(matcher) @@ -323,6 +413,7 @@ class HandlerActivationIndex: has_custom_rule=matcher_has_custom_rule(matcher), commands=commands, shortcuts=shortcuts, + alconna=alconna_descriptors, rules=rules, ) @@ -393,6 +484,8 @@ def matcher_is_command_like(matcher_cls: type[Matcher]) -> bool: return True if hasattr(matcher_cls, "command"): return True + if extract_matcher_alconna_descriptors(matcher_cls): + return True return extract_matcher_alconna_shortcuts(matcher_cls) is not None @@ -420,7 +513,10 @@ def classify_matcher_lane( if hasattr(matcher_cls, "command"): command_like = True commands = extract_matcher_command_literals(matcher_cls) or () + alconna_descriptors = extract_matcher_alconna_descriptors(matcher_cls) shortcuts = extract_matcher_alconna_shortcuts(matcher_cls) + if alconna_descriptors: + command_like = True if shortcuts is not None: command_like = True return classify_lane( @@ -437,11 +533,6 @@ def extract_matcher_rule_descriptors( matcher_cls: type[Matcher], ) -> tuple[ActivationRuleDescriptor, ...]: descriptors: list[ActivationRuleDescriptor] = [] - if hasattr(matcher_cls, "command"): - descriptors.append( - ActivationRuleDescriptor("matcher_command", command_like=True) - ) - rule = getattr(matcher_cls, "rule", None) checkers = getattr(rule, "checkers", ()) or () for checker in checkers: @@ -530,48 +621,10 @@ def extract_matcher_rule_descriptors( ): descriptors.append(ActivationRuleDescriptor("alconna", command_like=True)) else: - descriptors.append(_custom_rule_descriptor(call)) + descriptors.append(ActivationRuleDescriptor("custom")) return tuple(descriptors) -def _custom_rule_descriptor(call: object) -> ActivationRuleDescriptor: - keyword_regex = _extract_keyword_regex_pairs(call) - if keyword_regex: - return ActivationRuleDescriptor( - "keyword_regex", - keyword_regex, - deterministic_text=True, - ) - return ActivationRuleDescriptor("custom") - - -def _extract_keyword_regex_pairs( - call: object, -) -> tuple[tuple[str, str, int], ...]: - """Recognize generic keyword + regex custom rules without plugin coupling.""" - - source = getattr(call, "key_pattern_list", None) - if source is None: - source = getattr(call, "keyword_patterns", None) - if source is None: - source = getattr(call, "patterns", None) - if not isinstance(source, Iterable) or isinstance(source, str): - return () - - pairs: list[tuple[str, str, int]] = [] - for item in source: - if not isinstance(item, tuple | list) or len(item) < 2: - continue - keyword = str(item[0] or "").strip() - pattern_obj = item[1] - pattern = getattr(pattern_obj, "pattern", pattern_obj) - if not keyword or not isinstance(pattern, str) or not pattern: - continue - flags = int(getattr(pattern_obj, "flags", 0) or 0) - pairs.append((keyword, pattern, flags)) - return tuple(pairs) - - def normalize_rule_string_tuple(value: object) -> tuple[str, ...]: if isinstance(value, str): return (value,) @@ -633,13 +686,38 @@ def matcher_rule_matches_text( for descriptor in descriptors: kind = descriptor.kind - if kind == "regex": + if kind in {"custom", "alconna"}: + saw_unknown = True + continue + if kind in {"command", "shell_command"}: + saw_deterministic = True + commands: set[str] = set() + collect_command_literals(descriptor.value, commands) + normalized_commands = { + normalized + for item in commands + if (normalized := normalize_command(item)) + } + if any( + matcher_command_matches(text, command) + for text in plain_candidates + for command in normalized_commands + ): + matched_any = True + else: + return "miss" + elif kind in {"regex", "regex_fullmatch"}: saw_deterministic = True pattern = str(descriptor.value or "") if not pattern: continue try: - if re.search(pattern, message_text, descriptor.flags): + matched = ( + re.fullmatch(pattern, message_text, descriptor.flags) + if kind == "regex_fullmatch" + else re.search(pattern, message_text, descriptor.flags) + ) + if matched: matched_any = True else: return "miss" @@ -709,13 +787,6 @@ def matcher_rule_matches_text( matched_any = True else: return "miss" - elif kind == "keyword_regex": - saw_deterministic = True - values = descriptor.value if isinstance(descriptor.value, tuple) else () - if _keyword_regex_matches(values, plain_candidates): - matched_any = True - else: - return "miss" elif kind == "to_me": if not to_me: return "miss" @@ -729,43 +800,15 @@ def matcher_rule_matches_text( elif isinstance(types, tuple) and types: if not isinstance(event, types): return "miss" - elif kind in {"custom", "alconna", "matcher_command"}: - saw_unknown = True - if matched_any: - return "match" - if saw_deterministic: - return "miss" + return "unknown" if saw_unknown else "match" if saw_unknown: return "unknown" + if saw_deterministic: + return "miss" return "unknown" -def _keyword_regex_matches( - values: object, - candidates: tuple[str, ...], -) -> bool: - if not isinstance(values, tuple): - return False - for item in values: - if not isinstance(item, tuple | list) or len(item) < 3: - continue - keyword, pattern, flags = item[:3] - keyword_text = str(keyword or "") - pattern_text = str(pattern or "") - if not keyword_text or not pattern_text: - continue - for text in candidates: - if keyword_text not in text: - continue - try: - if re.search(pattern_text, text, int(flags or 0)): - return True - except re.error: - return False - return False - - def extract_matcher_command_literals( matcher_cls: type[Matcher], ) -> tuple[str, ...] | None: @@ -912,6 +955,293 @@ def collect_alconna_shortcuts(value: Any, target: set[str], depth: int = 0) -> N collect_alconna_shortcuts(nested, target, depth + 1) +def extract_matcher_alconna_descriptors( + matcher_cls: type[Matcher], +) -> tuple[AlconnaDescriptor, ...]: + descriptors: list[AlconnaDescriptor] = [] + rule = getattr(matcher_cls, "rule", None) + checkers = getattr(rule, "checkers", ()) or () + for checker in checkers: + call = getattr(checker, "call", None) + if call is None: + continue + if call.__class__.__name__ != "AlconnaRule": + continue + if not call.__class__.__module__.startswith("nonebot_plugin_alconna.rule"): + continue + command = resolve_maybe_weakref( + getattr(call, "command", None) or getattr(call, "alconna", None) + ) + descriptor = _build_alconna_descriptor(call, command) + if descriptor is not None: + descriptors.append(descriptor) + return tuple(descriptors) + + +def _build_alconna_descriptor(call: Any, command: Any) -> AlconnaDescriptor | None: + if command is None: + return None + command_text = str(getattr(command, "command", "") or "").strip() + aliases = tuple( + str(item).strip() + for item in getattr(command, "aliases", ()) or () + if str(item).strip() and str(item).strip() != command_text + ) + prefixes = tuple( + str(item) + for item in getattr(command, "prefixes", ()) or () + if isinstance(item, str) + ) + meta = getattr(command, "meta", None) + shortcuts = _extract_alconna_shortcut_descriptors(command) + return AlconnaDescriptor( + command=command_text, + aliases=aliases, + prefixes=prefixes, + shortcuts=shortcuts, + compact=bool(getattr(meta, "compact", False)), + skip_for_unmatch=bool(getattr(call, "skip", True)), + before_rule_count=_rule_checker_count(getattr(call, "before_rules", None)), + after_rule_count=_rule_checker_count(getattr(call, "after_rules", None)), + before_rule_known_safe=_alconna_rule_is_known_safe( + getattr(call, "before_rules", None) + ), + after_rule_known_safe=_alconna_rule_is_known_safe( + getattr(call, "after_rules", None) + ), + input_rewrite_extensions=_extract_alconna_input_rewrite_extensions(call), + ) + + +def _rule_checker_count(rule: Any) -> int: + checkers = getattr(rule, "checkers", ()) or () + with contextlib.suppress(TypeError): + return len(checkers) + return 1 + + +def _alconna_rule_is_known_safe(rule: Any) -> bool: + """Whether Alconna before/after Rule can be reasoned about statically. + + We still do not execute these rules here. A rule is considered safe only + when every checker is an official NoneBot rule whose negative result can be + reproduced by the selector. Custom rules stay fail-open. + """ + + checkers = getattr(rule, "checkers", ()) or () + for checker in checkers: + call = getattr(checker, "call", None) + if call is None: + return False + call_module = call.__class__.__module__ + call_name = call.__class__.__name__ + if not call_module.startswith("nonebot.rule"): + return False + if call_name not in KNOWN_SAFE_RULE_NAMES: + return False + return True + + +def _extract_alconna_shortcut_descriptors( + command: Any, +) -> tuple[AlconnaShortcutDescriptor, ...]: + shortcuts: list[AlconnaShortcutDescriptor] = [] + with contextlib.suppress(Exception): + from arclet.alconna import command_manager + + raw_shortcuts = command_manager.get_shortcut(command) # type: ignore[arg-type] + if isinstance(raw_shortcuts, dict): + for key, args in raw_shortcuts.items(): + pattern = str(key or "").strip() + if not pattern: + continue + shortcuts.append( + AlconnaShortcutDescriptor( + pattern=pattern, + fuzzy=bool(getattr(args, "fuzzy", False)), + flags=int(getattr(args, "flags", 0) or 0), + ) + ) + if shortcuts: + return tuple(shortcuts) + fallback: set[str] = set() + collect_alconna_shortcuts(command, fallback) + return tuple( + AlconnaShortcutDescriptor(pattern=item) + for item in sorted(fallback) + if item.strip() + ) + + +def _extract_alconna_input_rewrite_extensions(call: Any) -> tuple[str, ...]: + executor = getattr(call, "executor", None) + if executor is None: + return () + result: list[str] = [] + for attr in ("extensions", "_extensions", "exts", "context"): + extensions = getattr(executor, attr, None) + if not isinstance(extensions, list | tuple | set | frozenset): + continue + for extension in extensions: + overrides = getattr(extension.__class__, "_overrides", None) + if not isinstance(overrides, dict): + overrides = getattr(extension, "_overrides", None) + if not isinstance(overrides, dict): + continue + if not ( + bool(overrides.get("message_provider")) + or bool(overrides.get("receive_wrapper")) + ): + continue + name = extension.__class__.__name__ + if name not in result: + result.append(name) + return tuple(result) + + +def _alconna_can_prefilter(alconna: AlconnaDescriptor) -> bool: + if not (alconna.command or alconna.aliases or alconna.shortcuts): + return False + if not (alconna.before_rule_known_safe and alconna.after_rule_known_safe): + return False + if any(name != "ReplyMergeExtension" for name in alconna.input_rewrite_extensions): + return False + return True + + +def matcher_alconna_head_matches( + alconna: AlconnaDescriptor, + texts: Iterable[str], + context: ActivationContext, +) -> ActivationDecision: + if not _alconna_can_prefilter(alconna): + return "unknown" + if alconna.has_reply_merge_extension and _event_has_reply(context.event): + return "unknown" + + candidates = tuple(text.strip() for text in texts if text and text.strip()) + if not candidates: + return "unknown" + + saw_unknown = False + command_heads = (alconna.command, *alconna.aliases) + for text in candidates: + for command in command_heads: + if not command: + continue + decision = _alconna_command_head_matches( + text, + command, + alconna.prefixes, + compact=alconna.compact, + ) + if decision == "match": + return "match" + if decision == "unknown": + saw_unknown = True + for shortcut in alconna.shortcuts: + decision = _alconna_shortcut_matches(text, shortcut) + if decision == "match": + return "match" + if decision == "unknown": + saw_unknown = True + return "unknown" if saw_unknown else "miss" + + +def _alconna_command_head_matches( + text: str, + command: str, + prefixes: tuple[str, ...], + *, + compact: bool, +) -> ActivationDecision: + normalized = command.strip() + if not normalized: + return "unknown" + if normalized.startswith("re:"): + pattern = normalized.removeprefix("re:").strip() + if not pattern: + return "unknown" + for prefix in prefixes or ("",): + try: + if re.match(rf"^{re.escape(prefix)}(?:{pattern})", text): + return "match" + except re.error: + return "unknown" + return "miss" + for prefix in prefixes or ("",): + head = f"{prefix}{normalized}" + if _alconna_literal_head_matches(text, head, compact=compact): + return "match" + return "miss" + + +def _alconna_literal_head_matches(text: str, head: str, *, compact: bool) -> bool: + if not text or not head: + return False + if text == head: + return True + if not text.startswith(head): + return False + if len(text) == len(head): + return True + rest = text[len(head) :] + if rest and rest[0].isspace(): + return True + return bool(compact or not head[-1].isascii()) + + +def _alconna_shortcut_matches( + text: str, + shortcut: AlconnaShortcutDescriptor, +) -> ActivationDecision: + pattern = shortcut.pattern.strip() + if not pattern: + return "unknown" + normalized = normalize_shortcut_pattern(pattern) + placeholder_match = _placeholder_shortcut_decision(text, normalized) + if placeholder_match == "match": + return "match" + if placeholder_match == "unknown": + return "unknown" + if not is_regex_like_shortcut(normalized) and matcher_command_matches( + text, + normalized, + ): + return "match" + try: + if shortcut.fuzzy: + return "match" if re.match(f"^{pattern}", text, shortcut.flags) else "miss" + return "match" if re.fullmatch(pattern, text, shortcut.flags) else "miss" + except re.error: + return "unknown" + + +def _event_has_reply(event: object | None) -> bool: + if event is None: + return False + with contextlib.suppress(Exception): + message = event.get_message() # type: ignore[attr-defined] + for segment in message: + segment_type = getattr(segment, "type", None) + if segment_type == "reply": + return True + if isinstance(segment, dict) and segment.get("type") == "reply": + return True + for attr in ("reply", "reply_message", "source"): + with contextlib.suppress(Exception): + if getattr(event, attr, None) is not None: + return True + raw_text = "" + with contextlib.suppress(Exception): + raw_text = str(event.get_message()) # type: ignore[attr-defined] + lowered = raw_text.casefold() + return any( + marker in lowered + for marker in ("[cq:reply", "type=reply", '"reply"', "'reply'") + ) + + def resolve_maybe_weakref(value: Any) -> Any: if isinstance(value, weakref.ReferenceType): resolved = value() @@ -1026,8 +1356,12 @@ def shortcut_matches_text(text: str, shortcut: str) -> bool: def placeholder_shortcut_matches(text: str, pattern: str) -> bool: + return _placeholder_shortcut_decision(text, pattern) == "match" + + +def _placeholder_shortcut_decision(text: str, pattern: str) -> ActivationDecision: if "{" not in pattern or "}" not in pattern: - return False + return "miss" pieces: list[str] = [] last = 0 for match in re.finditer(r"\{[^{}]+\}", pattern): @@ -1035,12 +1369,16 @@ def placeholder_shortcut_matches(text: str, pattern: str) -> bool: pieces.append(r"\S+") last = match.end() if not pieces: - return False + return "miss" pieces.append(re.escape(pattern[last:])) try: - return re.match(rf"^{''.join(pieces)}(?:\s|$)", text) is not None + return ( + "match" + if re.match(rf"^{''.join(pieces)}(?:\s|$)", text) is not None + else "miss" + ) except re.error: - return False + return "unknown" def is_regex_like_shortcut(pattern: str) -> bool: @@ -1072,21 +1410,27 @@ def matcher_matches_ai_route_heads( return False -def _selection_is_guaranteed( +def _is_throttleable_broad_passive( descriptor: HandlerDescriptor, context: ActivationContext, ) -> bool: - """Return True for candidates that must not be budget-throttled.""" + """Only broad, no-rule passive message matchers may be budget-throttled.""" - if descriptor.temp or descriptor.lane == "system": - return True if context.event_type != "message": - return not descriptor.command_like + return False + if descriptor.temp or descriptor.lane == "system": + return False + if not descriptor.lane.startswith("passive_"): + return False + if descriptor.command_like or descriptor.deterministic_text: + return False + if descriptor.has_custom_rule or descriptor.rules: + return False if descriptor.lane == "passive_http" and ( context.has_url or _looks_like_rich_message(context.raw_text) ): - return True - return False + return False + return True def _looks_like_rich_message(text: str) -> bool: @@ -1106,34 +1450,11 @@ def _looks_like_rich_message(text: str) -> bool: ) -def _miss_reason(descriptor: HandlerDescriptor, context: ActivationContext) -> str: - matcher_type = descriptor.matcher_type - if matcher_type and matcher_type != context.event_type: - return "type_miss" - if context.event_type != "message" and descriptor.command_like: - return "non_message_command" - if descriptor.command_like: - return "command_or_rule_miss" - return "rule_miss" - - -def _uncertain_budget_lane(descriptor: HandlerDescriptor) -> str: - lane = descriptor.lane - if lane.startswith("passive_"): - return lane - # Unknown command-like matchers are fail-open for compatibility, but they - # should not fan out as unbounded command tasks when no deterministic signal - # matched. Put them into the cheapest passive bucket. - if lane.startswith("command_"): - return "passive_light" - return lane - - -def _consume_uncertain_budget( +def _consume_broad_passive_budget( descriptor: HandlerDescriptor, budget: dict[str, int], ) -> bool: - lane = _uncertain_budget_lane(descriptor) + lane = descriptor.lane if lane not in budget: return True if budget[lane] <= 0: diff --git a/zhenxun/builtin_plugins/hooks/auth_checker.py b/zhenxun/builtin_plugins/hooks/auth_checker.py index df6f1df1..60edee43 100644 --- a/zhenxun/builtin_plugins/hooks/auth_checker.py +++ b/zhenxun/builtin_plugins/hooks/auth_checker.py @@ -15,19 +15,9 @@ from nonebot_plugin_uninfo import Uninfo from zhenxun.configs.utils import PluginExtraData from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.user_console import UserConsole -from zhenxun.services.auth_observability import ( - append_auth_decision_log, - append_runtime_backpressure_log, - build_auth_observability_report, -) from zhenxun.services.cache.cache_containers import CacheDict -from zhenxun.services.cache.runtime_cache import ( - PluginInfoMemoryCache, - PluginLimitMemoryCache, -) -from zhenxun.services.data_access import DataAccess from zhenxun.services.log import logger -from zhenxun.services.message_load import is_overloaded, signal_overload +from zhenxun.services.message_load import signal_overload from zhenxun.utils.enum import GoldHandle, PluginType from zhenxun.utils.exception import InsufficientGold from zhenxun.utils.platform import PlatformUtils @@ -48,6 +38,7 @@ from .auth.context import ( set_route_modules, store_permission_context, ) +from .auth.data_provider import DEFAULT_PERMISSION_DATA_PROVIDER from .auth.exception import ( IsSuperuserException, PermissionExemption, @@ -55,7 +46,6 @@ from .auth.exception import ( ) from .auth_activation import ( ActivationContext, - ActivationResult, HandlerActivationIndex, classify_matcher_lane, extract_matcher_alconna_shortcuts, @@ -147,9 +137,7 @@ _ROUTE_COMMAND_MAP: dict[str, set[str]] = {} _ROUTE_PREFIX_MAP: dict[str, set[str]] = {} _ROUTE_MODULES_WITH_COMMANDS: set[str] = set() MATCHER_ROUTE_PREFILTER_TTL = AUTH_DISPATCH_RUNTIME_CONFIG.matcher_route_prefilter_ttl -PREFILTER_STATS_LOG_INTERVAL = AUTH_DISPATCH_RUNTIME_CONFIG.prefilter_stats_log_interval CACHE_SWEEP_INTERVAL = AUTH_DISPATCH_RUNTIME_CONFIG.cache_sweep_interval -DISPATCH_STATS_LOG_INTERVAL = AUTH_DISPATCH_RUNTIME_CONFIG.dispatch_stats_log_interval # 全局信号量与计数器 HOOKS_ACTIVE_COUNT = 0 @@ -182,56 +170,6 @@ _CHECK_MATCHER_ROUTE_CACHE = CacheDict( ) -_PREFILTER_STATS = { - "checked": 0, - "skipped": 0, - "before_task_checked": 0, - "before_task_skipped": 0, - "inside_task_checked": 0, - "inside_task_skipped": 0, - "type_miss": 0, - "route_miss": 0, - "command_miss": 0, - "empty_text": 0, -} -_PREFILTER_LAST_LOG = 0.0 -_DISPATCH_SELECTED = 0 -_DISPATCH_SKIPPED = 0 -_DISPATCH_SELECTED_BY_LANE: dict[str, int] = { - "command_exact": 0, - "command_shortcut": 0, - "command_regex": 0, - "system": 0, - "passive_light": 0, - "passive_db": 0, - "passive_http": 0, - "passive_ai": 0, - "passive_render": 0, -} -_DISPATCH_SKIPPED_BY_LANE: dict[str, int] = { - "command_exact": 0, - "command_shortcut": 0, - "command_regex": 0, - "system": 0, - "passive_light": 0, - "passive_db": 0, - "passive_http": 0, - "passive_ai": 0, - "passive_render": 0, -} -_DISPATCH_LANE_WAIT_MS: dict[str, float] = { - "command_exact": 0.0, - "command_shortcut": 0.0, - "command_regex": 0.0, - "system": 0.0, - "passive_light": 0.0, - "passive_db": 0.0, - "passive_http": 0.0, - "passive_ai": 0.0, - "passive_render": 0.0, -} -_DISPATCH_LAST_LOG = 0.0 -_DISPATCH_SHADOW_LAST_LOG = 0.0 _CACHE_SWEEP_TASK: asyncio.Task | None = None _BOT_WAKE_COMMAND_PATTERN = re.compile(r"^bot醒来(?:\s+\S+)?$", re.IGNORECASE) _BOT_WAKE_CANONICAL_PATTERN = re.compile( @@ -367,12 +305,6 @@ def _is_bot_wake_command(module: str, text: str | None) -> bool: ) -def _debug_log(message: str, *args, **kwargs) -> None: - if is_overloaded(): - return - logger.debug(message, *args, **kwargs) - - def _is_command_matcher_class(matcher_cls: type[Matcher]) -> bool: if matcher_cls in _MATCHER_COMMAND_TYPE_CACHE: return _MATCHER_COMMAND_TYPE_CACHE[matcher_cls] @@ -728,41 +660,6 @@ def _activation_context_from_dispatch( ) -def _record_dispatch_selection(lane: str, selected: bool, wait_ms: float = 0.0) -> None: - global _DISPATCH_LAST_LOG, _DISPATCH_SELECTED, _DISPATCH_SKIPPED - if lane == "command": - lane = "command_exact" - lane = lane if lane in _DISPATCH_SELECTED_BY_LANE else "passive_light" - if selected: - _DISPATCH_SELECTED += 1 - _DISPATCH_SELECTED_BY_LANE[lane] += 1 - _DISPATCH_LANE_WAIT_MS[lane] += wait_ms - else: - _DISPATCH_SKIPPED += 1 - _DISPATCH_SKIPPED_BY_LANE[lane] += 1 - - now = time.monotonic() - if now - _DISPATCH_LAST_LOG < DISPATCH_STATS_LOG_INTERVAL or is_overloaded(): - return - _DISPATCH_LAST_LOG = now - wait_snapshot = { - lane: round(wait, 2) for lane, wait in _DISPATCH_LANE_WAIT_MS.items() - } - lane_snapshot = " ".join( - f"{lane}={count}" for lane, count in _DISPATCH_SELECTED_BY_LANE.items() - ) - _debug_log( - ( - "dispatch stats: " - f"selected={_DISPATCH_SELECTED} " - f"skipped={_DISPATCH_SKIPPED} " - f"{lane_snapshot} " - f"wait_ms={wait_snapshot}" - ), - LOGGER_COMMAND, - ) - - def _new_dispatch_budget() -> dict[str, int]: return dict(_DISPATCH_LANE_LIMITS) @@ -775,52 +672,6 @@ def _merge_dispatch_budget( target[lane] = source.get(lane, target.get(lane, 0)) -def _record_activation_result(activation_result: ActivationResult) -> None: - for lane, count in activation_result.skipped_by_lane.items(): - for _ in range(count): - _record_dispatch_selection(lane, False) - - -def _compact_counter(counter: dict[str, int], limit: int = 6) -> str: - if not counter: - return "-" - items = sorted(counter.items(), key=lambda item: item[1], reverse=True)[:limit] - return ",".join(f"{key}={value}" for key, value in items) - - -def _debug_activation_shadow( - *, - priority: int, - activation_result: ActivationResult, - context: EventDispatchContext, -) -> None: - global _DISPATCH_SHADOW_LAST_LOG - if is_overloaded(): - return - now = time.monotonic() - if now - _DISPATCH_SHADOW_LAST_LOG < DISPATCH_STATS_LOG_INTERVAL: - return - _DISPATCH_SHADOW_LAST_LOG = now - text_hint = (context.plain_text or context.trie_raw_command or "")[:48] - _debug_log( - ( - "dispatch shadow: " - f"priority={priority} " - f"event={context.event_type} " - f"selected={activation_result.candidate_count}/" - f"{activation_result.total_descriptors} " - f"selected_lane={_compact_counter(activation_result.selected_by_lane)} " - f"skipped_lane={_compact_counter(activation_result.skipped_by_lane)} " - f"selected_reason=" - f"{_compact_counter(activation_result.selected_by_reason)} " - f"skipped_reason=" - f"{_compact_counter(activation_result.skipped_by_reason)} " - f"text={text_hint!r}" - ), - LOGGER_COMMAND, - ) - - def _auth_scope_key(context: EventContext) -> str: group_id = context.group_id or "" channel_id = context.channel_id or "" @@ -876,7 +727,6 @@ async def _dispatch_lane_section(lane: str): wait_ms = (time.perf_counter() - started) * 1000 if wait_ms >= AUTH_OVERLOAD_LANE_WAIT_MS: signal_overload(2.0) - _record_dispatch_selection(lane, True, wait_ms=wait_ms) try: yield finally: @@ -891,11 +741,6 @@ def get_dispatch_snapshot() -> dict[str, object]: value = getattr(semaphore, "_value", limit) lane_active[lane] = max(limit - int(value), 0) return { - "selected": _DISPATCH_SELECTED, - "skipped": _DISPATCH_SKIPPED, - "selected_by_lane": dict(_DISPATCH_SELECTED_BY_LANE), - "skipped_by_lane": dict(_DISPATCH_SKIPPED_BY_LANE), - "lane_wait_ms": dict(_DISPATCH_LANE_WAIT_MS), "lane_active": lane_active, "lane_limits": dict(_DISPATCH_LANE_LIMITS), } @@ -975,58 +820,6 @@ async def _run_selected_matcher( ) -def _record_prefilter_stats( - skipped: bool, - reason: str | None, - stage: str = "inside_task", -) -> None: - global _PREFILTER_LAST_LOG - _PREFILTER_STATS["checked"] += 1 - if skipped: - _PREFILTER_STATS["skipped"] += 1 - if stage == "before_task": - _PREFILTER_STATS["before_task_checked"] += 1 - if skipped: - _PREFILTER_STATS["before_task_skipped"] += 1 - else: - _PREFILTER_STATS["inside_task_checked"] += 1 - if skipped: - _PREFILTER_STATS["inside_task_skipped"] += 1 - if reason == "type_miss": - _PREFILTER_STATS["type_miss"] += 1 - elif reason == "route_miss": - _PREFILTER_STATS["route_miss"] += 1 - elif reason == "command_miss": - _PREFILTER_STATS["command_miss"] += 1 - elif reason == "empty_text": - _PREFILTER_STATS["empty_text"] += 1 - - if _PREFILTER_STATS["checked"] % 1024 == 0: - with contextlib.suppress(Exception): - _ = len(_CHECK_MATCHER_ROUTE_CACHE) - - now = time.monotonic() - if now - _PREFILTER_LAST_LOG < PREFILTER_STATS_LOG_INTERVAL or is_overloaded(): - return - _PREFILTER_LAST_LOG = now - _debug_log( - ( - "matcher prefilter stats: " - f"checked={_PREFILTER_STATS['checked']} " - f"skipped={_PREFILTER_STATS['skipped']} " - f"before_task={_PREFILTER_STATS['before_task_skipped']}/" - f"{_PREFILTER_STATS['before_task_checked']} " - f"inside_task={_PREFILTER_STATS['inside_task_skipped']}/" - f"{_PREFILTER_STATS['inside_task_checked']} " - f"type_miss={_PREFILTER_STATS['type_miss']} " - f"route_miss={_PREFILTER_STATS['route_miss']} " - f"command_miss={_PREFILTER_STATS['command_miss']} " - f"empty_text={_PREFILTER_STATS['empty_text']}" - ), - LOGGER_COMMAND, - ) - - _MAX_MATCHER_CACHE = 512 @@ -1038,8 +831,6 @@ _SELECTOR_DEPS = HandleEventSelectorDependencies( activation_context_from_dispatch=_activation_context_from_dispatch, new_dispatch_budget=_new_dispatch_budget, dispatch_lane_for_matcher=_dispatch_lane_for_matcher, - record_activation_result=_record_activation_result, - debug_activation_shadow=_debug_activation_shadow, merge_dispatch_budget=_merge_dispatch_budget, build_matcher_state=_build_matcher_state, run_selected_matcher=_run_selected_matcher, @@ -1068,15 +859,6 @@ async def _get_route_context(text: str, event_cache: dict | None) -> set[str]: return matched -def _get_auth_route_precheck_deps() -> dict: - return { - "route_modules_with_commands": _ROUTE_MODULES_WITH_COMMANDS, - "get_route_context": _get_route_context, - "is_command_matcher_class": _is_command_matcher_class, - "matcher_has_alconna_shortcuts": _matcher_has_alconna_shortcuts, - } - - async def _cache_sweep_loop() -> None: while True: await asyncio.sleep(CACHE_SWEEP_INTERVAL) @@ -1145,7 +927,8 @@ async def _has_limits_cached( if event_cache is not None: event_cache.setdefault("module_limits_ready", {})[module] = True return has_limits - limit_entries = PluginLimitMemoryCache.get_limits_if_ready(module) + provider = DEFAULT_PERMISSION_DATA_PROVIDER + limit_entries = provider.get_module_limits_if_ready(module) if limit_entries is not None: has_limits = bool(limit_entries) module_limit_cache[module] = has_limits @@ -1278,16 +1061,17 @@ async def _get_plugin_cache_first( *, allow_cache_load: bool, ) -> tuple[PluginInfo | None, bool]: + provider = DEFAULT_PERMISSION_DATA_PROVIDER plugin = None if event_cache is not None: plugin_cache = event_cache.setdefault("plugin_cache", {}) if module in plugin_cache: return cast(PluginInfo | None, plugin_cache[module]), False - plugin = PluginInfoMemoryCache.get_by_module_if_ready(module) - cache_miss = plugin is None and not PluginInfoMemoryCache.is_loaded() + plugin = provider.get_plugin_if_ready(module) + cache_miss = plugin is None and not provider.plugin_cache_loaded() if plugin is None and allow_cache_load: - plugin = await PluginInfoMemoryCache.get_by_module(module) + plugin = await provider.get_plugin(module) cache_miss = False if event_cache is not None: event_cache.setdefault("plugin_cache", {})[module] = plugin @@ -1351,7 +1135,6 @@ async def reserve_gold( ) except InsufficientGold: raise - await DataAccess(UserConsole).clear_cache(user_id=user_id) logger.debug(f"预扣功能花费金币: {cost_gold}", LOGGER_COMMAND, session=session) return reservation @@ -1384,14 +1167,12 @@ async def _record_backpressure( action: str, duration_ms: float = 0.0, ) -> None: - await append_runtime_backpressure_log( - scope_key=lane_context.scope_key, - reason=reason, - lane=lane_context.lane, - action=action, - queue_size=lane_context.queue_size, - active_count=HOOKS_ACTIVE_COUNT, - duration_ms=duration_ms, + logger.debug( + "auth backpressure: " + f"scope={lane_context.scope_key}, lane={lane_context.lane}, " + f"reason={reason}, action={action}, queue={lane_context.queue_size}, " + f"active={HOOKS_ACTIVE_COUNT}, duration_ms={duration_ms:.1f}", + LOGGER_COMMAND, ) @@ -1437,7 +1218,6 @@ async def _prepare_auth_state( context: EventContext, bot: Bot, event_cache: dict | None, - route_skip_checks: bool, skip_ban: bool, hook_recorder: HookTraceRecorder, state: dict | None, @@ -1502,7 +1282,6 @@ async def _prepare_auth_state( policy_context = PolicyContext( snapshot=snapshot, - route_skip_checks=route_skip_checks, allow_sleep_bypass=_is_bot_wake_command(module, context.plain_text), allow_group_sleep_bypass=_is_group_wake_command(plugin, context.plain_text), ) @@ -1522,7 +1301,6 @@ async def _prepare_auth_state_with_fallback( context: EventContext, bot: Bot, event_cache: dict | None, - route_skip_checks: bool, skip_ban: bool, hook_recorder: HookTraceRecorder, state: dict | None, @@ -1533,7 +1311,6 @@ async def _prepare_auth_state_with_fallback( context=context, bot=bot, event_cache=event_cache, - route_skip_checks=route_skip_checks, skip_ban=skip_ban, hook_recorder=hook_recorder, state=state, @@ -1548,7 +1325,6 @@ async def _prepare_auth_state_with_fallback( context=context, bot=bot, event_cache=event_cache, - route_skip_checks=route_skip_checks, skip_ban=skip_ban, hook_recorder=hook_recorder, state=state, @@ -1614,16 +1390,26 @@ async def _reserve_limit_side_effect( async def _resolve_cost_gold( *, prep: AuthPreparation, - route_skip_checks: bool, hook_recorder: HookTraceRecorder, session: Uninfo, ) -> int: plugin = prep.plugin - if route_skip_checks or prep.profile.cost_gold <= 0: + if prep.profile.cost_gold <= 0: hook_recorder.set("cost_gold", "skipped") return 0 cost_start = time.time() try: + if prep.user is None: + user_start = time.time() + prep.user = await with_timeout( + UserConsole.get_user( + prep.permission_context.user_id, + PlatformUtils.get_platform(session), + ), + name="get_cost_user", + ) + prep.permission_context.user = prep.user + hook_recorder.set("get_cost_user", f"{time.time() - user_start:.3f}s") cost_gold = await with_timeout( get_plugin_cost( prep.user, @@ -1649,7 +1435,6 @@ async def _run_auth_hooks( prep: AuthPreparation, session: Uninfo, event_cache: dict | None, - route_skip_checks: bool, lane_context: AuthLaneContext, hook_recorder: HookTraceRecorder, side_effect_commit: SideEffectCommit, @@ -1660,26 +1445,23 @@ async def _run_auth_hooks( await _enter_hooks_section(lane_context) hook_tasks = [] try: - if not route_skip_checks: - has_limits = await _has_limits_cached( - profile.module, - event_cache, - known=profile.has_limit, - ) - if has_limits: - hook_tasks.append( - time_hook( - _reserve_limit_side_effect( - prep=prep, - session=session, - side_effect_commit=side_effect_commit, - ), - "auth_limit", - hook_recorder, - ) + has_limits = await _has_limits_cached( + profile.module, + event_cache, + known=profile.has_limit, + ) + if has_limits: + hook_tasks.append( + time_hook( + _reserve_limit_side_effect( + prep=prep, + session=session, + side_effect_commit=side_effect_commit, + ), + "auth_limit", + hook_recorder, ) - else: - hook_recorder.set("auth_limit", "skipped") + ) else: hook_recorder.set("auth_limit", "skipped") @@ -1703,13 +1485,6 @@ async def _run_auth_hooks( return time.time() - hooks_start -async def build_auth_decision_backpressure_report( - *, - hours: float = 24.0, -) -> dict: - return await build_auth_observability_report(hours=hours) - - _AUTH_PIPELINE_DEPS = AuthPipelineDependencies( route_modules_with_commands=_ROUTE_MODULES_WITH_COMMANDS, get_route_context=_get_route_context, @@ -1726,7 +1501,6 @@ _AUTH_PIPELINE_DEPS = AuthPipelineDependencies( run_auth_hooks=_run_auth_hooks, bot_filter=bot_filter, reserve_gold=reserve_gold, - append_auth_decision_log=append_auth_decision_log, insufficient_gold_error=InsufficientGold, logger=logger, log_command=LOGGER_COMMAND, diff --git a/zhenxun/builtin_plugins/hooks/auth_event_selector.py b/zhenxun/builtin_plugins/hooks/auth_event_selector.py index 811c7ebd..c01fec31 100644 --- a/zhenxun/builtin_plugins/hooks/auth_event_selector.py +++ b/zhenxun/builtin_plugins/hooks/auth_event_selector.py @@ -31,8 +31,6 @@ class HandleEventSelectorDependencies: activation_context_from_dispatch: Callable[[EventDispatchContext, Event], Any] new_dispatch_budget: Callable[[], dict[str, int]] dispatch_lane_for_matcher: Callable[[type[Matcher], EventDispatchContext], str] - record_activation_result: Callable[[Any], None] - debug_activation_shadow: Callable[..., None] merge_dispatch_budget: Callable[[dict[str, int], dict[str, int]], None] build_matcher_state: Callable[[dict], dict] run_selected_matcher: Callable[..., Awaitable[None]] @@ -43,11 +41,81 @@ _ORIGINAL_HANDLE_EVENT: Callable[..., Awaitable[None]] | None = None _ORIGINAL_ADAPTER_HANDLE_EVENTS: dict[object, object] = {} +def _trim_leading_text(message: Any) -> None: + if not message: + return + segment = message[0] + if getattr(segment, "type", None) != "text": + return + data = getattr(segment, "data", None) + if not isinstance(data, dict): + return + data["text"] = str(data.get("text", "")).lstrip("\xa0").lstrip() + if not data["text"]: + del message[0] + + +def _is_self_mention_segment(segment: Any, bot: Bot) -> bool: + segment_type = getattr(segment, "type", None) + if segment_type not in {"mention_user", "group_mention_user"}: + return False + data = getattr(segment, "data", None) + if not isinstance(data, dict): + return False + if data.get("is_you") or data.get("is_bot"): + return True + user_id = data.get("user_id") + return user_id is not None and str(user_id) == str(bot.self_id) + + +def _ensure_nonempty_qq_message(message: Any) -> None: + if message: + return + with contextlib.suppress(Exception): + message_module = importlib.import_module("nonebot.adapters.qq.message") + MessageSegment = getattr(message_module, "MessageSegment") + message.append(MessageSegment.text("")) + + +def _normalize_qq_self_at_message(bot: Bot, event: Event) -> None: + """Remove the leading bot mention left by QQ official @ events. + + nonebot-adapter-qq's @ event branches can mark ``to_me`` but keep the + synthetic leading mention segment. Alconna command heads then see + ``<@bot>命令`` and fail to match, while regular ``event.get_plaintext()`` + still looks correct. Normalizing here keeps the runtime behavior aligned + with OneBot/standard to_me preprocessing without changing plugin code or + database state. + """ + if event.__class__.__name__ not in { + "AtMessageCreateEvent", + "GroupAtMessageCreateEvent", + }: + return + adapter = getattr(bot, "adapter", None) + adapter_name = "" + get_name = getattr(adapter, "get_name", None) + if callable(get_name): + with contextlib.suppress(Exception): + adapter_name = str(get_name()).lower() + if adapter_name != "qq": + return + with contextlib.suppress(Exception): + message = event.get_message() + if not message or not _is_self_mention_segment(message[0], bot): + return + message.pop(0) + setattr(event, "to_me", True) + _trim_leading_text(message) + _ensure_nonempty_qq_message(message) + + async def patched_handle_event( bot: Bot, event: Event, deps: HandleEventSelectorDependencies, ) -> None: + _normalize_qq_self_at_message(bot, event) show_log = True escape_tag = getattr(nb_message, "escape_tag") logger_ = getattr(nb_message, "logger") @@ -153,12 +221,6 @@ async def patched_handle_event( if activation_result is not None: selected_matchers = activation_result.selected - deps.record_activation_result(activation_result) - deps.debug_activation_shadow( - priority=priority, - activation_result=activation_result, - context=dispatch_context, - ) if ( activation_result.candidate_count > deps.overload_selected_threshold @@ -186,12 +248,6 @@ async def patched_handle_event( except Exception: single_result = None if single_result is not None: - deps.record_activation_result(single_result) - deps.debug_activation_shadow( - priority=priority, - activation_result=single_result, - context=dispatch_context, - ) deps.merge_dispatch_budget( priority_budget, single_budget, @@ -238,6 +294,7 @@ def install_handle_event_selector(deps: HandleEventSelectorDependencies) -> None for module_name in ( "nonebot.adapters.onebot.v11.bot", "nonebot.adapters.onebot.v12.bot", + "nonebot.adapters.qq.bot", "onebug.mixin.process", ): with contextlib.suppress(Exception): diff --git a/zhenxun/builtin_plugins/hooks/auth_hook.py b/zhenxun/builtin_plugins/hooks/auth_hook.py index c6fe0566..2d61d56b 100644 --- a/zhenxun/builtin_plugins/hooks/auth_hook.py +++ b/zhenxun/builtin_plugins/hooks/auth_hook.py @@ -26,13 +26,11 @@ from .auth.context import ( ) from .auth_checker import ( LimitManager, - _get_auth_route_precheck_deps, _get_route_context, auth, start_auth_runtime_tasks, stop_auth_runtime_tasks, ) -from .auth_route import route_precheck _SKIP_AUTH_PLUGINS = {"chat_history", "chat_message"} _BOT_CONNECT_TS: float | None = None @@ -113,9 +111,6 @@ async def _auth_preprocessor( ) set_route_modules(state, event_context, route_modules) - if await route_precheck(matcher, event_context, **_get_auth_route_precheck_deps()): - return - try: await auth( matcher, diff --git a/zhenxun/builtin_plugins/hooks/auth_pipeline.py b/zhenxun/builtin_plugins/hooks/auth_pipeline.py index a5513089..6c27a9b4 100644 --- a/zhenxun/builtin_plugins/hooks/auth_pipeline.py +++ b/zhenxun/builtin_plugins/hooks/auth_pipeline.py @@ -10,7 +10,6 @@ from nonebot.adapters import Bot, Event from nonebot.matcher import Matcher from nonebot_plugin_uninfo import Uninfo -from zhenxun.services.message_load import is_overloaded from zhenxun.utils.utils import EntityIDs from .auth.context import ( @@ -86,7 +85,6 @@ class AuthPipelineContext: event_cache: dict | None = None text: str = "" route_modules: set[str] | None = None - route_skip_checks: bool = False is_command_matcher: bool = False lane_context: AuthLaneContext | None = None side_effect_cache: PermissionSideEffectCache | None = None @@ -149,7 +147,6 @@ class AuthPipelineDependencies: run_auth_hooks: Callable[..., Awaitable[float]] bot_filter: Callable[..., None] reserve_gold: Callable[..., Awaitable[Any]] - append_auth_decision_log: Callable[..., Awaitable[None]] insufficient_gold_error: type[Exception] logger: Any log_command: str @@ -173,10 +170,7 @@ def apply_policy_precheck( hook_recorder.set("auth_core", f"policy:{decision.reason}") if decision.denied: raise_for_policy(decision, deps.policy_skip_message(decision.reason)) - if decision.allowed and decision.reason in { - "hidden_plugin_skip_auth", - "route_miss_skip_checks", - }: + if decision.allowed and decision.reason in {"hidden_plugin_skip_auth"}: flags.should_return_allowed = True return flags @@ -260,17 +254,16 @@ async def route_gate_stage( if ctx.route_modules is None: ctx.route_modules = await deps.get_route_context(ctx.text, ctx.event_cache) set_route_modules(ctx.state, ctx.event_context, ctx.route_modules) - ctx.route_skip_checks = ( + route_missed = ( ctx.is_command_matcher and ctx.module in deps.route_modules_with_commands and ctx.module not in ctx.route_modules and not deps.matcher_has_alconna_shortcuts(type(ctx.matcher)) ) - if ctx.route_skip_checks: + if route_missed: if ctx.event_cache is not None: - ctx.event_cache["route_skip"] = True + ctx.event_cache["route_miss_after_native_match"] = True _recorder(ctx).set("route", "miss") - ctx.stop(allowed=True, effect="allow", reason="route_miss_skip_checks") async def prepare_snapshot_stage( @@ -282,7 +275,6 @@ async def prepare_snapshot_stage( context=ctx.event_context, bot=ctx.bot, event_cache=ctx.event_cache, - route_skip_checks=ctx.route_skip_checks, skip_ban=ctx.skip_ban, hook_recorder=ctx.hook_recorder, state=ctx.state, @@ -305,7 +297,6 @@ async def policy_precheck_stage( context=ctx.event_context, bot=ctx.bot, event_cache=ctx.event_cache, - route_skip_checks=ctx.route_skip_checks, skip_ban=ctx.skip_ban, hook_recorder=ctx.hook_recorder, state=ctx.state, @@ -340,7 +331,6 @@ async def policy_precheck_stage( ) ctx.cost_gold = await deps.resolve_cost_gold( prep=ctx.prep, - route_skip_checks=ctx.route_skip_checks, hook_recorder=ctx.hook_recorder, session=ctx.session, ) @@ -356,7 +346,6 @@ async def legacy_hook_adapter_stage( prep=prep, session=ctx.session, event_cache=ctx.event_cache, - route_skip_checks=ctx.route_skip_checks, lane_context=_lane_context(ctx), hook_recorder=_recorder(ctx), side_effect_commit=_side_effect_commit(ctx), @@ -425,35 +414,12 @@ async def decision_log_stage( ctx.auth_allowed, None if ctx.auth_allowed else ctx.decision_reason, ) - side_effect_state = commit.snapshot() if commit is not None else None - shadow_effect = None - shadow_reason = None - if has_deferred_commit: - shadow_effect = "defer" - shadow_reason = "side_effect_pending:" + ",".join( - commit.pending_kinds if commit is not None else () - ) if ctx.entered_side_effect_lock and ctx.side_effect_lock is not None: try: ctx.side_effect_lock.release() except Exception: pass ctx.entered_side_effect_lock = False - latency_ms = (time.time() - ctx.start_time) * 1000 - await deps.append_auth_decision_log( - bot_id=ctx.event_context.bot_id, - platform=ctx.event_context.platform, - group_id=_entity(ctx).group_id, - user_id=_entity(ctx).user_id, - module=ctx.module, - effect=ctx.decision_effect or "error", - reason=ctx.decision_reason, - shadow_effect=shadow_effect, - shadow_reason=shadow_reason, - side_effect_state=side_effect_state, - latency_ms=latency_ms, - overloaded=is_overloaded(), - ) def build_auth_pipeline(deps: AuthPipelineDependencies) -> AuthPipeline: diff --git a/zhenxun/builtin_plugins/hooks/auth_policy.py b/zhenxun/builtin_plugins/hooks/auth_policy.py index 541d8d45..739ec166 100644 --- a/zhenxun/builtin_plugins/hooks/auth_policy.py +++ b/zhenxun/builtin_plugins/hooks/auth_policy.py @@ -60,7 +60,6 @@ class PolicyResource: @dataclass(frozen=True, slots=True) class PolicyContext: snapshot: AuthSnapshot - route_skip_checks: bool = False allow_sleep_bypass: bool = False allow_group_sleep_bypass: bool = False @@ -101,8 +100,6 @@ class PolicyDecisionPoint: profile = resource.profile if profile.hidden: return PolicyDecision("allow", "hidden_plugin_skip_auth") - if context.route_skip_checks: - return PolicyDecision("allow", "route_miss_skip_checks") if snapshot.ban_state is True and not principal.is_superuser: return PolicyDecision("deny", "user_or_group_banned") if profile.superuser_only and not principal.is_superuser: diff --git a/zhenxun/builtin_plugins/hooks/auth_profile.py b/zhenxun/builtin_plugins/hooks/auth_profile.py index d1c47c73..16128d2a 100644 --- a/zhenxun/builtin_plugins/hooks/auth_profile.py +++ b/zhenxun/builtin_plugins/hooks/auth_profile.py @@ -2,11 +2,13 @@ from __future__ import annotations from dataclasses import dataclass -from zhenxun.services.cache.runtime_cache import ( - PluginLimitMemoryCache, +from zhenxun.utils.enum import BlockType, PluginType + +from .auth.data_provider import ( + DEFAULT_PERMISSION_DATA_PROVIDER, + PermissionDataProvider, PluginLimitSnapshot, ) -from zhenxun.utils.enum import BlockType, PluginType @dataclass(frozen=True, slots=True) @@ -81,6 +83,7 @@ async def get_plugin_auth_profile( *, event_cache: dict | None = None, allow_cache_load: bool = True, + provider: PermissionDataProvider = DEFAULT_PERMISSION_DATA_PROVIDER, ) -> PluginAuthProfile: module = str(getattr(plugin, "module", "") or "") profile_cache: dict[str, PluginAuthProfile] = {} @@ -98,10 +101,10 @@ async def get_plugin_auth_profile( limits = limit_cache[module] limits_ready = True if limits is None: - limits = PluginLimitMemoryCache.get_limits_if_ready(module) + limits = provider.get_module_limits_if_ready(module) limits_ready = limits is not None if limits is None and allow_cache_load: - limits = await PluginLimitMemoryCache.get_limits(module) + limits = await provider.get_module_limits(module) limits_ready = True if limits is None: limits = [] diff --git a/zhenxun/builtin_plugins/hooks/auth_route.py b/zhenxun/builtin_plugins/hooks/auth_route.py deleted file mode 100644 index 88b9b43d..00000000 --- a/zhenxun/builtin_plugins/hooks/auth_route.py +++ /dev/null @@ -1,60 +0,0 @@ -from __future__ import annotations - -from collections.abc import Awaitable, Callable - -from nonebot.matcher import Matcher - -from zhenxun.utils.enum import PluginType - -from .auth.context import EventContext, set_route_modules - -RouteContextGetter = Callable[[str, dict | None], Awaitable[set[str]]] -CommandMatcherChecker = Callable[[type[Matcher]], bool] -AlconnaShortcutChecker = Callable[[type[Matcher]], bool] - - -async def route_precheck( - matcher: Matcher, - context: EventContext, - *, - route_modules_with_commands: set[str], - get_route_context: RouteContextGetter, - is_command_matcher_class: CommandMatcherChecker, - matcher_has_alconna_shortcuts: AlconnaShortcutChecker, -) -> bool: - """Skip expensive auth checks for command matchers proven to be off-route.""" - - module = matcher.plugin_name or "" - if not module: - return False - if _is_hidden_plugin(matcher): - return False - if not is_command_matcher_class(type(matcher)): - return False - - route_modules = context.route_modules if context.route_modules_loaded else None - if route_modules is None: - route_modules = await get_route_context( - context.plain_text, - context.event_cache, - ) - set_route_modules(None, context, route_modules) - - if module in route_modules_with_commands and module not in route_modules: - if matcher_has_alconna_shortcuts(type(matcher)): - return False - if context.event_cache is not None: - context.event_cache["route_skip"] = True - return True - return False - - -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 - - -__all__ = ["route_precheck"] diff --git a/zhenxun/builtin_plugins/hooks/auth_snapshot.py b/zhenxun/builtin_plugins/hooks/auth_snapshot.py index 4c79486a..9cc6e81c 100644 --- a/zhenxun/builtin_plugins/hooks/auth_snapshot.py +++ b/zhenxun/builtin_plugins/hooks/auth_snapshot.py @@ -1,24 +1,120 @@ from __future__ import annotations +import asyncio from dataclasses import dataclass, field +import time from typing import TYPE_CHECKING from zhenxun.services.cache.runtime_cache import ( - BanMemoryCache, - BotMemoryCache, BotSnapshot, - GroupMemoryCache, GroupSnapshot, - LevelUserMemoryCache, LevelUserSnapshot, ) +from zhenxun.services.log import logger +from .auth.config import LOGGER_COMMAND from .auth.context import EventContext +from .auth.data_provider import ( + DEFAULT_PERMISSION_DATA_PROVIDER, + PermissionDataProvider, +) from .auth_profile import PluginAuthProfile if TYPE_CHECKING: from nonebot.adapters import Bot +QQ_CLIENT_GROUP_REPAIR_TTL = 60 +_QQ_CLIENT_GROUP_REPAIR_FAILURES: dict[tuple[str, str], float] = {} +_QQ_CLIENT_GROUP_REPAIR_LOCKS: dict[tuple[str, str], asyncio.Lock] = {} + + +def _build_runtime_group_snapshot(context: EventContext) -> GroupSnapshot | None: + """Provide a non-persistent default group for QQ official runtime auth.""" + if context.platform_scope != "qq_api" or not context.group_id: + return None + return GroupSnapshot( + group_id=context.group_id, + channel_id=context.channel_id, + group_name="", + max_member_count=0, + member_count=0, + status=True, + level=5, + is_super=False, + group_flag=0, + block_plugin="", + superuser_block_plugin="", + block_task="", + superuser_block_task="", + platform=context.platform, + ) + + +def _qq_client_group_repair_key(context: EventContext) -> tuple[str, str] | None: + if context.platform_scope != "qq_client" or not context.group_id: + return None + return (context.group_id, context.channel_id or "") + + +def _qq_client_group_repair_on_cooldown(key: tuple[str, str]) -> bool: + expire_at = _QQ_CLIENT_GROUP_REPAIR_FAILURES.get(key) + if not expire_at: + return False + if expire_at <= time.time(): + _QQ_CLIENT_GROUP_REPAIR_FAILURES.pop(key, None) + return False + return True + + +async def _repair_missing_qq_client_group( + context: EventContext, + *, + provider: PermissionDataProvider, +) -> GroupSnapshot | None: + """Persist a minimal OneBot group when startup group sync returned empty.""" + key = _qq_client_group_repair_key(context) + if key is None or not provider.group_cache_loaded(): + return None + group_id, _ = key + if _qq_client_group_repair_on_cooldown(key): + return None + + try: + from zhenxun.models.group_console import GroupConsole + + lock = _QQ_CLIENT_GROUP_REPAIR_LOCKS.setdefault(key, asyncio.Lock()) + async with lock: + existing = provider.get_group_if_ready( + group_id, + context.channel_id, + ) + if existing is not None: + return existing + defaults = { + "group_name": "", + "max_member_count": 0, + "member_count": 0, + "group_flag": 1, + "platform": context.platform, + } + group, _ = await GroupConsole.get_or_create_root_group( + group_id=group_id, + defaults=defaults, + ) + from zhenxun.services.cache.runtime_cache import GroupMemoryCache + + await GroupMemoryCache.upsert_from_model(group) + return GroupSnapshot.from_model(group) + except Exception as exc: + _QQ_CLIENT_GROUP_REPAIR_FAILURES[key] = time.time() + QQ_CLIENT_GROUP_REPAIR_TTL + logger.warning( + "协议端群记录缺失自愈失败,已短期跳过重复修复", + LOGGER_COMMAND, + group_id=context.group_id, + e=exc, + ) + return None + @dataclass(slots=True) class AuthSnapshot: @@ -72,6 +168,7 @@ async def build_auth_snapshot( bot: "Bot", skip_ban: bool = False, allow_cache_load: bool = False, + provider: PermissionDataProvider = DEFAULT_PERMISSION_DATA_PROVIDER, ) -> AuthSnapshot: event_cache = context.event_cache entity = context.entity @@ -85,15 +182,15 @@ async def build_auth_snapshot( ): bot_data = event_cache.get("bot_data") else: - bot_data = BotMemoryCache.get_if_ready(bot.self_id) + bot_data = provider.get_bot_if_ready(bot.self_id) if bot_data is None: if allow_cache_load: - bot_data = await BotMemoryCache.get(bot.self_id) - elif not BotMemoryCache.is_loaded(): + bot_data = await provider.get_bot(bot.self_id) + elif not provider.bot_cache_loaded(): cache_misses.add("bot") if event_cache is not None: event_cache["bot_data"] = bot_data - event_cache["bot_cache_ready"] = BotMemoryCache.is_loaded() + event_cache["bot_cache_ready"] = provider.bot_cache_loaded() group = None if entity.group_id: @@ -104,14 +201,31 @@ async def build_auth_snapshot( ): group = event_cache.get("group") else: - group = GroupMemoryCache.get_if_ready(entity.group_id, entity.channel_id) - if group is None and not GroupMemoryCache.is_loaded(): + group = provider.get_group_if_ready(entity.group_id, entity.channel_id) + if group is None and not provider.group_cache_loaded(): cache_misses.add("group") elif group is None and allow_cache_load: - group = await GroupMemoryCache.get(entity.group_id, entity.channel_id) + group = await provider.get_group(entity.group_id, entity.channel_id) if event_cache is not None: event_cache["group"] = group - event_cache["group_cache_ready"] = GroupMemoryCache.is_loaded() + event_cache["group_cache_ready"] = provider.group_cache_loaded() + if group is None: + group = await _repair_missing_qq_client_group( + context, + provider=provider, + ) + if group is None and (runtime_group := _build_runtime_group_snapshot(context)): + group = runtime_group + cache_misses.discard("group") + if event_cache is not None: + event_cache["group"] = group + event_cache["group_cache_ready"] = True + event_cache["group_runtime_virtual"] = True + elif group is not None: + cache_misses.discard("group") + if event_cache is not None: + event_cache["group"] = group + event_cache["group_cache_ready"] = True admin_levels = None if profile.need_admin: @@ -122,13 +236,13 @@ async def build_auth_snapshot( ): admin_levels = event_cache.get("admin_levels") else: - admin_levels = LevelUserMemoryCache.get_levels_if_ready( + admin_levels = provider.get_admin_levels_if_ready( entity.user_id, entity.group_id, ) if admin_levels is None: if allow_cache_load: - admin_levels = await LevelUserMemoryCache.get_levels( + admin_levels = await provider.get_admin_levels( entity.user_id, entity.group_id, ) @@ -136,19 +250,19 @@ async def build_auth_snapshot( cache_misses.add("admin_levels") if event_cache is not None: event_cache["admin_levels"] = admin_levels - event_cache["admin_cache_ready"] = LevelUserMemoryCache.is_loaded() + event_cache["admin_cache_ready"] = provider.admin_cache_loaded() ban_state = None if not skip_ban: if event_cache is not None and "ban_state" in event_cache: ban_state = event_cache.get("ban_state") - elif BanMemoryCache.is_loaded(): - ban_state = BanMemoryCache.is_banned(entity.user_id, entity.group_id) + elif provider.ban_cache_loaded(): + ban_state = provider.is_banned(entity.user_id, entity.group_id) if event_cache is not None: event_cache["ban_state"] = ban_state elif allow_cache_load: - await BanMemoryCache.ensure_loaded() - ban_state = BanMemoryCache.is_banned(entity.user_id, entity.group_id) + await provider.ensure_ban_loaded() + ban_state = provider.is_banned(entity.user_id, entity.group_id) if event_cache is not None: event_cache["ban_state"] = ban_state else: @@ -174,6 +288,7 @@ async def get_or_build_auth_snapshot( bot: "Bot", skip_ban: bool = False, allow_cache_load: bool = False, + provider: PermissionDataProvider = DEFAULT_PERMISSION_DATA_PROVIDER, ) -> AuthSnapshot: event_cache = context.event_cache module = profile.module @@ -190,6 +305,7 @@ async def get_or_build_auth_snapshot( bot=bot, skip_ban=skip_ban, allow_cache_load=allow_cache_load, + provider=provider, ) if event_cache is not None: event_cache.setdefault("auth_snapshots", {})[module] = snapshot diff --git a/zhenxun/builtin_plugins/hooks/call_hook.py b/zhenxun/builtin_plugins/hooks/call_hook.py index ac40fa37..6fbf251e 100644 --- a/zhenxun/builtin_plugins/hooks/call_hook.py +++ b/zhenxun/builtin_plugins/hooks/call_hook.py @@ -1,11 +1,9 @@ +from collections.abc import Mapping from typing import Any from nonebot.adapters import Bot, Message -from zhenxun.configs.config import Config -from zhenxun.models.bot_message_store import BotMessageStore from zhenxun.services.log import logger -from zhenxun.utils.enum import BotSentType from zhenxun.utils.log_sanitizer import sanitize_for_logging from zhenxun.utils.manager.message_manager import MessageManager from zhenxun.utils.platform import PlatformUtils @@ -45,13 +43,15 @@ def replace_message(message: Message) -> str: async def handle_api_result( bot: Bot, exception: Exception | None, api: str, data: dict[str, Any], result: Any ): - if exception or api != "send_msg": + if ( + exception + or api != "send_msg" + or PlatformUtils.get_platform_scope(bot) != "qq_client" + ): return user_id = data.get("user_id") - group_id = data.get("group_id") - message_id = result.get("message_id") + message_id = result.get("message_id") if isinstance(result, Mapping) else None message: Message = data.get("message", "") - message_type = data.get("message_type") try: if user_id and message_id: MessageManager.add(str(user_id), str(message_id)) @@ -62,27 +62,5 @@ async def handle_api_result( logger.warning( f"收集消息id发生错误...data: {data}, result: {result}", LOG_COMMAND, e=e ) - if not Config.get_config("hook", "RECORD_BOT_SENT_MESSAGES"): - return - try: - await BotMessageStore.append_buffered( - bot_id=bot.self_id, - user_id=user_id, - group_id=group_id, - sent_type=BotSentType.GROUP - if message_type == "group" - else BotSentType.PRIVATE, - text=replace_message(message), - plain_text=message.extract_plain_text() - if isinstance(message, Message) - else replace_message(message), - platform=PlatformUtils.get_platform(bot), - ) - sanitized_message = sanitize_for_logging(message, context="nonebot_message") - logger.debug(f"消息发送记录,message: {sanitized_message}") - except Exception as e: - logger.warning( - f"消息发送记录发生错误...data: {data}, result: {result}", - LOG_COMMAND, - e=e, - ) + sanitized_message = sanitize_for_logging(message, context="nonebot_message") + logger.debug(f"消息发送记录,message: {sanitized_message}") diff --git a/zhenxun/builtin_plugins/init/__init__.py b/zhenxun/builtin_plugins/init/__init__.py index b5574233..241a4d68 100644 --- a/zhenxun/builtin_plugins/init/__init__.py +++ b/zhenxun/builtin_plugins/init/__init__.py @@ -30,7 +30,7 @@ async def _(bot: Bot): 参数: bot: Bot """ - if PlatformUtils.get_platform(bot) != "qq": + if PlatformUtils.get_platform_scope(bot) != "qq_client": return logger.debug(f"更新Bot: {bot.self_id} 的群认证...", "群认证同步") @@ -44,6 +44,13 @@ async def _(bot: Bot): ) return + if not current_group_list: + logger.warning( + f"Bot: {bot.self_id} 未获取到任何群组," + "本次不会创建群认证;后续群消息将尝试按事件自愈。", + "群认证同步", + ) + db_group_list: list[str] = await GroupConsole.all().values_list( "group_id", flat=True ) # pyright: ignore[reportAssignmentType] diff --git a/zhenxun/builtin_plugins/init/init_task.py b/zhenxun/builtin_plugins/init/init_task.py index 1606e41c..8d78eda3 100644 --- a/zhenxun/builtin_plugins/init/init_task.py +++ b/zhenxun/builtin_plugins/init/init_task.py @@ -8,7 +8,7 @@ from nonebot_plugin_apscheduler import scheduler from zhenxun.configs.utils import PluginExtraData, Task from zhenxun.models.group_console import GroupConsole from zhenxun.models.task_info import TaskInfo -from zhenxun.services.cache.runtime_cache import TaskInfoMemoryCache +from zhenxun.services.cache.runtime_cache import GroupMemoryCache, TaskInfoMemoryCache from zhenxun.services.log import logger from zhenxun.utils.common_utils import CommonUtils from zhenxun.utils.manager.priority_manager import PriorityLifecycle @@ -63,6 +63,8 @@ async def update_to_group(create_list: list[tuple[bool, TaskInfo]]): ) group.block_task = CommonUtils.convert_module_format(block_tasks) await GroupConsole.bulk_update(group_list, ["block_task"], 10) + for group in group_list: + await GroupMemoryCache.upsert_from_model(group) async def to_db( @@ -146,8 +148,10 @@ async def _(): for plugin in get_loaded_plugins(): await _handle_setting(plugin, task_info_list, task_list) if not task_info_list: - await TaskInfo.all().update(load_status=False) - await TaskInfoMemoryCache.refresh() + logger.warning( + "未扫描到任何被动技能,跳过 TaskInfo.load_status 全量关闭," + "避免插件加载异常时误关闭全部被动技能。", + ) return module_dict = {t[1]: t[0] for t in await TaskInfo.all().values_list("id", "module")} load_task = [] diff --git a/zhenxun/builtin_plugins/platform/__init__.py b/zhenxun/builtin_plugins/platform/__init__.py index 8ded37ef..e8c006d1 100644 --- a/zhenxun/builtin_plugins/platform/__init__.py +++ b/zhenxun/builtin_plugins/platform/__init__.py @@ -2,6 +2,7 @@ from pathlib import Path import nonebot +from zhenxun.configs.config import BotConfig from zhenxun.services.log import logger path = Path(__file__).parent @@ -15,11 +16,12 @@ except ImportError: logger.warning("未安装 onebot-adapter,无法加载QQ平台专用插件...") -try: - from nonebot.adapters.qq import ( # noqa: F401 # pyright: ignore [reportMissingImports] - Bot, - ) +if BotConfig.qq_adapter_load: + try: + from nonebot.adapters.qq import ( # noqa: F401 # pyright: ignore [reportMissingImports] + Bot, + ) - nonebot.load_plugins(str((path / "qq_api").resolve())) -except ImportError: - logger.warning("未安装 qq-adapter,无法加载QQ官平台专用插件...") + nonebot.load_plugins(str((path / "qq_api").resolve())) + except ImportError: + logger.warning("未安装 qq-adapter,无法加载QQ官平台专用插件...") diff --git a/zhenxun/builtin_plugins/platform/qq/group_handle/__init__.py b/zhenxun/builtin_plugins/platform/qq/group_handle/__init__.py index 4ddcc47e..06be733a 100644 --- a/zhenxun/builtin_plugins/platform/qq/group_handle/__init__.py +++ b/zhenxun/builtin_plugins/platform/qq/group_handle/__init__.py @@ -108,9 +108,7 @@ async def _( ): if session.user.id == bot.self_id: """新成员为bot本身""" - group, _ = await GroupConsole.get_or_create( - group_id=str(event.group_id), channel_id__isnull=True - ) + group, _ = await GroupConsole.get_or_create_root_group(str(event.group_id)) try: await GroupManager.add_bot( bot, str(event.operator_id), str(event.group_id), group 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 751dfb33..0981ff01 100644 --- a/zhenxun/builtin_plugins/platform/qq/group_handle/data_source.py +++ b/zhenxun/builtin_plugins/platform/qq/group_handle/data_source.py @@ -124,7 +124,7 @@ class GroupManager: group_id=group_id, ) return - await GroupConsole.update_or_create( + await GroupConsole.get_or_create_root_group( group_id=group_info["group_id"], defaults={ "group_name": group_info["group_name"], @@ -134,6 +134,7 @@ class GroupManager: "block_plugin": block_plugin, "platform": "qq", }, + update_defaults=True, ) @classmethod diff --git a/zhenxun/builtin_plugins/platform/qq_api/ug_watch.py b/zhenxun/builtin_plugins/platform/qq_api/ug_watch.py index a27e9fb5..329188db 100644 --- a/zhenxun/builtin_plugins/platform/qq_api/ug_watch.py +++ b/zhenxun/builtin_plugins/platform/qq_api/ug_watch.py @@ -1,14 +1,13 @@ +"""QQ official platform observer. + +Official QQ identifiers are not in the same namespace as OneBot QQ numbers. +This observer intentionally avoids writing legacy identity tables; runtime auth +uses a non-persistent group snapshot when needed. +""" + from nonebot import on_message from nonebot_plugin_uninfo import Uninfo -from zhenxun.models.friend_user import FriendUser -from zhenxun.models.group_console import GroupConsole -from zhenxun.models.group_member_info import GroupInfoUser -from zhenxun.services.hot_query_cache import ( - invalidate_group_members, - invalidate_member_names, -) -from zhenxun.services.log import logger from zhenxun.utils.platform import PlatformUtils @@ -20,21 +19,5 @@ _matcher = on_message(priority=999, block=False, rule=rule) @_matcher.handle() -async def _(session: Uninfo): - platform = PlatformUtils.get_platform(session) - if session.group: - if not await GroupConsole.exists(group_id=session.group.id): - await GroupConsole.create(group_id=session.group.id) - logger.info("添加当前群组ID信息", session=session) - await GroupInfoUser.update_or_create( - user_id=session.user.id, - group_id=session.group.id, - platform=PlatformUtils.get_platform(session), - ) - await invalidate_group_members(session.group.id, [session.user.id]) - await invalidate_member_names([session.user.id]) - elif not await FriendUser.exists(user_id=session.user.id, platform=platform): - await FriendUser.create( - user_id=session.user.id, platform=PlatformUtils.get_platform(session) - ) - logger.info("添加当前好友用户信息", "", session=session) +async def _(): + return diff --git a/zhenxun/builtin_plugins/record_request.py b/zhenxun/builtin_plugins/record_request.py index 6c2034d0..5683b6f8 100644 --- a/zhenxun/builtin_plugins/record_request.py +++ b/zhenxun/builtin_plugins/record_request.py @@ -203,7 +203,7 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio session=event.user_id, target=event.group_id, ) - group, _ = await GroupConsole.update_or_create( + group, _ = await GroupConsole.get_or_create_root_group( group_id=str(event.group_id), defaults={ "group_name": "", @@ -211,6 +211,7 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio "member_count": 0, "group_flag": 1, }, + update_defaults=True, ) await bot.set_group_add_request( flag=event.flag, sub_type="invite", approve=True diff --git a/zhenxun/builtin_plugins/scheduler/auto_update_group.py b/zhenxun/builtin_plugins/scheduler/auto_update_group.py index bae378ad..2f98ceb2 100644 --- a/zhenxun/builtin_plugins/scheduler/auto_update_group.py +++ b/zhenxun/builtin_plugins/scheduler/auto_update_group.py @@ -18,6 +18,8 @@ async def _(): return bots = nonebot.get_bots() for bot in bots.values(): + if PlatformUtils.get_platform_scope(bot) != "qq_client": + continue try: await PlatformUtils.update_group(bot) except Exception as e: @@ -36,6 +38,8 @@ async def _(): return bots = nonebot.get_bots() for bot in bots.values(): + if PlatformUtils.get_platform_scope(bot) != "qq_client": + continue try: await PlatformUtils.update_friend(bot) except Exception as e: diff --git a/zhenxun/builtin_plugins/scheduler/chat_check.py b/zhenxun/builtin_plugins/scheduler/chat_check.py index 8e878486..ac5c17f8 100644 --- a/zhenxun/builtin_plugins/scheduler/chat_check.py +++ b/zhenxun/builtin_plugins/scheduler/chat_check.py @@ -52,8 +52,8 @@ async def _(): if last_message: now = datetime.now(pytz.timezone("Asia/Shanghai")) if now - timedelta(days=2) > last_message.create_time: - _group, _ = await GroupConsole.get_or_create( - group_id=group.group_id, channel_id__isnull=True + _group, _ = await GroupConsole.get_or_create_root_group( + group.group_id ) modules = [f"<{module}" for module in modules] _group.block_task = ",".join(modules) + "," # type: ignore diff --git a/zhenxun/builtin_plugins/superuser/group_manage.py b/zhenxun/builtin_plugins/superuser/group_manage.py index 7d42282c..9cbc914e 100644 --- a/zhenxun/builtin_plugins/superuser/group_manage.py +++ b/zhenxun/builtin_plugins/superuser/group_manage.py @@ -147,7 +147,7 @@ def CheckGroupId(): @_matcher.assign("modify-level", parameterless=[CheckGroupId()]) async def _(session: EventSession, arparma: Arparma, state: T_State, level: int): gid = state["group_id"] - group, _ = await GroupConsole.get_or_create(group_id=gid) + group, _ = await GroupConsole.get_or_create_root_group(gid) old_level = group.level group.level = level await group.save(update_fields=["level"]) @@ -176,10 +176,10 @@ async def _(session: EventSession, arparma: Arparma, state: T_State): @_matcher.assign("auth-handle", parameterless=[CheckGroupId()]) async def _(session: EventSession, arparma: Arparma, state: T_State): gid = state["group_id"] - await GroupConsole.update_or_create( + await GroupConsole.get_or_create_root_group( group_id=gid, - channel_id__isnull=True, defaults={"group_flag": 0 if arparma.find("delete") else 1}, + update_defaults=True, ) s = "删除" if arparma.find("delete") else "添加" await MessageUtils.build_message(f"{s}群认证成功!").send(reply_to=True) diff --git a/zhenxun/builtin_plugins/superuser/update_fg_info.py b/zhenxun/builtin_plugins/superuser/update_fg_info.py index 4932a976..2664db93 100644 --- a/zhenxun/builtin_plugins/superuser/update_fg_info.py +++ b/zhenxun/builtin_plugins/superuser/update_fg_info.py @@ -54,6 +54,11 @@ async def _( arparma: Arparma, ): try: + if PlatformUtils.get_platform_scope(bot) != "qq_client": + await MessageUtils.build_message( + "当前平台不支持旧群组信息同步,仅 OneBot 协议端可用。" + ).send() + return num = await PlatformUtils.update_group(bot) logger.info( f"更新群聊信息完成,共更新了 {num} 个群组的信息!", @@ -75,6 +80,11 @@ async def _( arparma: Arparma, ): try: + if PlatformUtils.get_platform_scope(bot) != "qq_client": + await MessageUtils.build_message( + "当前平台不支持旧好友信息同步,仅 OneBot 协议端可用。" + ).send() + return num = await PlatformUtils.update_friend(bot) logger.info( f"更新好友信息完成,共更新了 {num} 个好友的信息!", 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 89d8a6ce..94a23d05 100644 --- a/zhenxun/builtin_plugins/web_ui/api/tabs/manage/__init__.py +++ b/zhenxun/builtin_plugins/web_ui/api/tabs/manage/__init__.py @@ -213,9 +213,10 @@ async def _(param: HandleRequest) -> Result: group.group_flag = 1 await group.save(update_fields=["group_flag"]) else: - await GroupConsole.update_or_create( + await GroupConsole.get_or_create_root_group( group_id=req.group_id, defaults={"group_flag": 1}, + update_defaults=True, ) try: await FgRequest.approve(bot, param.id) @@ -240,8 +241,7 @@ async def _(param: HandleRequest) -> Result: async def _(param: LeaveGroup) -> Result: try: bot = nonebot.get_bot(param.bot_id) - platform = PlatformUtils.get_platform(bot) - if platform != "qq": + if PlatformUtils.get_platform_scope(bot) != "qq_client": return Result.warning_("该平台不支持退群操作...") group_list, _ = await PlatformUtils.get_group_list(bot) if param.group_id not in [g.group_id for g in group_list]: @@ -265,8 +265,7 @@ async def _(param: LeaveGroup) -> Result: async def _(param: DeleteFriend) -> Result: try: bot = nonebot.get_bot(param.bot_id) - platform = PlatformUtils.get_platform(bot) - if platform != "qq": + if PlatformUtils.get_platform_scope(bot) != "qq_client": return Result.warning_("该平台不支持删除好友操作...") friend_list, _ = await PlatformUtils.get_friend_list(bot) if param.user_id not in [f.user_id for f in friend_list]: diff --git a/zhenxun/builtin_plugins/withdraw.py b/zhenxun/builtin_plugins/withdraw.py index eb0cd0dc..2b4b7bdd 100644 --- a/zhenxun/builtin_plugins/withdraw.py +++ b/zhenxun/builtin_plugins/withdraw.py @@ -40,7 +40,7 @@ def reply_check() -> Rule: if event.get_type() == "message": return ( bool(await reply_fetch(event, bot)) - and PlatformUtils.get_platform(session) == "qq" + and PlatformUtils.get_platform_scope(session) == "qq_client" ) return False diff --git a/zhenxun/models/_bot_message_buffer.py b/zhenxun/models/_bot_message_buffer.py deleted file mode 100644 index d2e30129..00000000 --- a/zhenxun/models/_bot_message_buffer.py +++ /dev/null @@ -1,113 +0,0 @@ -from __future__ import annotations - -import asyncio -from collections import deque -import contextlib -import time -from typing import TYPE_CHECKING - -from zhenxun.services.log import logger - -if TYPE_CHECKING: - from .bot_message_store import BotMessageStore - -LOG_COMMAND = "BotMessageStore" - -_BUFFER_MAX_RETAIN = 20_000 -_FLUSH_TRIGGER_SIZE = 64 -_FLUSH_BATCH_SIZE = 500 -_FLUSH_INTERVAL_SECONDS = 5.0 -_DROP_LOG_INTERVAL_SECONDS = 10.0 - -_buffer: deque[BotMessageStore] = deque() -_buffer_lock = asyncio.Lock() -_flush_lock = asyncio.Lock() -_flush_task: asyncio.Task[None] | None = None -_dropped = 0 -_last_drop_log_at = 0.0 - - -def _ensure_flush_task() -> None: - global _flush_task - if _flush_task is not None and not _flush_task.done(): - return - try: - loop = asyncio.get_running_loop() - except RuntimeError: - return - _flush_task = loop.create_task(_flush_loop()) - - -def _record_drop() -> None: - global _dropped, _last_drop_log_at - _dropped += 1 - now = time.monotonic() - if now - _last_drop_log_at < _DROP_LOG_INTERVAL_SECONDS: - return - _last_drop_log_at = now - logger.warning( - f"bot_message_store buffer full, dropped {_dropped} records, " - f"backlog={len(_buffer)}", - LOG_COMMAND, - ) - - -async def _flush_loop() -> None: - while True: - await asyncio.sleep(_FLUSH_INTERVAL_SECONDS) - try: - await flush_bot_message_store_buffer("定时") - except asyncio.CancelledError: - raise - except Exception as exc: - logger.warning("定时批量写入 Bot 发送记录失败", LOG_COMMAND, e=exc) - - -async def append_bot_message_store_record(record: BotMessageStore) -> None: - _ensure_flush_task() - async with _buffer_lock: - if len(_buffer) >= _BUFFER_MAX_RETAIN: - _buffer.popleft() - _record_drop() - _buffer.append(record) - should_flush = len(_buffer) >= _FLUSH_TRIGGER_SIZE and not _flush_lock.locked() - if should_flush: - await flush_bot_message_store_buffer("缓冲区触发") - - -async def flush_bot_message_store_buffer(reason: str) -> int: - from .bot_message_store import BotMessageStore - - async with _flush_lock: - written = 0 - while True: - batch: list[BotMessageStore] = [] - async with _buffer_lock: - while _buffer and len(batch) < _FLUSH_BATCH_SIZE: - batch.append(_buffer.popleft()) - if not batch: - break - try: - await BotMessageStore.bulk_create(batch, batch_size=_FLUSH_BATCH_SIZE) - except Exception as exc: - async with _buffer_lock: - retain_count = max(_BUFFER_MAX_RETAIN - len(_buffer), 0) - for record in reversed(batch[-retain_count:]): - _buffer.appendleft(record) - logger.error(f"{reason}批量写入 Bot 发送记录失败", LOG_COMMAND, e=exc) - return written - written += len(batch) - if written: - logger.debug(f"{reason}批量写入 Bot 发送记录 {written} 条", LOG_COMMAND) - return written - - -async def stop_bot_message_store_buffer() -> int: - global _flush_task - task = _flush_task - _flush_task = None - if task is not None: - task.cancel() - with contextlib.suppress(BaseException): - await task - return await flush_bot_message_store_buffer("关闭") diff --git a/zhenxun/models/auth_decision_log.py b/zhenxun/models/auth_decision_log.py deleted file mode 100644 index 33578a9e..00000000 --- a/zhenxun/models/auth_decision_log.py +++ /dev/null @@ -1,49 +0,0 @@ -from typing import ClassVar - -from tortoise import fields - -from zhenxun.services.db_context import Model - - -class AuthDecisionLog(Model): - id = fields.IntField(pk=True, generated=True, auto_increment=True) - """自增id""" - bot_id = fields.CharField(255, null=True, description="Bot ID") - """Bot ID""" - platform = fields.CharField(64, null=True, description="平台") - """平台""" - group_id = fields.CharField(255, null=True, description="群组id") - """群组id""" - user_id = fields.CharField(255, null=True, description="用户id") - """用户id""" - module = fields.CharField(255, null=True, description="插件模块") - """插件模块""" - effect = fields.CharField(32, description="决策结果") - """决策结果 allow/deny/skip/defer/error""" - reason = fields.CharField(255, null=True, description="原因") - """原因""" - shadow_effect = fields.CharField(32, null=True, description="影子决策结果") - """影子决策结果""" - shadow_reason = fields.CharField(255, null=True, description="影子决策原因") - """影子决策原因""" - side_effect_state = fields.TextField(null=True, description="副作用状态") - """副作用状态 JSON 摘要""" - latency_ms = fields.FloatField(default=0, description="耗时毫秒") - """耗时毫秒""" - overloaded = fields.BooleanField(default=False, description="是否过载") - """是否过载""" - create_time = fields.DatetimeField(auto_now_add=True, description="创建时间") - """创建时间""" - - class Meta: # pyright: ignore [reportIncompatibleVariableOverride] - table = "auth_decision_log" - table_description = "权限决策追加审计日志" - indexes: ClassVar = [ - ("create_time",), - ("module", "create_time"), - ("effect", "create_time"), - ] - - @classmethod - async def _run_script(cls): - return [] diff --git a/zhenxun/models/ban_console.py b/zhenxun/models/ban_console.py index 2080486e..12302096 100644 --- a/zhenxun/models/ban_console.py +++ b/zhenxun/models/ban_console.py @@ -3,6 +3,7 @@ from typing import ClassVar from typing_extensions import Self from tortoise import fields +from tortoise.expressions import Q from zhenxun.services.cache import CacheException, CacheRegistry, CacheRoot from zhenxun.services.cache.runtime_cache import BanMemoryCache @@ -89,15 +90,18 @@ class BanConsole(Model): cls._ensure_cache_registered() if not user_id and not group_id: raise UserAndGroupIsNone() - dao = DataAccess(cls) if user_id: + dao = DataAccess(cls) return ( await dao.safe_get_or_none(user_id=user_id, group_id=group_id) if group_id else await dao.safe_get_or_none(user_id=user_id, group_id__isnull=True) ) else: - return await dao.safe_get_or_none(user_id="", group_id=group_id) + return await cls.safe_get_or_none( + Q(user_id__isnull=True) | Q(user_id=""), + group_id=group_id, + ) @classmethod async def check_ban_level( diff --git a/zhenxun/models/bot_console.py b/zhenxun/models/bot_console.py index 176e8e29..1180a799 100644 --- a/zhenxun/models/bot_console.py +++ b/zhenxun/models/bot_console.py @@ -26,7 +26,7 @@ class BotConsole(Model): available_plugins = fields.TextField(default="", description="可用插件") """可用插件""" available_tasks = fields.TextField(default="", description="可用被动技能") - """可用被动技能""" + """可用被动技能管理镜像,不作为运行白名单。""" class Meta: # pyright: ignore [reportIncompatibleVariableOverride] table = "bot_console" @@ -87,7 +87,11 @@ class BotConsole(Model): @classmethod async def get_tasks(cls, bot_id: str | None = None, status: bool | None = True): """ - 获取bot被动技能 + 获取bot被动技能管理镜像。 + + available_tasks 只服务管理命令和展示,不参与运行白名单判断。 + 被动运行真源是 TaskInfo.status/load_status、BotConsole.block_tasks、 + GroupConsole.block_task/superuser_block_task。 参数: bot_id (str | None, optional): bot_id. Defaults to None. diff --git a/zhenxun/models/bot_message_store.py b/zhenxun/models/bot_message_store.py deleted file mode 100644 index 6159e094..00000000 --- a/zhenxun/models/bot_message_store.py +++ /dev/null @@ -1,55 +0,0 @@ -from tortoise import fields - -from zhenxun.services.db_context import Model -from zhenxun.utils.enum import BotSentType - -from ._bot_message_buffer import append_bot_message_store_record - - -class BotMessageStore(Model): - id = fields.IntField(pk=True, generated=True, auto_increment=True) - """自增id""" - bot_id = fields.CharField(255, null=True) - """bot id""" - user_id = fields.CharField(255, null=True) - """目标id""" - group_id = fields.CharField(255, null=True) - """群组id""" - sent_type = fields.CharEnumField(BotSentType) - """类型""" - text = fields.TextField(null=True) - """文本内容""" - plain_text = fields.TextField(null=True) - """纯文本""" - platform = fields.CharField(255, null=True) - """平台""" - create_time = fields.DatetimeField(auto_now_add=True) - """创建时间""" - - class Meta: # pyright: ignore [reportIncompatibleVariableOverride] - table = "bot_message_store" - table_description = "Bot发送消息列表" - - @classmethod - async def append_buffered( - cls, - *, - bot_id: str | None = None, - user_id: str | None = None, - group_id: str | None = None, - sent_type: BotSentType, - text: str | None = None, - plain_text: str | None = None, - platform: str | None = None, - ) -> None: - await append_bot_message_store_record( - cls( - bot_id=bot_id, - user_id=user_id, - group_id=group_id, - sent_type=sent_type, - text=text, - plain_text=plain_text, - platform=platform, - ) - ) diff --git a/zhenxun/models/fg_request.py b/zhenxun/models/fg_request.py index 2342abbb..0b80d241 100644 --- a/zhenxun/models/fg_request.py +++ b/zhenxun/models/fg_request.py @@ -151,8 +151,10 @@ class FgRequest(Model): "添加好友自动发送BOT自我介绍图片", session=req.user_id ) else: - await GroupConsole.update_or_create( - group_id=req.group_id, defaults={"group_flag": 1} + await GroupConsole.get_or_create_root_group( + group_id=req.group_id, + defaults={"group_flag": 1}, + update_defaults=True, ) if req.flag == "0": # 用户手动申请入群,创建群认证后提醒用户拉群 diff --git a/zhenxun/models/group_console.py b/zhenxun/models/group_console.py index c52650fb..19001f13 100644 --- a/zhenxun/models/group_console.py +++ b/zhenxun/models/group_console.py @@ -1,3 +1,4 @@ +import asyncio from typing import TYPE_CHECKING, Any, ClassVar, cast, overload from typing_extensions import Self @@ -103,6 +104,8 @@ class GroupConsole(Model): """缓存键字段""" enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT] """开启锁""" + _root_group_locks: ClassVar[dict[str, asyncio.Lock]] = {} + """普通群记录应用层锁,规避 channel_id=NULL 唯一键语义差异。""" @classmethod async def _get_task_modules(cls, *, default_status: bool) -> list[str]: @@ -250,6 +253,143 @@ class GroupConsole(Model): return group, is_create + @classmethod + def _clean_root_group_defaults(cls, defaults: dict | None) -> dict[str, Any]: + cleaned = {} + for field, value in (defaults or {}).items(): + if field in {"id", "group_id", "channel_id", "channel_id__isnull"}: + continue + if value is None: + continue + cleaned[field] = value + return cleaned + + @classmethod + def _root_group_score(cls, group: Self) -> tuple[int, int]: + score = 0 + score += 8 if group.group_name else 0 + score += 4 if group.max_member_count else 0 + score += 4 if group.member_count else 0 + score += 4 if group.group_flag else 0 + score += 4 if group.is_super else 0 + score += 3 if not group.status else 0 + score += 3 if group.level != 5 else 0 + score += len(convert_module_format(group.block_plugin)) + score += len(convert_module_format(group.superuser_block_plugin)) + score += len(convert_module_format(group.block_task)) + score += len(convert_module_format(group.superuser_block_task)) + return score, int(group.id or 0) + + @classmethod + def _merge_module_field(cls, groups: list[Self], field: str) -> str: + modules: list[str] = [] + seen = set() + for group in groups: + value = getattr(group, field, "") or "" + for module in cast(list[str], convert_module_format(value)): + if module not in seen: + seen.add(module) + modules.append(module) + return cast(str, convert_module_format(modules)) + + @classmethod + async def _deduplicate_root_group_records(cls, groups: list[Self]) -> Self: + if len(groups) == 1: + return groups[0] + + keep = max(groups, key=cls._root_group_score) + newest_first = sorted( + groups, key=lambda group: int(group.id or 0), reverse=True + ) + + merged = { + "group_name": next( + (g.group_name for g in newest_first if g.group_name), "" + ), + "max_member_count": max(g.max_member_count for g in groups), + "member_count": max(g.member_count for g in groups), + "status": all(g.status for g in groups), + "level": min(g.level for g in groups), + "is_super": any(g.is_super for g in groups), + "group_flag": max(g.group_flag for g in groups), + "block_plugin": cls._merge_module_field(groups, "block_plugin"), + "superuser_block_plugin": cls._merge_module_field( + groups, "superuser_block_plugin" + ), + "block_task": cls._merge_module_field(groups, "block_task"), + "superuser_block_task": cls._merge_module_field( + groups, "superuser_block_task" + ), + "platform": next((g.platform for g in newest_first if g.platform), "qq"), + } + + update_fields = [] + for field, value in merged.items(): + if getattr(keep, field) != value: + setattr(keep, field, value) + update_fields.append(field) + if update_fields: + await keep.save(update_fields=update_fields) + + for group in groups: + if group.id != keep.id: + await group.delete() + await GroupMemoryCache.upsert_from_model(keep) + return keep + + @classmethod + async def get_or_create_root_group( + cls, + group_id: str | int, + defaults: dict | None = None, + *, + update_defaults: bool = False, + ) -> tuple[Self, bool]: + """获取或创建普通群记录,并收敛 channel_id=NULL 重复数据。 + + 普通群固定使用 ``channel_id IS NULL``;频道记录必须继续显式传 + ``channel_id`` 走原有 get_or_create/update_or_create。 + """ + gid = str(group_id).strip() + if not gid: + raise ValueError("group_id cannot be empty") + + lock = cls._root_group_locks.setdefault(gid, asyncio.Lock()) + async with lock: + defaults = cls._clean_root_group_defaults(defaults) + records = await cls.filter(group_id=gid, channel_id__isnull=True).all() + if records: + group = await cls._deduplicate_root_group_records(records) + if update_defaults: + update_fields = [] + for field, value in defaults.items(): + if hasattr(group, field) and getattr(group, field) != value: + setattr(group, field, value) + update_fields.append(field) + if update_fields: + await group.save(update_fields=update_fields) + await GroupMemoryCache.upsert_from_model(group) + return group, False + + group = await cls.create(group_id=gid, channel_id=None, **defaults) + return group, True + + @classmethod + async def _get_or_create_group_for_write( + cls, + group_id: str, + channel_id: str | None, + defaults: dict | None = None, + ) -> tuple[Self, bool]: + defaults = cls._clean_root_group_defaults(defaults) + if channel_id: + return await cls.get_or_create( + group_id=group_id, + channel_id=channel_id, + defaults=defaults, + ) + return await cls.get_or_create_root_group(group_id, defaults=defaults) + async def save(self, *args, **kwargs): await super().save(*args, **kwargs) await GroupMemoryCache.upsert_from_model(self) @@ -326,6 +466,7 @@ class GroupConsole(Model): module: str, is_superuser: bool = False, platform: str | None = None, + channel_id: str | None = None, ): """禁用群组插件 @@ -335,8 +476,10 @@ class GroupConsole(Model): is_superuser: 是否为超级用户 platform: 平台 """ - group, _ = await cls.get_or_create( - group_id=group_id, defaults={"platform": platform} + group, _ = await cls._get_or_create_group_for_write( + group_id=group_id, + channel_id=channel_id, + defaults={"platform": platform}, ) update_fields = [] if is_superuser: @@ -365,6 +508,7 @@ class GroupConsole(Model): module: str, is_superuser: bool = False, platform: str | None = None, + channel_id: str | None = None, ): """禁用群组插件 @@ -374,8 +518,10 @@ class GroupConsole(Model): is_superuser: 是否为超级用户 platform: 平台 """ - group, _ = await cls.get_or_create( - group_id=group_id, defaults={"platform": platform} + group, _ = await cls._get_or_create_group_for_write( + group_id=group_id, + channel_id=channel_id, + defaults={"platform": platform}, ) update_fields = [] if is_superuser: @@ -447,6 +593,7 @@ class GroupConsole(Model): task: str, is_superuser: bool = False, platform: str | None = None, + channel_id: str | None = None, ): """禁用群组插件 @@ -456,13 +603,15 @@ class GroupConsole(Model): is_superuser: 是否为超级用户 platform: 平台 """ - group, _ = await cls.get_or_create( - group_id=group_id, defaults={"platform": platform} + group, _ = await cls._get_or_create_group_for_write( + group_id=group_id, + channel_id=channel_id, + defaults={"platform": platform}, ) update_fields = [] if is_superuser: superuser_block_task = convert_module_format(group.superuser_block_task) - if task not in group.superuser_block_task: + if task not in superuser_block_task: superuser_block_task.append(task) group.superuser_block_task = convert_module_format(superuser_block_task) update_fields.append("superuser_block_task") @@ -484,6 +633,7 @@ class GroupConsole(Model): task: str, is_superuser: bool = False, platform: str | None = None, + channel_id: str | None = None, ): """禁用群组插件 @@ -493,8 +643,10 @@ class GroupConsole(Model): is_superuser: 是否为超级用户 platform: 平台 """ - group, _ = await cls.get_or_create( - group_id=group_id, defaults={"platform": platform} + group, _ = await cls._get_or_create_group_for_write( + group_id=group_id, + channel_id=channel_id, + defaults={"platform": platform}, ) update_fields = [] if is_superuser: diff --git a/zhenxun/models/runtime_backpressure_log.py b/zhenxun/models/runtime_backpressure_log.py deleted file mode 100644 index 3b8f050f..00000000 --- a/zhenxun/models/runtime_backpressure_log.py +++ /dev/null @@ -1,39 +0,0 @@ -from typing import ClassVar - -from tortoise import fields - -from zhenxun.services.db_context import Model - - -class RuntimeBackpressureLog(Model): - id = fields.IntField(pk=True, generated=True, auto_increment=True) - """自增id""" - scope_key = fields.CharField(255, null=True, description="作用域") - """作用域""" - reason = fields.CharField(255, null=True, description="原因") - """原因""" - lane = fields.CharField(64, null=True, description="调度通道") - """调度通道""" - action = fields.CharField(64, description="处理动作") - """处理动作 execute/skip/defer/signal""" - queue_size = fields.IntField(default=0, description="队列长度") - """队列长度""" - active_count = fields.IntField(default=0, description="活跃数量") - """活跃数量""" - duration_ms = fields.FloatField(default=0, description="持续耗时毫秒") - """持续耗时毫秒""" - create_time = fields.DatetimeField(auto_now_add=True, description="创建时间") - """创建时间""" - - class Meta: # pyright: ignore [reportIncompatibleVariableOverride] - table = "runtime_backpressure_log" - table_description = "运行时背压追加审计日志" - indexes: ClassVar = [ - ("create_time",), - ("scope_key", "create_time"), - ("lane", "create_time"), - ] - - @classmethod - async def _run_script(cls): - return [] diff --git a/zhenxun/models/user_console.py b/zhenxun/models/user_console.py index bcc8a6a1..a25de811 100644 --- a/zhenxun/models/user_console.py +++ b/zhenxun/models/user_console.py @@ -219,6 +219,7 @@ class UserConsole(Model): ) if not updated: raise InsufficientGold() + await cls.invalidate_user_cache(user_id) return GoldReservation( user_id=user_id, gold=gold, diff --git a/zhenxun/services/auth_observability.py b/zhenxun/services/auth_observability.py deleted file mode 100644 index e4b76fff..00000000 --- a/zhenxun/services/auth_observability.py +++ /dev/null @@ -1,547 +0,0 @@ -from __future__ import annotations - -import asyncio -from collections import deque -import contextlib -from dataclasses import dataclass -from datetime import datetime, timedelta -import json -import random -import time -from typing import Any, TypeVar - -from tortoise import Tortoise - -from zhenxun.builtin_plugins.hooks.auth_runtime_config import ( - AUTH_OBSERVABILITY_RUNTIME_CONFIG, -) -from zhenxun.models.auth_decision_log import AuthDecisionLog -from zhenxun.models.runtime_backpressure_log import RuntimeBackpressureLog -from zhenxun.services.log import logger -from zhenxun.utils.manager.priority_manager import PriorityLifecycle - -LOG_COMMAND = "AuthObservability" - -_BUFFER_MAX_RETAIN = AUTH_OBSERVABILITY_RUNTIME_CONFIG.buffer_max_retain -_FLUSH_TRIGGER_SIZE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.flush_trigger_size -_FLUSH_BATCH_SIZE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.flush_batch_size -_FLUSH_INTERVAL_SECONDS = AUTH_OBSERVABILITY_RUNTIME_CONFIG.flush_interval_seconds -_DROP_LOG_INTERVAL_SECONDS = AUTH_OBSERVABILITY_RUNTIME_CONFIG.drop_log_interval_seconds -_ALLOW_SAMPLE_RATE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.allow_sample_rate -_OVERLOADED_ALLOW_SAMPLE_RATE = ( - AUTH_OBSERVABILITY_RUNTIME_CONFIG.overloaded_allow_sample_rate -) -_NON_ALLOW_SAMPLE_RATE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.non_allow_sample_rate -_BACKPRESSURE_SAMPLE_RATE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.backpressure_sample_rate -_BACKPRESSURE_SEVERE_ACTIVE_THRESHOLD = ( - AUTH_OBSERVABILITY_RUNTIME_CONFIG.backpressure_severe_active_threshold -) - - -@dataclass(slots=True) -class AuthDecisionLogRecord: - bot_id: str | None - platform: str | None - group_id: str | None - user_id: str | None - module: str | None - effect: str - reason: str | None = None - shadow_effect: str | None = None - shadow_reason: str | None = None - side_effect_state: dict[str, Any] | None = None - latency_ms: float = 0.0 - overloaded: bool = False - - def to_model(self) -> AuthDecisionLog: - return AuthDecisionLog( - bot_id=self.bot_id, - platform=self.platform, - group_id=self.group_id, - user_id=self.user_id, - module=self.module, - effect=self.effect, - reason=self.reason, - shadow_effect=self.shadow_effect, - shadow_reason=self.shadow_reason, - side_effect_state=json.dumps( - self.side_effect_state, - ensure_ascii=False, - separators=(",", ":"), - )[:4000] - if self.side_effect_state - else None, - latency_ms=self.latency_ms, - overloaded=self.overloaded, - ) - - -@dataclass(slots=True) -class RuntimeBackpressureLogRecord: - scope_key: str | None - reason: str | None - lane: str | None - action: str - queue_size: int = 0 - active_count: int = 0 - duration_ms: float = 0.0 - - def to_model(self) -> RuntimeBackpressureLog: - return RuntimeBackpressureLog( - scope_key=self.scope_key, - reason=self.reason, - lane=self.lane, - action=self.action, - queue_size=self.queue_size, - active_count=self.active_count, - duration_ms=self.duration_ms, - ) - - -_auth_decision_buffer: deque[AuthDecisionLogRecord] = deque() -_backpressure_buffer: deque[RuntimeBackpressureLogRecord] = deque() -_buffer_lock = asyncio.Lock() -_flush_lock = asyncio.Lock() -_flush_task: asyncio.Task[None] | None = None -_dropped = 0 -_last_drop_log_at = 0.0 -_last_schema_repair_at = 0.0 -_SCHEMA_REPAIR_INTERVAL_SECONDS = 300.0 - -T = TypeVar("T") - - -def _ensure_flush_task() -> None: - global _flush_task - if _flush_task is not None and not _flush_task.done(): - return - _flush_task = asyncio.create_task(_flush_loop()) - - -def _record_drop() -> None: - global _dropped, _last_drop_log_at - _dropped += 1 - now = time.monotonic() - if now - _last_drop_log_at < _DROP_LOG_INTERVAL_SECONDS: - return - _last_drop_log_at = now - logger.warning( - "auth observability buffer full, dropped " - f"{_dropped} records, auth_backlog={len(_auth_decision_buffer)}, " - f"backpressure_backlog={len(_backpressure_buffer)}", - LOG_COMMAND, - ) - - -def _sample(rate: float) -> bool: - if rate >= 1: - return True - if rate <= 0: - return False - return random.random() < rate - - -def _auth_decision_sample_rate(effect: str, overloaded: bool) -> float: - if effect != "allow": - return _NON_ALLOW_SAMPLE_RATE - if overloaded: - return _OVERLOADED_ALLOW_SAMPLE_RATE - return _ALLOW_SAMPLE_RATE - - -def _backpressure_sample_rate(record: RuntimeBackpressureLogRecord) -> float: - if record.reason and record.reason.startswith("hooks_"): - return 1.0 - if record.active_count >= _BACKPRESSURE_SEVERE_ACTIVE_THRESHOLD: - return 1.0 - if record.action in {"skip", "defer"}: - return _BACKPRESSURE_SAMPLE_RATE - return min(_BACKPRESSURE_SAMPLE_RATE, 0.02) - - -async def _append_auth_decision_record(record: AuthDecisionLogRecord) -> None: - _ensure_flush_task() - async with _buffer_lock: - total = len(_auth_decision_buffer) + len(_backpressure_buffer) - if total >= _BUFFER_MAX_RETAIN: - if len(_auth_decision_buffer) >= len(_backpressure_buffer): - with contextlib.suppress(IndexError): - _auth_decision_buffer.popleft() - else: - with contextlib.suppress(IndexError): - _backpressure_buffer.popleft() - _record_drop() - _auth_decision_buffer.append(record) - should_flush = ( - len(_auth_decision_buffer) + len(_backpressure_buffer) - >= _FLUSH_TRIGGER_SIZE - and not _flush_lock.locked() - ) - if should_flush: - # Fire-and-forget keeps auth hot path independent of database stalls. - asyncio.create_task(flush_auth_observability_buffer("缓冲区触发")) # noqa: RUF006 - - -async def _append_backpressure_record(record: RuntimeBackpressureLogRecord) -> None: - _ensure_flush_task() - async with _buffer_lock: - total = len(_auth_decision_buffer) + len(_backpressure_buffer) - if total >= _BUFFER_MAX_RETAIN: - if len(_auth_decision_buffer) >= len(_backpressure_buffer): - with contextlib.suppress(IndexError): - _auth_decision_buffer.popleft() - else: - with contextlib.suppress(IndexError): - _backpressure_buffer.popleft() - _record_drop() - _backpressure_buffer.append(record) - should_flush = ( - len(_auth_decision_buffer) + len(_backpressure_buffer) - >= _FLUSH_TRIGGER_SIZE - and not _flush_lock.locked() - ) - if should_flush: - # Fire-and-forget keeps auth hot path independent of database stalls. - asyncio.create_task(flush_auth_observability_buffer("缓冲区触发")) # noqa: RUF006 - - -async def append_auth_decision_log( - *, - bot_id: str | None, - platform: str | None, - group_id: str | None, - user_id: str | None, - module: str | None, - effect: str, - reason: str | None = None, - shadow_effect: str | None = None, - shadow_reason: str | None = None, - side_effect_state: dict[str, Any] | None = None, - latency_ms: float = 0.0, - overloaded: bool = False, -) -> None: - if shadow_effect is None and not _sample( - _auth_decision_sample_rate(effect, overloaded) - ): - return - record = AuthDecisionLogRecord( - bot_id=bot_id, - platform=platform, - group_id=group_id, - user_id=user_id, - module=module, - effect=effect, - reason=(reason or "")[:255] or None, - shadow_effect=(shadow_effect or "")[:32] or None, - shadow_reason=(shadow_reason or "")[:255] or None, - side_effect_state=side_effect_state, - latency_ms=latency_ms, - overloaded=overloaded, - ) - await _append_auth_decision_record(record) - - -async def append_runtime_backpressure_log( - *, - scope_key: str | None, - reason: str | None, - lane: str | None, - action: str, - queue_size: int = 0, - active_count: int = 0, - duration_ms: float = 0.0, -) -> None: - record = RuntimeBackpressureLogRecord( - scope_key=(scope_key or "")[:255] or None, - reason=(reason or "")[:255] or None, - lane=(lane or "")[:64] or None, - action=action, - queue_size=queue_size, - active_count=active_count, - duration_ms=duration_ms, - ) - if not _sample(_backpressure_sample_rate(record)): - return - await _append_backpressure_record(record) - - -async def _flush_loop() -> None: - while True: - await asyncio.sleep(_FLUSH_INTERVAL_SECONDS) - try: - await flush_auth_observability_buffer("定时") - except asyncio.CancelledError: - raise - except Exception as exc: - logger.warning("定时批量写入权限观测日志失败", LOG_COMMAND, e=exc) - - -async def _drain_batch(buffer: deque[T]) -> list[T]: - batch: list[T] = [] - async with _buffer_lock: - while buffer and len(batch) < _FLUSH_BATCH_SIZE: - batch.append(buffer.popleft()) - return batch - - -async def _restore_batch(buffer: deque[T], batch: list[T]) -> None: - async with _buffer_lock: - retain_count = max(_BUFFER_MAX_RETAIN - len(buffer), 0) - for record in reversed(batch[-retain_count:]): - buffer.appendleft(record) - - -def _is_schema_mismatch_error(exc: Exception) -> bool: - message = str(exc).lower() - return any( - marker in message - for marker in ( - "no column named", - "unknown column", - "column does not exist", - "no such column", - ) - ) - - -async def _try_repair_auth_schema_once() -> bool: - global _last_schema_repair_at - now = time.monotonic() - if now - _last_schema_repair_at < _SCHEMA_REPAIR_INTERVAL_SECONDS: - return False - _last_schema_repair_at = now - try: - from zhenxun.services.db_context.schema_guard import repair_table_schema - - await repair_table_schema("auth_decision_log") - await repair_table_schema("runtime_backpressure_log") - return True - except Exception as exc: - logger.warning("权限观测日志表结构自修复失败", LOG_COMMAND, e=exc) - return False - - -async def flush_auth_observability_buffer(reason: str) -> int: - async with _flush_lock: - written = 0 - while True: - auth_batch = await _drain_batch(_auth_decision_buffer) - backpressure_batch = await _drain_batch(_backpressure_buffer) - if not auth_batch and not backpressure_batch: - break - try: - if auth_batch: - await AuthDecisionLog.bulk_create( - [record.to_model() for record in auth_batch], - _FLUSH_BATCH_SIZE, - ) - written += len(auth_batch) - if backpressure_batch: - await RuntimeBackpressureLog.bulk_create( - [record.to_model() for record in backpressure_batch], - _FLUSH_BATCH_SIZE, - ) - written += len(backpressure_batch) - except Exception as exc: - if _is_schema_mismatch_error(exc): - if await _try_repair_auth_schema_once(): - try: - if auth_batch: - await AuthDecisionLog.bulk_create( - [record.to_model() for record in auth_batch], - _FLUSH_BATCH_SIZE, - ) - written += len(auth_batch) - if backpressure_batch: - await RuntimeBackpressureLog.bulk_create( - [ - record.to_model() - for record in backpressure_batch - ], - _FLUSH_BATCH_SIZE, - ) - written += len(backpressure_batch) - continue - except Exception as retry_exc: - exc = retry_exc - dropped = len(auth_batch) + len(backpressure_batch) - logger.warning( - f"{reason}批量写入权限观测日志遇到表结构不匹配," - f"已丢弃低优先级观测日志 {dropped} 条,等待下次启动修复", - LOG_COMMAND, - e=exc, - ) - return written - await _restore_batch(_auth_decision_buffer, auth_batch) - await _restore_batch(_backpressure_buffer, backpressure_batch) - logger.error(f"{reason}批量写入权限观测日志失败", LOG_COMMAND, e=exc) - return written - if written: - logger.debug(f"{reason}批量写入权限观测日志 {written} 条", LOG_COMMAND) - return written - - -async def stop_auth_observability_buffer() -> int: - global _flush_task - task = _flush_task - _flush_task = None - if task is not None: - task.cancel() - with contextlib.suppress(BaseException): - await task - return await flush_auth_observability_buffer("关闭") - - -def _percentile(values: list[float], ratio: float) -> float: - if not values: - return 0.0 - ordered = sorted(values) - index = min(max(round((len(ordered) - 1) * ratio), 0), len(ordered) - 1) - return round(ordered[index], 3) - - -def _bucket_counts(rows: list[dict[str, Any]], field: str) -> dict[str, int]: - counts: dict[str, int] = {} - for row in rows: - key = str(row.get(field) or "") - counts[key] = counts.get(key, 0) + 1 - return counts - - -def _lane_budget_advice(rows: list[dict[str, Any]]) -> dict[str, dict[str, Any]]: - buckets: dict[str, list[dict[str, Any]]] = {} - for row in rows: - lane = str(row.get("lane") or "") - buckets.setdefault(lane, []).append(row) - advice: dict[str, dict[str, Any]] = {} - for lane, items in buckets.items(): - if lane == "": - continue - durations = [float(item.get("duration_ms") or 0.0) for item in items] - slow_waits = sum(1 for value in durations if value >= 200.0) - active_max = max( - (int(item.get("active_count") or 0) for item in items), default=0 - ) - total = len(items) - if not total: - continue - pressure_ratio = slow_waits / total - if pressure_ratio >= 0.2 or active_max >= 5: - action = "increase_or_split" - elif pressure_ratio == 0 and active_max <= 1 and total >= 20: - action = "can_reduce" - else: - action = "keep" - advice[lane] = { - "samples": total, - "slow_waits": slow_waits, - "pressure_ratio": round(pressure_ratio, 3), - "active_max": active_max, - "p95_duration_ms": _percentile(durations, 0.95), - "action": action, - } - return advice - - -def _query_placeholder() -> str: - try: - connection = Tortoise.get_connection("default") - if ( - getattr(connection, "capabilities", None) - and getattr( - connection.capabilities, - "dialect", - "", - ) - == "postgres" - ): - return "$1" - except Exception: - return "?" - return "?" - - -async def build_auth_observability_report(*, hours: float = 24.0) -> dict[str, Any]: - since = datetime.now() - timedelta(hours=hours) - db = Tortoise.get_connection("default") - placeholder = _query_placeholder() - auth_rows = await db.execute_query_dict( - "SELECT module, effect, reason, shadow_effect, shadow_reason, latency_ms, " - f"overloaded FROM auth_decision_log WHERE create_time >= {placeholder} " - "ORDER BY create_time DESC LIMIT 100000", - [since], - ) - backpressure_rows = await db.execute_query_dict( - "SELECT scope_key, lane, reason, action, queue_size, active_count, duration_ms " - f"FROM runtime_backpressure_log WHERE create_time >= {placeholder} " - "ORDER BY create_time DESC LIMIT 100000", - [since], - ) - - module_buckets: dict[str, list[dict[str, Any]]] = {} - for row in auth_rows: - module_buckets.setdefault(str(row.get("module") or ""), []).append(row) - module_stats: list[dict[str, Any]] = [] - for module, items in module_buckets.items(): - latencies = [float(item.get("latency_ms") or 0.0) for item in items] - module_stats.append( - { - "module": module, - "total": len(items), - "effects": _bucket_counts(items, "effect"), - "shadow_effects": _bucket_counts(items, "shadow_effect"), - "avg_latency_ms": round(sum(latencies) / len(latencies), 3) - if latencies - else 0.0, - "p95_latency_ms": _percentile(latencies, 0.95), - "overloaded": sum(1 for item in items if bool(item.get("overloaded"))), - } - ) - - backpressure_buckets: dict[str, list[dict[str, Any]]] = {} - for row in backpressure_rows: - key = f"{row.get('lane') or ''}:{row.get('reason') or ''}" - backpressure_buckets.setdefault(key, []).append(row) - backpressure_stats: list[dict[str, Any]] = [] - for key, items in backpressure_buckets.items(): - durations = [float(item.get("duration_ms") or 0.0) for item in items] - backpressure_stats.append( - { - "key": key, - "total": len(items), - "actions": _bucket_counts(items, "action"), - "avg_duration_ms": round(sum(durations) / len(durations), 3) - if durations - else 0.0, - "p95_duration_ms": _percentile(durations, 0.95), - } - ) - - return { - "created_at": datetime.now().isoformat(timespec="seconds"), - "window_hours": hours, - "auth_decisions": { - "total": len(auth_rows), - "effects": _bucket_counts(auth_rows, "effect"), - "shadow_effects": _bucket_counts(auth_rows, "shadow_effect"), - "top_modules_by_p95": sorted( - module_stats, - key=lambda item: (item["p95_latency_ms"], item["total"]), - reverse=True, - )[:30], - }, - "backpressure": { - "total": len(backpressure_rows), - "lane_budget_advice": _lane_budget_advice(backpressure_rows), - "top_reasons": sorted( - backpressure_stats, - key=lambda item: (item["total"], item["p95_duration_ms"]), - reverse=True, - )[:30], - }, - } - - -@PriorityLifecycle.on_shutdown(priority=90) -async def _flush_auth_observability_buffer_on_shutdown() -> None: - await stop_auth_observability_buffer() diff --git a/zhenxun/services/cache/runtime_cache.py b/zhenxun/services/cache/runtime_cache.py index 4eb0cfd7..d43cc1b7 100644 --- a/zhenxun/services/cache/runtime_cache.py +++ b/zhenxun/services/cache/runtime_cache.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +from contextvars import ContextVar from dataclasses import dataclass, field import json import os @@ -35,6 +36,8 @@ def _coerce_int(value, default: int) -> int: # RuntimeCache 以模型 save/delete 主动失效为主,周期 refresh 只做兜底。 # 这些默认值避免低压力运行时频繁全量扫表。 +# 权限检查热路径以 RuntimeCache/AuthSnapshot 为唯一数据入口;普通业务的 +# DataAccess/CacheRoot 缓存不能替代这里的运行态快照。 PLUGININFO_MEM_REFRESH_INTERVAL = 3600 # 60分钟 BAN_MEM_REFRESH_INTERVAL = 300 BAN_MEM_CLEAN_INTERVAL = 60 @@ -52,10 +55,16 @@ LIMIT_MEM_REFRESH_INTERVAL = 900 # 15分钟 LIMIT_MEM_NEGATIVE_TTL = 30 RUNTIME_CACHE_SYNC_ENABLED = True RUNTIME_CACHE_SYNC_CHANNEL = "ZHENXUN_RUNTIME_CACHE_SYNC" +RUNTIME_CACHE_LOAD_RETRY_SECONDS = 1.0 +RUNTIME_CACHE_STARTUP_REFRESH_SKIP_SECONDS = 5.0 INSTANCE_ID = uuid.uuid4().hex _CACHE_READY_EVENT = asyncio.Event() +_APPLYING_REMOTE_CACHE_EVENT: ContextVar[bool] = ContextVar( + "APPLYING_REMOTE_RUNTIME_CACHE_EVENT", + default=False, +) def _env_get(name: str, default: str | None = None) -> str | None: @@ -180,6 +189,65 @@ class PluginInfoSnapshot: plugin._saved_in_db = True return plugin + def to_payload(self) -> dict[str, Any]: + return { + "id": self.id, + "module": self.module, + "module_path": self.module_path, + "name": self.name, + "status": self.status, + "block_type": self.block_type.value if self.block_type else None, + "load_status": self.load_status, + "author": self.author, + "version": self.version, + "level": self.level, + "default_status": self.default_status, + "limit_superuser": self.limit_superuser, + "menu_type": self.menu_type, + "plugin_type": self.plugin_type.value if self.plugin_type else None, + "cost_gold": self.cost_gold, + "admin_level": self.admin_level, + "ignore_prompt": self.ignore_prompt, + "is_delete": self.is_delete, + "parent": self.parent, + "is_show": self.is_show, + "ignore_statistics": self.ignore_statistics, + "impression": self.impression, + } + + @classmethod + def from_payload(cls, payload: dict[str, Any]) -> "PluginInfoSnapshot": + block_type = payload.get("block_type") + if block_type is not None and not isinstance(block_type, BlockType): + block_type = BlockType(block_type) + plugin_type = payload.get("plugin_type") + if plugin_type is not None and not isinstance(plugin_type, PluginType): + plugin_type = PluginType(plugin_type) + return cls( + id=int(payload.get("id", 0) or 0), + module=str(payload.get("module", "") or ""), + module_path=str(payload.get("module_path", "") or ""), + name=str(payload.get("name", "") or ""), + status=bool(payload.get("status", True)), + block_type=block_type, + load_status=bool(payload.get("load_status", True)), + author=payload.get("author"), + version=payload.get("version"), + level=int(payload.get("level", 0) or 0), + default_status=bool(payload.get("default_status", True)), + limit_superuser=bool(payload.get("limit_superuser", False)), + menu_type=str(payload.get("menu_type", "") or ""), + plugin_type=plugin_type, + cost_gold=int(payload.get("cost_gold", 0) or 0), + admin_level=payload.get("admin_level"), + ignore_prompt=bool(payload.get("ignore_prompt", False)), + is_delete=bool(payload.get("is_delete", False)), + parent=payload.get("parent"), + is_show=bool(payload.get("is_show", True)), + ignore_statistics=bool(payload.get("ignore_statistics", False)), + impression=float(payload.get("impression", 0) or 0), + ) + @dataclass(frozen=True) class BanEntry: @@ -648,18 +716,87 @@ class RuntimeCacheSync: 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 == "task": - await TaskInfoMemoryCache.apply_sync_event(action, data) - elif cache_type == "plugin_limit": - await PluginLimitMemoryCache.apply_sync_event(action, data) + token = _APPLYING_REMOTE_CACHE_EVENT.set(True) + try: + 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 == "task": + await TaskInfoMemoryCache.apply_sync_event(action, data) + elif cache_type == "plugin_limit": + await PluginLimitMemoryCache.apply_sync_event(action, data) + elif cache_type == "plugin": + await PluginInfoMemoryCache.apply_sync_event(action, data) + finally: + _APPLYING_REMOTE_CACHE_EVENT.reset(token) + + +class RuntimeCacheMutation: + """Small helpers for runtime cache mutation bookkeeping. + + Cache classes still own their storage layout. This helper centralizes the + shared mutation side effects: health markers, negative-cache cleanup and + cross-process publish. + """ + + _load_locks: ClassVar[dict[str, asyncio.Lock]] = {} + _retry_after: ClassVar[dict[str, float]] = {} + + @classmethod + async def ensure_loaded(cls, cache_cls: type, label: str) -> None: + if getattr(cache_cls, "_loaded", False): + return + now = time.monotonic() + if cls._retry_after.get(label, 0.0) > now: + return + lock = cls._load_locks.setdefault(label, asyncio.Lock()) + async with lock: + if getattr(cache_cls, "_loaded", False): + return + now = time.monotonic() + if cls._retry_after.get(label, 0.0) > now: + return + try: + await cache_cls.refresh() + except Exception as exc: + cls.mark_error(cache_cls, exc) + cls._retry_after[label] = ( + time.monotonic() + RUNTIME_CACHE_LOAD_RETRY_SECONDS + ) + raise + + @staticmethod + def mark_refreshed(cache_cls: type) -> None: + setattr(cache_cls, "_loaded", True) + setattr(cache_cls, "_last_refresh", time.time()) + setattr(cache_cls, "_last_error", None) + + @staticmethod + def mark_error(cache_cls: type, exc: Exception) -> None: + setattr(cache_cls, "_last_error", f"{type(exc).__name__}: {exc}") + + @staticmethod + def clear_negative_key(cache_cls: type, key: object) -> None: + negative = getattr(cache_cls, "_negative", None) + if isinstance(negative, dict): + negative.pop(key, None) + + @staticmethod + def clear_negative_all(cache_cls: type) -> None: + negative = getattr(cache_cls, "_negative", None) + if isinstance(negative, dict): + negative.clear() + + @staticmethod + def publish(cache_type: str, action: str, data: dict[str, Any]) -> None: + if _APPLYING_REMOTE_CACHE_EVENT.get(): + return + RuntimeCacheSync.publish_event(cache_type, action, data) class PluginInfoMemoryCache: @@ -669,6 +806,7 @@ class PluginInfoMemoryCache: _loaded: ClassVar[bool] = False _refresh_task: ClassVar[asyncio.Task | None] = None _last_refresh: ClassVar[float] = 0.0 + _last_error: ClassVar[str | None] = None @classmethod def _to_model(cls, snapshot: PluginInfoSnapshot | None) -> "PluginInfo | None": @@ -687,6 +825,10 @@ class PluginInfoMemoryCache: cls._by_module.pop(old.module, None) cls._by_module_path[snapshot.module_path] = snapshot + @staticmethod + def _module_snapshot_rank(snapshot: PluginInfoSnapshot) -> tuple[int, int]: + return (1 if snapshot.load_status else 0, snapshot.id) + @classmethod async def refresh(cls) -> None: from zhenxun.models.plugin_info import PluginInfo @@ -698,13 +840,16 @@ class PluginInfoMemoryCache: for plugin in plugins: snapshot = PluginInfoSnapshot.from_model(plugin) if snapshot.module: - by_module[snapshot.module] = snapshot + current = by_module.get(snapshot.module) + if current is None or cls._module_snapshot_rank( + snapshot + ) >= cls._module_snapshot_rank(current): + by_module[snapshot.module] = snapshot if snapshot.module_path: by_module_path[snapshot.module_path] = snapshot cls._by_module = by_module cls._by_module_path = by_module_path - cls._loaded = True - cls._last_refresh = time.time() + RuntimeCacheMutation.mark_refreshed(cls) logger.debug( f"plugin cache refreshed: {len(by_module)} entries", LOG_COMMAND ) @@ -713,7 +858,7 @@ class PluginInfoMemoryCache: async def ensure_loaded(cls) -> None: if cls._loaded: return - await cls.refresh() + await RuntimeCacheMutation.ensure_loaded(cls, "plugin") @classmethod def is_loaded(cls) -> bool: @@ -749,8 +894,6 @@ class PluginInfoMemoryCache: return snapshot = PluginInfoSnapshot.from_model(plugin) cls._store_snapshot(snapshot) - cls._loaded = True - cls._last_refresh = time.time() @classmethod def remove_by_module(cls, module: str) -> None: @@ -765,8 +908,16 @@ class PluginInfoMemoryCache: async with cls._lock: snapshot = PluginInfoSnapshot.from_model(plugin) cls._store_snapshot(snapshot) - cls._loaded = True - cls._last_refresh = time.time() + RuntimeCacheMutation.publish("plugin", "upsert", snapshot.to_payload()) + + @classmethod + async def upsert_from_payload(cls, payload: dict[str, Any]) -> None: + try: + snapshot = PluginInfoSnapshot.from_payload(payload) + except Exception: + return + async with cls._lock: + cls._store_snapshot(snapshot) @classmethod async def remove( @@ -783,6 +934,20 @@ class PluginInfoMemoryCache: snapshot = cls._by_module_path.pop(module_path, None) if snapshot and snapshot.module: cls._by_module.pop(snapshot.module, None) + RuntimeCacheMutation.publish( + "plugin", + "delete", + {"module": module, "module_path": module_path}, + ) + + @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("module"), data.get("module_path")) + elif action == "refresh": + await cls.refresh() @classmethod async def _refresh_loop(cls, interval: int) -> None: @@ -815,6 +980,8 @@ class BotMemoryCache: _negative: ClassVar[dict[str, float]] = {} _loaded: ClassVar[bool] = False _refresh_task: ClassVar[asyncio.Task | None] = None + _last_refresh: ClassVar[float] = 0.0 + _last_error: ClassVar[str | None] = None @classmethod def _normalize(cls, bot_id: str | None) -> str | None: @@ -833,7 +1000,7 @@ class BotMemoryCache: if not expire_at: return False if expire_at <= time.time(): - cls._negative.pop(bot_id, None) + RuntimeCacheMutation.clear_negative_key(cls, bot_id) return False return True @@ -851,15 +1018,15 @@ class BotMemoryCache: 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 + RuntimeCacheMutation.clear_negative_all(cls) + RuntimeCacheMutation.mark_refreshed(cls) 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() + await RuntimeCacheMutation.ensure_loaded(cls, "bot") @classmethod def is_loaded(cls) -> bool: @@ -918,15 +1085,16 @@ class BotMemoryCache: available_tasks=entry.available_tasks, ) cls._by_id[bot_id] = updated - RuntimeCacheSync.publish_event("bot", "upsert", updated.to_payload()) + RuntimeCacheMutation.clear_negative_key(cls, bot_id) + RuntimeCacheMutation.publish("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()) + RuntimeCacheMutation.clear_negative_key(cls, entry.bot_id) + RuntimeCacheMutation.publish("bot", "upsert", entry.to_payload()) @classmethod async def upsert_from_payload(cls, payload: dict[str, Any]) -> None: @@ -935,7 +1103,7 @@ class BotMemoryCache: return async with cls._lock: cls._by_id[entry.bot_id] = entry - cls._negative.pop(entry.bot_id, None) + RuntimeCacheMutation.clear_negative_key(cls, entry.bot_id) @classmethod async def remove(cls, bot_id: str | None) -> None: @@ -944,7 +1112,8 @@ class BotMemoryCache: return async with cls._lock: cls._by_id.pop(bot_id, None) - RuntimeCacheSync.publish_event("bot", "delete", {"bot_id": bot_id}) + RuntimeCacheMutation.clear_negative_key(cls, bot_id) + RuntimeCacheMutation.publish("bot", "delete", {"bot_id": bot_id}) @classmethod async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None: @@ -986,6 +1155,8 @@ class GroupMemoryCache: _negative: ClassVar[dict[tuple[str, str], float]] = {} _loaded: ClassVar[bool] = False _refresh_task: ClassVar[asyncio.Task | None] = None + _last_refresh: ClassVar[float] = 0.0 + _last_error: ClassVar[str | None] = None @classmethod def _normalize(cls, value: str | None) -> str | None: @@ -1014,7 +1185,7 @@ class GroupMemoryCache: if not expire_at: return False if expire_at <= time.time(): - cls._negative.pop(key, None) + RuntimeCacheMutation.clear_negative_key(cls, key) return False return True @@ -1038,15 +1209,15 @@ class GroupMemoryCache: if key: by_key[key] = entry cls._by_key = by_key - cls._negative = {} - cls._loaded = True + RuntimeCacheMutation.clear_negative_all(cls) + RuntimeCacheMutation.mark_refreshed(cls) 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() + await RuntimeCacheMutation.ensure_loaded(cls, "group") @classmethod def is_loaded(cls) -> bool: @@ -1094,8 +1265,8 @@ class GroupMemoryCache: return async with cls._lock: cls._by_key[key] = entry - cls._negative.pop(key, None) - RuntimeCacheSync.publish_event("group", "upsert", entry.to_payload()) + RuntimeCacheMutation.clear_negative_key(cls, key) + RuntimeCacheMutation.publish("group", "upsert", entry.to_payload()) @classmethod async def upsert_from_payload(cls, payload: dict[str, Any]) -> None: @@ -1105,7 +1276,7 @@ class GroupMemoryCache: return async with cls._lock: cls._by_key[key] = entry - cls._negative.pop(key, None) + RuntimeCacheMutation.clear_negative_key(cls, key) @classmethod async def remove(cls, group_id: str | None, channel_id: str | None = None) -> None: @@ -1114,7 +1285,8 @@ class GroupMemoryCache: return async with cls._lock: cls._by_key.pop(key, None) - RuntimeCacheSync.publish_event( + RuntimeCacheMutation.clear_negative_key(cls, key) + RuntimeCacheMutation.publish( "group", "delete", {"group_id": key[0], "channel_id": key[1] or None} ) @@ -1186,7 +1358,7 @@ class LevelUserMemoryCache: if not expire_at: return False if expire_at <= time.time(): - cls._negative.pop(key, None) + RuntimeCacheMutation.clear_negative_key(cls, key) return False return True @@ -1215,16 +1387,15 @@ class LevelUserMemoryCache: by_user_max[entry.user_id] = entry.user_level cls._by_key = by_key cls._by_user_max = by_user_max - cls._negative = {} - cls._loaded = True - cls._last_refresh = time.time() + RuntimeCacheMutation.clear_negative_all(cls) + RuntimeCacheMutation.mark_refreshed(cls) 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() + await RuntimeCacheMutation.ensure_loaded(cls, "level") @classmethod def is_loaded(cls) -> bool: @@ -1310,13 +1481,13 @@ class LevelUserMemoryCache: async with cls._lock: prev = cls._by_key.get(key) cls._by_key[key] = entry - cls._negative.pop(key, None) current = cls._by_user_max.get(entry.user_id, 0) if entry.user_level >= current: cls._by_user_max[entry.user_id] = entry.user_level elif prev and prev.user_level == current and entry.user_level < current: cls._recalc_user_max(entry.user_id) - RuntimeCacheSync.publish_event("level", "upsert", entry.to_payload()) + RuntimeCacheMutation.clear_negative_key(cls, key) + RuntimeCacheMutation.publish("level", "upsert", entry.to_payload()) @classmethod async def upsert_from_payload(cls, payload: dict[str, Any]) -> None: @@ -1327,12 +1498,12 @@ class LevelUserMemoryCache: async with cls._lock: prev = cls._by_key.get(key) cls._by_key[key] = entry - cls._negative.pop(key, None) current = cls._by_user_max.get(entry.user_id, 0) if entry.user_level >= current: cls._by_user_max[entry.user_id] = entry.user_level elif prev and prev.user_level == current and entry.user_level < current: cls._recalc_user_max(entry.user_id) + RuntimeCacheMutation.clear_negative_key(cls, key) @classmethod async def remove(cls, user_id: str | None, group_id: str | None) -> None: @@ -1343,7 +1514,8 @@ class LevelUserMemoryCache: removed = cls._by_key.pop(key, None) if removed and cls._by_user_max.get(removed.user_id) == removed.user_level: cls._recalc_user_max(removed.user_id) - RuntimeCacheSync.publish_event( + RuntimeCacheMutation.clear_negative_key(cls, key) + RuntimeCacheMutation.publish( "level", "delete", {"user_id": key[0], "group_id": key[1] or None} ) @@ -1399,6 +1571,8 @@ class TaskInfoMemoryCache: _negative: ClassVar[dict[str, float]] = {} _loaded: ClassVar[bool] = False _refresh_task: ClassVar[asyncio.Task | None] = None + _last_refresh: ClassVar[float] = 0.0 + _last_error: ClassVar[str | None] = None @classmethod def _normalize(cls, module: str | None) -> str | None: @@ -1417,7 +1591,7 @@ class TaskInfoMemoryCache: if not expire_at: return False if expire_at <= time.time(): - cls._negative.pop(module, None) + RuntimeCacheMutation.clear_negative_key(cls, module) return False return True @@ -1443,8 +1617,8 @@ class TaskInfoMemoryCache: by_name[entry.name] = entry cls._by_module = by_module cls._by_name = by_name - cls._negative = {} - cls._loaded = True + RuntimeCacheMutation.clear_negative_all(cls) + RuntimeCacheMutation.mark_refreshed(cls) logger.debug( f"task info cache refreshed: {len(cls._by_module)} entries", LOG_COMMAND, @@ -1454,7 +1628,7 @@ class TaskInfoMemoryCache: async def ensure_loaded(cls) -> None: if cls._loaded: return - await cls.refresh() + await RuntimeCacheMutation.ensure_loaded(cls, "task") @classmethod async def get(cls, module: str | None) -> TaskInfoSnapshot | None: @@ -1488,10 +1662,21 @@ class TaskInfoMemoryCache: @classmethod async def is_disabled(cls, module: str | None) -> bool: + """Backward-compatible runtime disabled check for passive tasks.""" + return await cls.is_runtime_disabled(module) + + @classmethod + async def is_runtime_disabled(cls, module: str | None) -> bool: + """Return whether a passive task is unavailable at runtime. + + Runtime passive availability is defined by TaskInfo.status and + TaskInfo.load_status. Bot/group scoped block lists are checked by + CommonUtils.task_is_block(). + """ entry = await cls.get(module) if not entry: return False - return not entry.status + return not entry.status or not entry.load_status @classmethod async def upsert_from_model(cls, record) -> None: @@ -1500,8 +1685,8 @@ class TaskInfoMemoryCache: cls._by_module[entry.module] = entry if entry.name: cls._by_name[entry.name] = entry - cls._negative.pop(entry.module, None) - RuntimeCacheSync.publish_event("task", "upsert", entry.to_payload()) + RuntimeCacheMutation.clear_negative_key(cls, entry.module) + RuntimeCacheMutation.publish("task", "upsert", entry.to_payload()) @classmethod async def upsert_from_payload(cls, payload: dict[str, Any]) -> None: @@ -1512,7 +1697,7 @@ class TaskInfoMemoryCache: cls._by_module[entry.module] = entry if entry.name: cls._by_name[entry.name] = entry - cls._negative.pop(entry.module, None) + RuntimeCacheMutation.clear_negative_key(cls, entry.module) @classmethod async def remove(cls, module: str | None) -> None: @@ -1525,7 +1710,8 @@ class TaskInfoMemoryCache: current = cls._by_name.get(removed.name) if current and current.module == removed.module: cls._by_name.pop(removed.name, None) - RuntimeCacheSync.publish_event("task", "delete", {"module": module}) + RuntimeCacheMutation.clear_negative_key(cls, module) + RuntimeCacheMutation.publish("task", "delete", {"module": module}) @classmethod async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None: @@ -1568,6 +1754,8 @@ class PluginLimitMemoryCache: _negative: ClassVar[dict[str, float]] = {} _loaded: ClassVar[bool] = False _refresh_task: ClassVar[asyncio.Task | None] = None + _last_refresh: ClassVar[float] = 0.0 + _last_error: ClassVar[str | None] = None @classmethod def _normalize(cls, value: str | None) -> str | None: @@ -1586,7 +1774,7 @@ class PluginLimitMemoryCache: if not expire_at: return False if expire_at <= time.time(): - cls._negative.pop(module, None) + RuntimeCacheMutation.clear_negative_key(cls, module) return False return True @@ -1611,8 +1799,8 @@ class PluginLimitMemoryCache: by_module.setdefault(entry.module, []).append(entry) cls._by_id = by_id cls._by_module = by_module - cls._negative = {} - cls._loaded = True + RuntimeCacheMutation.clear_negative_all(cls) + RuntimeCacheMutation.mark_refreshed(cls) logger.debug( f"plugin limit cache refreshed: {len(by_id)} entries", LOG_COMMAND, @@ -1622,7 +1810,7 @@ class PluginLimitMemoryCache: async def ensure_loaded(cls) -> None: if cls._loaded: return - await cls.refresh() + await RuntimeCacheMutation.ensure_loaded(cls, "plugin_limit") @classmethod def is_loaded(cls) -> bool: @@ -1668,7 +1856,7 @@ class PluginLimitMemoryCache: 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()) + RuntimeCacheMutation.publish("plugin_limit", "upsert", entry.to_payload()) @classmethod async def upsert_from_payload(cls, payload: dict[str, Any]) -> None: @@ -1695,6 +1883,7 @@ class PluginLimitMemoryCache: for item in cls._by_module.get(entry.module, []) if item.id != entry.id ] + RuntimeCacheMutation.clear_negative_key(cls, entry.module) return cls._by_id[entry.id] = entry module_limits = [ @@ -1704,7 +1893,7 @@ class PluginLimitMemoryCache: ] module_limits.append(entry) cls._by_module[entry.module] = module_limits - cls._negative.pop(entry.module, None) + RuntimeCacheMutation.clear_negative_key(cls, entry.module) @classmethod async def remove_by_id(cls, limit_id: int | None) -> None: @@ -1718,7 +1907,8 @@ class PluginLimitMemoryCache: for item in cls._by_module.get(entry.module, []) if item.id != entry.id ] - RuntimeCacheSync.publish_event("plugin_limit", "delete", {"id": int(limit_id)}) + RuntimeCacheMutation.clear_negative_key(cls, entry.module) + RuntimeCacheMutation.publish("plugin_limit", "delete", {"id": int(limit_id)}) @classmethod async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None: @@ -1764,6 +1954,8 @@ class BanMemoryCache: _refresh_task: ClassVar[asyncio.Task | None] = None _cleanup_task: ClassVar[asyncio.Task | None] = None _remove_tasks: ClassVar[set[asyncio.Task]] = set() + _last_refresh: ClassVar[float] = 0.0 + _last_error: ClassVar[str | None] = None @classmethod def _normalize_id(cls, value: str | None) -> str | None: @@ -1788,7 +1980,7 @@ class BanMemoryCache: if not expire_at: return False if expire_at <= time.time(): - cls._negative.pop(key, None) + RuntimeCacheMutation.clear_negative_key(cls, key) return False return True @@ -1842,8 +2034,8 @@ class BanMemoryCache: cls._by_user = by_user cls._by_group = by_group cls._by_user_group = by_user_group - cls._negative = {} - cls._loaded = True + RuntimeCacheMutation.clear_negative_all(cls) + RuntimeCacheMutation.mark_refreshed(cls) logger.debug( "ban cache refreshed: " f"user={len(by_user)} group={len(by_group)} " @@ -1855,7 +2047,7 @@ class BanMemoryCache: async def ensure_loaded(cls) -> None: if cls._loaded: return - await cls.refresh() + await RuntimeCacheMutation.ensure_loaded(cls, "ban") @classmethod def is_loaded(cls) -> bool: @@ -1873,13 +2065,13 @@ class BanMemoryCache: 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()) + RuntimeCacheMutation.clear_negative_all(cls) + RuntimeCacheMutation.publish("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( + RuntimeCacheMutation.publish( "ban", "delete", {"user_id": user_id, "group_id": group_id} ) @@ -1894,7 +2086,7 @@ class BanMemoryCache: cls._by_user.pop(user_id, None) elif group_id: cls._by_group.pop(group_id, None) - cls._negative = {} + RuntimeCacheMutation.clear_negative_all(cls) @classmethod def _get_entry(cls, user_id: str | None, group_id: str | None) -> BanEntry | None: @@ -1995,7 +2187,7 @@ class BanMemoryCache: elif entry.group_id: cls._by_group.pop(entry.group_id, None) if expired: - cls._negative = {} + RuntimeCacheMutation.clear_negative_all(cls) if not delete_db or not expired: return from tortoise.expressions import Q @@ -2064,7 +2256,7 @@ class BanMemoryCache: cls._by_user[entry.user_id] = entry elif entry.group_id: cls._by_group[entry.group_id] = entry - cls._negative = {} + RuntimeCacheMutation.clear_negative_all(cls) @classmethod async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None: @@ -2078,12 +2270,126 @@ class BanMemoryCache: async def _safe_refresh(cache_cls: type, label: str) -> None: """安全地刷新单个缓存,异常不影响其他缓存。""" + if getattr(cache_cls, "_loaded", False): + last_refresh = float(getattr(cache_cls, "_last_refresh", 0.0) or 0.0) + if time.time() - last_refresh <= RUNTIME_CACHE_STARTUP_REFRESH_SKIP_SECONDS: + logger.debug(f"{label} cache startup refresh skipped", LOG_COMMAND) + return try: await cache_cls.refresh() except Exception as exc: + RuntimeCacheMutation.mark_error(cache_cls, exc) logger.error(f"{label} cache init failed", LOG_COMMAND, e=exc) +def _cache_health( + cache_cls: type, + *, + entry_count: int, + negative_count: int = 0, +) -> dict[str, Any]: + return { + "loaded": bool(getattr(cache_cls, "_loaded", False)), + "entry_count": entry_count, + "last_refresh": float(getattr(cache_cls, "_last_refresh", 0.0) or 0.0), + "negative_count": negative_count, + "last_error": getattr(cache_cls, "_last_error", None), + } + + +def health_snapshot() -> dict[str, dict[str, Any]]: + """Return in-memory runtime cache health without touching the database.""" + return { + "plugin": _cache_health( + PluginInfoMemoryCache, + entry_count=len(PluginInfoMemoryCache._by_module), + ), + "bot": _cache_health( + BotMemoryCache, + entry_count=len(BotMemoryCache._by_id), + negative_count=len(BotMemoryCache._negative), + ), + "group": _cache_health( + GroupMemoryCache, + entry_count=len(GroupMemoryCache._by_key), + negative_count=len(GroupMemoryCache._negative), + ), + "level": _cache_health( + LevelUserMemoryCache, + entry_count=len(LevelUserMemoryCache._by_key), + negative_count=len(LevelUserMemoryCache._negative), + ), + "task": _cache_health( + TaskInfoMemoryCache, + entry_count=len(TaskInfoMemoryCache._by_module), + negative_count=len(TaskInfoMemoryCache._negative), + ), + "plugin_limit": _cache_health( + PluginLimitMemoryCache, + entry_count=len(PluginLimitMemoryCache._by_id), + negative_count=len(PluginLimitMemoryCache._negative), + ), + "ban": _cache_health( + BanMemoryCache, + entry_count=( + len(BanMemoryCache._by_user) + + len(BanMemoryCache._by_group) + + len(BanMemoryCache._by_user_group) + ), + negative_count=len(BanMemoryCache._negative), + ), + } + + +def passive_status_snapshot(max_modules: int = 50) -> dict[str, Any]: + """Return passive-task state from in-memory caches only. + + This is a local diagnostic helper: it does not query or write the database, + and it is not used by runtime decisions. + """ + tasks = list(TaskInfoMemoryCache._by_module.values()) + disabled = sorted(task.module for task in tasks if not task.status) + unloaded = sorted(task.module for task in tasks if not task.load_status) + runtime_enabled = [ + task.module for task in tasks if task.status and task.load_status + ] + bot_block_total = sum( + len(_parse_block_modules(bot.block_tasks)) + for bot in BotMemoryCache._by_id.values() + ) + group_block_total = sum( + len(group.block_task_set) + len(group.superuser_block_task_set) + for group in GroupMemoryCache._by_key.values() + ) + return { + "cache": health_snapshot(), + "passive_tasks": { + "total": len(tasks), + "status_enabled": sum(1 for task in tasks if task.status), + "load_status_enabled": sum(1 for task in tasks if task.load_status), + "runtime_enabled": len(runtime_enabled), + "disabled_modules": disabled[:max_modules], + "disabled_modules_total": len(disabled), + "unloaded_modules": unloaded[:max_modules], + "unloaded_modules_total": len(unloaded), + }, + "scoped_blocks": { + "bot_block_tasks_total": bot_block_total, + "group_block_tasks_total": group_block_total, + }, + "semantics": { + "available_tasks": "management_display_mirror_not_runtime_whitelist", + "runtime_truth": [ + "TaskInfo.status", + "TaskInfo.load_status", + "BotConsole.block_tasks", + "GroupConsole.block_task", + "GroupConsole.superuser_block_task", + ], + }, + } + + @PriorityLifecycle.on_startup(priority=6) async def _init_runtime_cache(): await RuntimeCacheSync.start() diff --git a/zhenxun/services/data_access.py b/zhenxun/services/data_access.py index ec0ef714..c3713f1c 100644 --- a/zhenxun/services/data_access.py +++ b/zhenxun/services/data_access.py @@ -12,6 +12,10 @@ T = TypeVar("T", bound=Model) class DataAccess(Generic[T]): """数据访问兼容层,根据配置保留单点缓存读取和清理能力 + 边界说明:DataAccess 面向普通业务查询和低频管理链路。权限检查热路径 + 必须优先使用 RuntimeCache/AuthSnapshot,避免高并发消息处理时触发 + DB/cache-aside 读放大。 + 新的高频运行态路径应优先使用 RuntimeCache 或 BoundedTTLCache。 这里不再把 filter/all/create/update_or_create 结果写入通用缓存, create/update_or_create 只负责清理旧缓存,避免旧值残留。 diff --git a/zhenxun/services/db_context/config.py b/zhenxun/services/db_context/config.py index c4889ffe..98d7d8c0 100644 --- a/zhenxun/services/db_context/config.py +++ b/zhenxun/services/db_context/config.py @@ -5,6 +5,9 @@ from pydantic import BaseModel # 数据库操作超时设置(秒) DB_TIMEOUT_SECONDS = 3.0 +# 启动期自动补齐字段/索引可能需要等待数据库锁或扫描较大的表,单独放宽超时 +DB_SCHEMA_GUARD_TIMEOUT_SECONDS = 30.0 + # 性能监控阈值(秒) SLOW_QUERY_THRESHOLD = 0.5 diff --git a/zhenxun/services/db_context/schema_guard.py b/zhenxun/services/db_context/schema_guard.py index fd00ae96..2ef9955a 100644 --- a/zhenxun/services/db_context/schema_guard.py +++ b/zhenxun/services/db_context/schema_guard.py @@ -13,7 +13,7 @@ from tortoise.exceptions import OperationalError from zhenxun.services.log import logger -from .config import DB_TIMEOUT_SECONDS, LOG_COMMAND +from .config import DB_SCHEMA_GUARD_TIMEOUT_SECONDS, LOG_COMMAND Dialect = Literal["sqlite", "postgres", "mysql", "unknown"] @@ -466,7 +466,7 @@ async def repair_safe_schema_drift() -> SchemaGuardResult: try: await asyncio.wait_for( connection.execute_query_dict(sql), - timeout=DB_TIMEOUT_SECONDS, + timeout=DB_SCHEMA_GUARD_TIMEOUT_SECONDS, ) columns[source] = ColumnInfo(name=source, data_type="") result.repaired_columns += 1 @@ -521,7 +521,7 @@ async def repair_safe_schema_drift() -> SchemaGuardResult: try: await asyncio.wait_for( connection.execute_query_dict(sql), - timeout=DB_TIMEOUT_SECONDS, + timeout=DB_SCHEMA_GUARD_TIMEOUT_SECONDS, ) existing_indexes.add(index_columns) result.repaired_indexes += 1 @@ -529,6 +529,16 @@ async def repair_safe_schema_drift() -> SchemaGuardResult: f"SchemaGuard 已补齐索引: {table}.{index_columns}", LOG_COMMAND, ) + except TimeoutError as exc: + result.warnings += 1 + result.skipped_indexes += 1 + logger.warning( + "SchemaGuard 补齐索引超时,已跳过: " + f"{table}.{index_columns} " + f"({DB_SCHEMA_GUARD_TIMEOUT_SECONDS}s)", + LOG_COMMAND, + e=exc, + ) except OperationalError as exc: err = str(exc).lower() if any( diff --git a/zhenxun/services/runtime_bootstrap.py b/zhenxun/services/runtime_bootstrap.py index 91cdd993..a813fea9 100644 --- a/zhenxun/services/runtime_bootstrap.py +++ b/zhenxun/services/runtime_bootstrap.py @@ -140,9 +140,6 @@ def register_runtime_bootstrap(_driver) -> None: global _thread_executor await _stop_launcher_watchdog() await stop_send_queue() - from zhenxun.models._bot_message_buffer import stop_bot_message_store_buffer - - await stop_bot_message_store_buffer() await stop_memory_governor() executor = _thread_executor _thread_executor = None diff --git a/zhenxun/services/scheduler/engine.py b/zhenxun/services/scheduler/engine.py index c5652514..9a3f7747 100644 --- a/zhenxun/services/scheduler/engine.py +++ b/zhenxun/services/scheduler/engine.py @@ -27,6 +27,7 @@ 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.platform import PlatformUtils from zhenxun.utils.pydantic_compat import parse_as from .repository import ScheduleRepository @@ -37,6 +38,37 @@ SCHEDULE_CONCURRENCY_KEY = "all_groups_concurrency_limit" _LAST_PRESSURE_SKIP = 0.0 +def _resolve_scheduler_bot(bot_id: str | None, log_target: str) -> Bot | None: + if bot_id: + try: + return nonebot.get_bot(bot_id) + except KeyError: + logger.warning(f"{log_target} 需要的 Bot {bot_id} 不在线,本次执行跳过。") + return None + + bots = list(nonebot.get_bots().values()) + if not bots: + logger.warning(f"{log_target} 当前没有可用 Bot,本次执行跳过。") + return None + if len(bots) == 1: + return bots[0] + + qq_client_bots = [ + bot for bot in bots if PlatformUtils.get_platform_scope(bot) == "qq_client" + ] + if len(qq_client_bots) == 1: + bot = qq_client_bots[0] + logger.warning( + f"{log_target} 未指定 Bot,多 Bot 在线,自动选择 OneBot {bot.self_id}。" + ) + return bot + + logger.warning( + f"{log_target} 未指定 Bot 且多 Bot 在线,无法安全选择," "本次执行跳过。" + ) + return None + + class APSchedulerAdapter: """封装对 APScheduler 的操作""" @@ -343,7 +375,12 @@ async def _execute_job( return try: - bot = nonebot.get_bot() + bot = _resolve_scheduler_bot( + context_override.bot_id, + f"临时任务 {plugin_name}", + ) + if bot is None: + return logger.info(f"开始执行临时任务: {plugin_name}") injected_params = {"context": context_override} state: T_State = {ScheduleContext: context_override} @@ -380,18 +417,9 @@ async def _execute_job( logger.warning(f"定时任务 {schedule_id} 不存在或已禁用,跳过执行。") return - try: - bot = ( - nonebot.get_bot(schedule.bot_id) - if schedule.bot_id - else nonebot.get_bot() - ) - except (KeyError, ValueError): - logger.warning( - f"任务 {schedule_id} 需要的 Bot {schedule.bot_id} " - f"不在线,本次执行跳过。" - ) - raise + bot = _resolve_scheduler_bot(schedule.bot_id, f"任务 {schedule_id}") + if bot is None: + return resolver = scheduler_manager._target_resolvers.get(schedule.target_type) if not resolver: diff --git a/zhenxun/services/uninfo_patch.py b/zhenxun/services/uninfo_patch.py index 2793f53b..08d43743 100644 --- a/zhenxun/services/uninfo_patch.py +++ b/zhenxun/services/uninfo_patch.py @@ -1,6 +1,7 @@ import asyncio from collections.abc import Awaitable, Callable import contextlib +import importlib from typing import Any, cast from nonebot.adapters import Bot, Event @@ -10,6 +11,9 @@ from nonebot.log import logger _PATCHED = False _ORIGINAL_FETCH: Callable[..., Awaitable[Any]] | None = None _ORIGINAL_ONEBOT11_GROUP_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None +_ORIGINAL_QQ_C2C_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None +_ORIGINAL_QQ_GROUP_AT_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None +_ORIGINAL_QQ_GUILD_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None def _sender_value(sender: Any, key: str, default: Any = None) -> Any: @@ -81,6 +85,88 @@ async def _fast_onebot11_group_message(bot: Bot, event: Event) -> dict[str, Any] } +def _qq_bot_app_id(bot: Bot) -> str: + bot_info = getattr(bot, "bot_info", None) + app_id = getattr(bot_info, "id", None) + return str(app_id or getattr(bot, "self_id", "")) + + +async def _fast_qq_c2c_message(bot: Bot, event: Event) -> dict[str, Any]: + """Build Uninfo session for QQ official C2C messages from event fields.""" + + author = _event_value(event, "author") + user_id = str( + _sender_value(author, "user_openid") + or _sender_value(author, "id") + or _event_value(event, "user_id", "") + ) + username = str(_sender_value(author, "username", "") or "") + return { + "user_id": user_id, + "name": username, + "nickname": username, + "avatar": f"https://q.qlogo.cn/qqapp/{_qq_bot_app_id(bot)}/{user_id}/100", + } + + +async def _fast_qq_group_at_message(bot: Bot, event: Event) -> dict[str, Any]: + """Build Uninfo session for QQ official group-at messages from event fields.""" + + author = _event_value(event, "author") + user_id = str( + _sender_value(author, "member_openid") + or _sender_value(author, "id") + or _event_value(event, "user_id", "") + ) + username = str(_sender_value(author, "username", "") or "") + group_id = str( + _event_value(event, "group_openid") or _event_value(event, "group_id") or "" + ) + return { + "user_id": user_id, + "name": username, + "nickname": username, + "avatar": f"https://q.qlogo.cn/qqapp/{_qq_bot_app_id(bot)}/{user_id}/100", + "group_id": group_id, + } + + +async def _fast_qq_guild_message(bot: Bot, event: Event) -> dict[str, Any]: + """Build Uninfo session for QQ official guild/channel messages locally. + + nonebot-plugin-uninfo enriches guild messages through remote guild/channel + APIs. Runtime auth only needs stable scene/user ids, so avoid remote calls + during matcher fanout. + """ + + author = _event_value(event, "author") + member = _event_value(event, "member") + guild_id = str(_event_value(event, "guild_id", "") or "") + channel_id = str(_event_value(event, "channel_id", "") or "") + user_id = str(_sender_value(author, "id", "") or "") + nickname = str(_sender_value(member, "nick", "") or "") + username = str(_sender_value(author, "username", "") or "") + base: dict[str, Any] = { + "user_id": user_id, + "name": username, + "nickname": nickname or username, + "avatar": _sender_value(author, "avatar"), + "guild_id": guild_id, + "channel_id": channel_id, + "guild_name": "", + "guild_avatar": None, + "channel_name": "", + "channel_type": -1, + } + roles = _sender_value(member, "roles") + if roles is not None: + base["roles"] = roles + joined_at = _sender_value(member, "joined_at") + if joined_at is not None: + base["joined_at"] = joined_at + return base + + async def _singleflight_fetch(self: Any, bot: Bot, event: Event) -> Any: original = _ORIGINAL_FETCH if original is None: @@ -114,6 +200,8 @@ async def _singleflight_fetch(self: Any, bot: Bot, event: Event) -> Any: def apply_uninfo_onebot11_patch() -> None: global _ORIGINAL_FETCH, _ORIGINAL_ONEBOT11_GROUP_MESSAGE, _PATCHED + global _ORIGINAL_QQ_C2C_MESSAGE, _ORIGINAL_QQ_GROUP_AT_MESSAGE + global _ORIGINAL_QQ_GUILD_MESSAGE if _PATCHED: return @@ -129,6 +217,59 @@ def apply_uninfo_onebot11_patch() -> None: setattr(_fast_onebot11_group_message, "__zhenxun_fast_onebot11__", True) fetcher.endpoint[GroupMessageEvent] = _fast_onebot11_group_message + with contextlib.suppress(Exception): + qq_event_module = importlib.import_module("nonebot.adapters.qq.event") + AtMessageCreateEvent = getattr(qq_event_module, "AtMessageCreateEvent") + C2CMessageCreateEvent = getattr(qq_event_module, "C2CMessageCreateEvent") + DirectMessageCreateEvent = getattr(qq_event_module, "DirectMessageCreateEvent") + GroupAtMessageCreateEvent = getattr( + qq_event_module, + "GroupAtMessageCreateEvent", + ) + GroupMessageCreateEvent = getattr( + qq_event_module, + "GroupMessageCreateEvent", + ) + MessageCreateEvent = getattr(qq_event_module, "MessageCreateEvent") + from nonebot_plugin_uninfo.adapters.qq.main import fetcher as qq_fetcher + + original_c2c = qq_fetcher.endpoint.get(C2CMessageCreateEvent) + if not getattr(original_c2c, "__zhenxun_fast_qq__", False): + _ORIGINAL_QQ_C2C_MESSAGE = cast( + Callable[..., Awaitable[dict[str, Any]]] | None, + original_c2c, + ) + setattr(_fast_qq_c2c_message, "__zhenxun_fast_qq__", True) + qq_fetcher.endpoint[C2CMessageCreateEvent] = _fast_qq_c2c_message + + for event_type in (GroupMessageCreateEvent, GroupAtMessageCreateEvent): + original_group_at = qq_fetcher.endpoint.get(event_type) + if getattr(original_group_at, "__zhenxun_fast_qq__", False): + continue + if _ORIGINAL_QQ_GROUP_AT_MESSAGE is None and original_group_at is not None: + _ORIGINAL_QQ_GROUP_AT_MESSAGE = cast( + Callable[..., Awaitable[dict[str, Any]]], + original_group_at, + ) + setattr(_fast_qq_group_at_message, "__zhenxun_fast_qq__", True) + qq_fetcher.endpoint[event_type] = _fast_qq_group_at_message + + for event_type in ( + MessageCreateEvent, + AtMessageCreateEvent, + DirectMessageCreateEvent, + ): + original_guild = qq_fetcher.endpoint.get(event_type) + if getattr(original_guild, "__zhenxun_fast_qq__", False): + continue + if _ORIGINAL_QQ_GUILD_MESSAGE is None and original_guild is not None: + _ORIGINAL_QQ_GUILD_MESSAGE = cast( + Callable[..., Awaitable[dict[str, Any]]], + original_guild, + ) + setattr(_fast_qq_guild_message, "__zhenxun_fast_qq__", True) + qq_fetcher.endpoint[event_type] = _fast_qq_guild_message + try: from nonebot_plugin_uninfo.fetch import InfoFetcher except Exception as e: @@ -146,4 +287,4 @@ def apply_uninfo_onebot11_patch() -> None: setattr(_singleflight_fetch, "__zhenxun_singleflight__", True) setattr(InfoFetcher, "fetch", _singleflight_fetch) _PATCHED = True - logger.debug("Uninfo OneBot11 fast fetch and singleflight patch applied") + logger.debug("Uninfo fast fetch and singleflight patch applied") diff --git a/zhenxun/utils/common_utils.py b/zhenxun/utils/common_utils.py index 40cdd6d2..39318701 100644 --- a/zhenxun/utils/common_utils.py +++ b/zhenxun/utils/common_utils.py @@ -23,29 +23,31 @@ class CommonUtils: async def task_is_block( cls, session: Uninfo | Bot, module: str, group_id: str | None = None ) -> bool: - """判断被动技能是否可以发送 + """判断被动技能是否被阻断。 + + 运行真源固定为 TaskInfo.status/load_status、BotConsole.block_tasks、 + GroupConsole.block_task/superuser_block_task,以及 bot/group ban 状态。 + BotConsole.available_tasks 只用于管理展示,不作为运行白名单。 参数: module: 被动技能模块名 group_id: 群组id 返回: - bool: 是否可以发送 + bool: True 表示被动技能应被阻断,False 表示允许继续执行 """ if isinstance(session, Bot): if interface := get_interface(session): info = interface.basic_info() if info["scope"] == SupportScope.qq_api: - logger.info("q官bot放弃所有被动技能发言...") - """q官bot放弃所有被动技能发言""" - return False - if session.scene == SupportScope.qq_api: - """q官bot放弃所有被动技能发言""" - logger.info("q官bot放弃所有被动技能发言...") - return False + logger.debug("q官bot放弃所有被动技能发言...") + return True + if isinstance(session, Session) and session.scope == SupportScope.qq_api: + logger.debug("q官bot放弃所有被动技能发言...") + return True if not group_id and isinstance(session, Session): group_id = session.group.id if session.group else None - if await TaskInfoMemoryCache.is_disabled(module): + if await TaskInfoMemoryCache.is_runtime_disabled(module): """被动全局状态""" return True bot_snapshot = await BotMemoryCache.get(session.self_id) diff --git a/zhenxun/utils/enum.py b/zhenxun/utils/enum.py index 0412ee4a..76b8eff8 100644 --- a/zhenxun/utils/enum.py +++ b/zhenxun/utils/enum.py @@ -13,11 +13,6 @@ class PriorityLifecycleType(StrEnum): """关闭""" -class BotSentType(StrEnum): - GROUP = "GROUP" - PRIVATE = "PRIVATE" - - class BankHandleType(StrEnum): DEPOSIT = "DEPOSIT" """存款""" diff --git a/zhenxun/utils/platform.py b/zhenxun/utils/platform.py index 416b37a5..12dfc830 100644 --- a/zhenxun/utils/platform.py +++ b/zhenxun/utils/platform.py @@ -1,5 +1,6 @@ import asyncio from collections.abc import Awaitable, Callable +import contextlib import random from typing import cast @@ -37,6 +38,23 @@ def _adapter_name(bot: Bot) -> str: return adapter.__class__.__name__.lower() +def _scope_name(scope: object) -> str: + """Normalize uninfo/alconna scope values without changing legacy platform.""" + if scope is None: + return "" + raw = getattr(scope, "name", None) or getattr(scope, "value", scope) + text = str(raw or "").strip().lower() + if not text: + text = str(getattr(scope, "name", "") or "").strip().lower() + text = text.replace("-", "_").replace(" ", "_") + compact = "".join(ch for ch in text if ch.isalnum()) + if compact.endswith("qqclient"): + return "qq_client" + if compact.endswith("qqapi"): + return "qq_api" + return text + + class UserData(BaseModel): name: str """昵称""" @@ -57,6 +75,31 @@ class UserData(BaseModel): class PlatformUtils: + @classmethod + def _resolve_unique_qq_client_bot(cls, log_cmd: str | None = None) -> Bot | None: + bots = list(nonebot.get_bots().values()) + if not bots: + logger.warning("当前没有可用的 OneBot 协议端 Bot,已跳过。", log_cmd) + return None + qq_client_bots = [ + bot for bot in bots if cls.get_platform_scope(bot) == "qq_client" + ] + if len(qq_client_bots) == 1: + bot = qq_client_bots[0] + if len(bots) > 1: + logger.warning( + f"多 Bot 在线且未指定 Bot,自动选择 OneBot {bot.self_id}。", + log_cmd, + ) + return bot + if not qq_client_bots: + logger.warning("未找到 OneBot 协议端 Bot,已跳过。", log_cmd) + else: + logger.warning( + "存在多个 OneBot 协议端 Bot,无法安全选择,已跳过。", log_cmd + ) + return None + @classmethod def is_qbot(cls, session: Uninfo | Bot) -> bool: """判断bot是否为qq官bot @@ -68,7 +111,11 @@ class PlatformUtils: bool: 是否为官bot """ if isinstance(session, Bot): + if cls.get_platform_scope(session) == "qq_api": + return True return bool(BotConfig.get_qbot_uid(session.self_id)) + if cls.get_platform_scope(session) == "qq_api": + return True if BotConfig.get_qbot_uid(session.self_id): return True return session.scope == SupportScope.qq_api @@ -83,7 +130,7 @@ class PlatformUtils: group_id: 群组id duration: 禁言时长(分钟) """ - if cls.get_platform(bot) == "qq": + if cls.get_platform_scope(bot) == "qq_client": await bot.set_group_ban( group_id=int(group_id), user_id=int(user_id), @@ -111,7 +158,9 @@ class PlatformUtils: Receipt | None: Receipt """ if not bot: - bot = nonebot.get_bot() + bot = cls._resolve_unique_qq_client_bot("PlatformUtils:send_superuser") + if bot is None: + return [] superuser_ids = [] if superuser_id: superuser_ids.append(superuser_id) @@ -397,6 +446,11 @@ class PlatformUtils: def get_platform_scope(cls, t: Bot | Uninfo | object) -> str: """获取细粒度平台作用域,不改变旧 get_platform 返回值。""" if isinstance(t, Bot): + if interface := get_interface(t): + with contextlib.suppress(Exception): + scope = _scope_name(interface.basic_info().get("scope")) + if scope: + return scope adapter_name = _adapter_name(t) if "onebot" in adapter_name: return "qq_client" @@ -406,17 +460,13 @@ class PlatformUtils: return "qq_api" return adapter_name or cls.get_platform(t) - scope = str(getattr(t, "scope", "") or "").lower() + scope = _scope_name(getattr(t, "scope", "") or "") if not scope: basic = getattr(t, "basic", None) if isinstance(basic, dict): - scope = str(basic.get("scope") or "").lower() - if "qq_client" in scope: - return "qq_client" - if "qq_api" in scope: - return "qq_api" - if scope.startswith("qq"): - return "qq" + scope = _scope_name(basic.get("scope")) + if scope: + return scope adapter = getattr(t, "adapter", None) if adapter is not None: @@ -499,6 +549,9 @@ class PlatformUtils: 返回: int: 更新个数 """ + if cls.get_platform_scope(bot) == "qq_api": + logger.warning("QQ 官方适配器不支持旧好友同步,已跳过。", "更新好友信息") + return 0 create_list = [] friend_list, platform = await cls.get_friend_list(bot) if friend_list: @@ -521,6 +574,11 @@ class PlatformUtils: 返回: list[FriendUser]: 好友列表 """ + if cls.get_platform_scope(bot) == "qq_api": + logger.warning( + "QQ 官方适配器不支持旧好友列表查询,已返回空列表。", "好友列表" + ) + return [], cls.get_platform(bot) if interface := get_interface(bot): user_list = await interface.get_users() return [ @@ -602,14 +660,9 @@ class BroadcastEngine: except KeyError: logger.warning(f"Bot:{i} 对象未连接或不存在", log_cmd) if not self.bot_list: - try: - bot = nonebot.get_bot() + bot = PlatformUtils._resolve_unique_qq_client_bot(log_cmd) + if bot is not None: self.bot_list.append(bot) - logger.warning( - f"广播任务未传入Bot对象,使用默认Bot {bot.self_id}", log_cmd - ) - except Exception as e: - raise ValueError("当前没有可用的Bot对象...", log_cmd) from e async def call_check(self, bot: Bot, group_id: str) -> bool: """运行发送检测函数