From 5d92ccd3b09d0af9da25efd0b0060a816d59487f Mon Sep 17 00:00:00 2001 From: Copaan <98086483+Copaan@users.noreply.github.com> Date: Sun, 26 Apr 2026 15:50:15 +0800 Subject: [PATCH] =?UTF-8?q?=E6=80=A7=E8=83=BD=E4=BC=98=E5=8C=96=20(#2126)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * 性能优化 * 代码改进 * 优化浏览器代际切换逻辑 * 统一缓存与生命周期 * 添加aiomysql依赖 * 优化插件路径处理逻辑,简化条件判断;在虚拟环境包管理器中添加编码和错误处理参数以增强稳定性 * :rotating_light: auto fix by pre-commit hooks * 优化Windows下的关闭逻辑 * 代码优化 * bugfix:修复配置重载问题 * bugfix:修复插件加载启动竞态问题 * 收敛事件入口和权限上下文 * 优化 Windows launcher 关闭重启兜底 --------- Co-authored-by: HibiKier <775757368@qq.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> --- pyproject.toml | 1 + uv.lock | 23 + .../admin/plugin_switch/data_source.py | 5 +- .../admin/plugin_switch/strategy.py | 4 +- .../chat_history/chat_message.py | 35 +- .../builtin_plugins/hooks/auth/auth_admin.py | 17 +- .../builtin_plugins/hooks/auth/auth_ban.py | 55 +- .../builtin_plugins/hooks/auth/auth_bot.py | 6 + .../builtin_plugins/hooks/auth/auth_cost.py | 9 +- .../builtin_plugins/hooks/auth/auth_group.py | 8 + .../builtin_plugins/hooks/auth/auth_limit.py | 51 +- .../builtin_plugins/hooks/auth/auth_plugin.py | 23 +- .../builtin_plugins/hooks/auth/bot_filter.py | 17 +- zhenxun/builtin_plugins/hooks/auth/context.py | 321 +++++++++ zhenxun/builtin_plugins/hooks/auth_checker.py | 608 ++++++++++++------ zhenxun/builtin_plugins/hooks/auth_hook.py | 145 ++--- zhenxun/builtin_plugins/hooks/chkdsk_hook.py | 38 +- zhenxun/builtin_plugins/init/__init_cache.py | 1 - zhenxun/builtin_plugins/init/init_plugin.py | 2 +- .../plugin_store/data_source.py | 49 +- zhenxun/builtin_plugins/shop/_data_source.py | 12 +- .../builtin_plugins/sign_in/_data_source.py | 31 +- .../statistics/statistics_hook.py | 45 +- .../superuser/reload_setting.py | 78 ++- zhenxun/builtin_plugins/web_ui/__init__.py | 24 +- .../web_ui/api/logs/log_manager.py | 88 ++- .../builtin_plugins/web_ui/api/logs/logs.py | 12 +- .../web_ui/api/tabs/main/__init__.py | 65 +- .../web_ui/api/tabs/plugin_manage/__init__.py | 40 +- .../api/tabs/plugin_manage/data_source.py | 11 +- zhenxun/cli.py | 142 +++- zhenxun/configs/utils/__init__.py | 127 +++- zhenxun/models/group_member_info.py | 53 +- zhenxun/models/user_console.py | 26 +- zhenxun/services/avatar_service.py | 5 + zhenxun/services/buffered_writers.py | 128 ++++ zhenxun/services/cache/__init__.py | 16 +- zhenxun/services/cache/bounded_ttl.py | 218 +++++++ zhenxun/services/cache/cache_containers.py | 86 ++- zhenxun/services/cache/runtime_cache.py | 211 ++++-- zhenxun/services/data_access.py | 77 +-- zhenxun/services/db_context/base_model.py | 2 +- zhenxun/services/group_settings_service.py | 9 +- zhenxun/services/memory_governor.py | 307 +++++++++ zhenxun/services/message_load.py | 11 + zhenxun/services/renderer/engine.py | 85 ++- zhenxun/services/renderer/result_cache.py | 64 +- zhenxun/services/renderer/service.py | 12 + zhenxun/services/renderer/theme.py | 19 +- zhenxun/services/runtime_bootstrap.py | 66 +- zhenxun/services/send_queue.py | 104 ++- zhenxun/services/uninfo_patch.py | 149 +++++ zhenxun/utils/enum.py | 2 - zhenxun/utils/http_utils.py | 46 +- zhenxun/utils/manager/message_manager.py | 68 +- .../manager/virtual_env_package_manager.py | 41 ++ 56 files changed, 3092 insertions(+), 806 deletions(-) create mode 100644 zhenxun/builtin_plugins/hooks/auth/context.py create mode 100644 zhenxun/services/buffered_writers.py create mode 100644 zhenxun/services/cache/bounded_ttl.py create mode 100644 zhenxun/services/memory_governor.py create mode 100644 zhenxun/services/uninfo_patch.py diff --git a/pyproject.toml b/pyproject.toml index 52d154c0..e508cfe4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -45,6 +45,7 @@ dependencies = [ "alibabacloud-devops20210625>=5.0.2,<6.0.0", "uvloop>=0.21.0; sys_platform != 'win32'", "pytest-timeout>=2.4.0", + "aiomysql>=0.3.2", ] [project.scripts] diff --git a/uv.lock b/uv.lock index 6fa05389..dd9d37d1 100644 --- a/uv.lock +++ b/uv.lock @@ -158,6 +158,18 @@ wheels = [ { url = "https://mirrors.aliyun.com/pypi/packages/62/29/2f8418269e46454a26171bfdd6a055d74febf32234e474930f2f60a17145/aiohttp-3.13.5-cp314-cp314t-win_amd64.whl", hash = "sha256:18a2f6c1182c51baa1d28d68fea51513cb2a76612f038853c0ad3c145423d3d9" }, ] +[[package]] +name = "aiomysql" +version = "0.3.2" +source = { registry = "https://mirrors.aliyun.com/pypi/simple/" } +dependencies = [ + { name = "pymysql" }, +] +sdist = { url = "https://mirrors.aliyun.com/pypi/packages/29/e0/302aeffe8d90853556f47f3106b89c16cc2ec2a4d269bdfd82e3f4ae12cc/aiomysql-0.3.2.tar.gz", hash = "sha256:72d15ef5cfc34c03468eb41e1b90adb9fd9347b0b589114bd23ead569a02ac1a" } +wheels = [ + { url = "https://mirrors.aliyun.com/pypi/packages/4c/af/aae0153c3e28712adaf462328f6c7a3c196a1c1c27b491de4377dd3e6b52/aiomysql-0.3.2-py3-none-any.whl", hash = "sha256:c82c5ba04137d7afd5c693a258bea8ead2aad77101668044143a991e04632eb2" }, +] + [[package]] name = "aiosignal" version = "1.4.0" @@ -2598,6 +2610,15 @@ wheels = [ { url = "https://mirrors.aliyun.com/pypi/packages/f7/27/a2fc51a4a122dfd1015e921ae9d22fee3d20b0b8080d9a704578bf9deece/pymdown_extensions-10.21.2-py3-none-any.whl", hash = "sha256:5c0fd2a2bea14eb39af8ff284f1066d898ab2187d81b889b75d46d4348c01638" }, ] +[[package]] +name = "pymysql" +version = "1.1.2" +source = { registry = "https://mirrors.aliyun.com/pypi/simple/" } +sdist = { url = "https://mirrors.aliyun.com/pypi/packages/f5/ae/1fe3fcd9f959efa0ebe200b8de88b5a5ce3e767e38c7ac32fb179f16a388/pymysql-1.1.2.tar.gz", hash = "sha256:4961d3e165614ae65014e361811a724e2044ad3ea3739de9903ae7c21f539f03" } +wheels = [ + { url = "https://mirrors.aliyun.com/pypi/packages/7c/4c/ad33b92b9864cbde84f259d5df035a6447f91891f5be77788e2a3892bce3/pymysql-1.1.2-py3-none-any.whl", hash = "sha256:e6b1d89711dd51f8f74b1631fe08f039e7d76cf67a42a323d3178f0f25762ed9" }, +] + [[package]] name = "pypika-tortoise" version = "0.1.6" @@ -4157,6 +4178,7 @@ source = { editable = "." } dependencies = [ { name = "aiocache", extra = ["redis"] }, { name = "aiofiles" }, + { name = "aiomysql" }, { name = "alibabacloud-devops20210625" }, { name = "asyncpg" }, { name = "beautifulsoup4" }, @@ -4212,6 +4234,7 @@ dev = [ requires-dist = [ { name = "aiocache", extras = ["redis"], specifier = ">=0.12.3" }, { name = "aiofiles", specifier = ">=23.2.1" }, + { name = "aiomysql", specifier = ">=0.3.2" }, { name = "alibabacloud-devops20210625", specifier = ">=5.0.2,<6.0.0" }, { name = "asyncpg", specifier = ">=0.20.0" }, { name = "beautifulsoup4", specifier = ">=4.12.3,<5.0.0" }, diff --git a/zhenxun/builtin_plugins/admin/plugin_switch/data_source.py b/zhenxun/builtin_plugins/admin/plugin_switch/data_source.py index 7118e74d..cabea9b2 100644 --- a/zhenxun/builtin_plugins/admin/plugin_switch/data_source.py +++ b/zhenxun/builtin_plugins/admin/plugin_switch/data_source.py @@ -1,10 +1,9 @@ from nonebot.adapters import Bot from zhenxun.models.group_console import GroupConsole -from zhenxun.services.cache import CacheRoot from zhenxun.services.cache.runtime_cache import GroupMemoryCache from zhenxun.utils.common_utils import CommonUtils -from zhenxun.utils.enum import BlockType, CacheType +from zhenxun.utils.enum import BlockType from zhenxun.utils.platform import PlatformUtils from .strategy import get_strategy @@ -134,7 +133,6 @@ class PluginManager: await GroupConsole.bulk_update( update_list, [norm_field, su_field], batch_size=500 ) - await CacheRoot.clear(CacheType.GROUPS) for group in update_list: await GroupMemoryCache.upsert_from_model(group) @@ -318,7 +316,6 @@ class PluginManager: status=False ) - await CacheRoot.clear(CacheType.GROUPS) await GroupMemoryCache.refresh() action_str = "醒来" if status else "休眠" diff --git a/zhenxun/builtin_plugins/admin/plugin_switch/strategy.py b/zhenxun/builtin_plugins/admin/plugin_switch/strategy.py index d4c3e17d..fb1c484e 100644 --- a/zhenxun/builtin_plugins/admin/plugin_switch/strategy.py +++ b/zhenxun/builtin_plugins/admin/plugin_switch/strategy.py @@ -4,12 +4,11 @@ from typing import Any, cast from zhenxun.models.group_console import GroupConsole from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.task_info import TaskInfo -from zhenxun.services.cache import CacheRoot from zhenxun.services.cache.runtime_cache import ( PluginInfoMemoryCache, TaskInfoMemoryCache, ) -from zhenxun.utils.enum import BlockType, CacheType, PluginType +from zhenxun.utils.enum import BlockType, PluginType class SwitchStrategy(ABC): @@ -135,7 +134,6 @@ class PluginStrategy(SwitchStrategy): await self.refresh_cache() async def refresh_cache(self) -> None: - await CacheRoot.invalidate_cache(CacheType.PLUGINS) await PluginInfoMemoryCache.refresh() diff --git a/zhenxun/builtin_plugins/chat_history/chat_message.py b/zhenxun/builtin_plugins/chat_history/chat_message.py index aead64d4..60780964 100644 --- a/zhenxun/builtin_plugins/chat_history/chat_message.py +++ b/zhenxun/builtin_plugins/chat_history/chat_message.py @@ -10,6 +10,7 @@ from nonebot_plugin_uninfo import Uninfo from zhenxun.configs.config import Config from zhenxun.configs.utils import PluginExtraData, RegisterConfig from zhenxun.models.chat_history import ChatHistory +from zhenxun.services.db_context import with_db_timeout from zhenxun.services.log import logger from zhenxun.services.message_load import is_overloaded, should_pause_tasks from zhenxun.utils.enum import PluginType @@ -47,6 +48,9 @@ _HISTORY_QUEUE: asyncio.Queue[ChatHistory] = asyncio.Queue(maxsize=5000) _DROP_COUNT = 0 _LAST_DROP_LOG = 0.0 _DROP_LOG_INTERVAL = 10.0 +_FLUSH_BATCH_SIZE = 200 +_FLUSH_MAX_PER_TICK = 1000 +_FLUSH_DB_TIMEOUT = 5.0 @chat_history.handle() @@ -80,19 +84,34 @@ async def _(message: UniMsg, session: Uninfo): @scheduler.scheduled_job( "interval", minutes=1, + max_instances=1, + coalesce=True, ) async def _(): try: if should_pause_tasks(): return - message_list: list[ChatHistory] = [] - while True: - try: - message_list.append(_HISTORY_QUEUE.get_nowait()) - except asyncio.QueueEmpty: + flushed = 0 + while flushed < _FLUSH_MAX_PER_TICK: + message_list: list[ChatHistory] = [] + limit = min(_FLUSH_BATCH_SIZE, _FLUSH_MAX_PER_TICK - flushed) + for _ in range(limit): + try: + message_list.append(_HISTORY_QUEUE.get_nowait()) + except asyncio.QueueEmpty: + break + if not message_list: break - if message_list: - await ChatHistory.bulk_create(message_list) - logger.debug(f"批量添加聊天记录 {len(message_list)} 条", "定时任务") + await with_db_timeout( + ChatHistory.bulk_create(message_list, _FLUSH_BATCH_SIZE), + timeout=_FLUSH_DB_TIMEOUT, + operation=f"ChatHistory.bulk_create[{len(message_list)}]", + source="chat_history", + ) + flushed += len(message_list) + if flushed: + backlog = _HISTORY_QUEUE.qsize() + suffix = f",剩余队列 {backlog} 条" if backlog else "" + logger.debug(f"批量添加聊天记录 {flushed} 条{suffix}", "定时任务") except Exception as e: logger.warning("存储聊天记录失败", "chat_history", e=e) diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_admin.py b/zhenxun/builtin_plugins/hooks/auth/auth_admin.py index aed470e7..9c307580 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_admin.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_admin.py @@ -7,9 +7,10 @@ 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 get_entity_ids +from zhenxun.utils.utils import EntityIDs, get_entity_ids from .config import LOGGER_COMMAND, WARNING_THRESHOLD +from .context import PermissionContext from .exception import SkipPluginException @@ -20,6 +21,9 @@ async def auth_admin( LevelUser | LevelUserSnapshot | None, LevelUser | LevelUserSnapshot | None ] | None = None, + *, + context: PermissionContext | None = None, + entity: EntityIDs | None = None, ): """管理员命令 个人权限 @@ -33,7 +37,12 @@ async def auth_admin( return try: - entity = get_entity_ids(session) + if context is not None: + entity = context.entity + if cached_levels is None: + cached_levels = context.admin_levels + if entity is None: + entity = get_entity_ids(session) global_user: LevelUser | LevelUserSnapshot | None = None group_users: LevelUser | LevelUserSnapshot | None = None @@ -42,7 +51,7 @@ async def auth_admin( global_user, group_users = cached_levels else: global_user, group_users = await LevelUserMemoryCache.get_levels( - session.user.id, entity.group_id + entity.user_id, entity.group_id ) user_level = global_user.user_level if global_user else 0 @@ -53,7 +62,7 @@ async def auth_admin( raise SkipPluginException( f"{plugin.name}({plugin.module}) 管理员权限不足...", tip_message=[ - At(flag="user", target=session.user.id), + At(flag="user", target=entity.user_id), f"你的权限不足喔,该功能需要的权限等级: {plugin.admin_level}", ], tip_check_tag=entity.user_id, diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_ban.py b/zhenxun/builtin_plugins/hooks/auth/auth_ban.py index 6b1ee105..bc1f660a 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_ban.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_ban.py @@ -1,4 +1,3 @@ -import asyncio import time from nonebot.matcher import Matcher @@ -8,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.cache_containers import CacheDict from zhenxun.services.cache.runtime_cache import BanMemoryCache from zhenxun.services.db_context import DB_TIMEOUT_SECONDS from zhenxun.services.log import logger @@ -16,6 +14,7 @@ from zhenxun.utils.enum import PluginType from zhenxun.utils.utils import EntityIDs, get_entity_ids from .config import LOGGER_COMMAND, WARNING_THRESHOLD +from .context import PermissionContext from .exception import SkipPluginException from .utils import freq @@ -25,37 +24,6 @@ Config.add_plugin_config( "才不会给你发消息.", help="对被ban用户发送的消息", ) -BAN_CACHE_TTL = 2 -BAN_CACHE_TTL_POSITIVE = 30 -BAN_CACHE_TTL_NEGATIVE = 5 - -BAN_CACHE = ( - CacheDict("AUTH_BAN_CACHE", expire=0) - if max(BAN_CACHE_TTL_POSITIVE, BAN_CACHE_TTL_NEGATIVE) > 0 - else None -) - - -def _ban_cache_key(user_id: str | None, group_id: str | None) -> str: - return f"{user_id or ''}:{group_id or ''}" - - -def _ban_cache_get(key: str) -> int | None: - if not BAN_CACHE: - return None - try: - return BAN_CACHE[key] - except KeyError: - return None - - -def _ban_cache_set(key: str, value: int) -> None: - if not BAN_CACHE: - return - ttl = BAN_CACHE_TTL_POSITIVE if value else BAN_CACHE_TTL_NEGATIVE - if ttl <= 0: - return - BAN_CACHE.set(key, value, expire=ttl) async def calculate_ban_time(ban_record: BanConsole | None) -> int: @@ -214,6 +182,7 @@ async def auth_ban( session: Uninfo, plugin: PluginInfo, *, + context: PermissionContext | None = None, entity: EntityIDs | None = None, is_superuser: bool = False, ) -> None: @@ -229,28 +198,18 @@ async def auth_ban( return if not matcher.plugin_name: return + if context is not None: + entity = context.entity + is_superuser = context.is_superuser if entity is None: entity = get_entity_ids(session) if is_superuser: return if entity.group_id: - try: - await asyncio.wait_for( - group_handle(entity.group_id), timeout=DB_TIMEOUT_SECONDS - ) - except asyncio.TimeoutError: - logger.error(f"群组ban检查超时: {entity.group_id}", LOGGER_COMMAND) - # 超时时不阻塞,继续执行 + await group_handle(entity.group_id) if entity.user_id: - try: - await asyncio.wait_for( - user_handle(plugin, entity, session), - timeout=DB_TIMEOUT_SECONDS, - ) - except asyncio.TimeoutError: - logger.error(f"用户ban检查超时: {entity.user_id}", LOGGER_COMMAND) - # 超时时不阻塞,继续执行 + await user_handle(plugin, entity, session) finally: # 记录总执行时间 elapsed = time.time() - start_time diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_bot.py b/zhenxun/builtin_plugins/hooks/auth/auth_bot.py index 20e64efb..6a0b30b6 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_bot.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_bot.py @@ -7,6 +7,7 @@ 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 .exception import SkipPluginException @@ -16,6 +17,8 @@ async def auth_bot( bot_data: BotConsole | BotSnapshot | None = None, skip_fetch: bool = False, allow_sleep_bypass: bool = False, + *, + context: PermissionContext | None = None, ): """bot层面的权限检查 @@ -30,6 +33,9 @@ async def auth_bot( start_time = time.time() try: + 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) diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_cost.py b/zhenxun/builtin_plugins/hooks/auth/auth_cost.py index 8f2ea47e..1017edbd 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_cost.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_cost.py @@ -7,13 +7,18 @@ from zhenxun.models.user_console import UserConsole from zhenxun.services.log import logger from .config import LOGGER_COMMAND, WARNING_THRESHOLD +from .context import PermissionContext from .exception import SkipPluginException DEFAULT_GOLD = 100 async def auth_cost( - user: UserConsole | None, plugin: PluginInfo, session: Uninfo + user: UserConsole | None, + plugin: PluginInfo, + session: Uninfo, + *, + context: PermissionContext | None = None, ) -> int: """检测是否满足金币条件 @@ -28,6 +33,8 @@ async def auth_cost( start_time = time.time() try: + if context is not None and user is None: + user = context.user user_gold = user.gold if user else DEFAULT_GOLD if user_gold < plugin.cost_gold: """插件消耗金币不足""" diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_group.py b/zhenxun/builtin_plugins/hooks/auth/auth_group.py index e59bc206..29f2608a 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_group.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_group.py @@ -7,6 +7,7 @@ from zhenxun.services.cache.runtime_cache import GroupSnapshot from zhenxun.services.log import logger from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum +from .context import PermissionContext from .exception import SkipPluginException _GROUP_WAKE_PATTERN = re.compile(r"^醒来$", re.IGNORECASE) @@ -34,6 +35,8 @@ async def auth_group( group: GroupConsole | GroupSnapshot | None, text: str | None, group_id: str | None, + *, + context: PermissionContext | None = None, ): """群黑名单检测 群总开关检测 @@ -42,6 +45,11 @@ async def auth_group( group: GroupConsole message: UniMsg """ + if context is not None: + group = context.group or group + text = context.plain_text + group_id = context.group_id + if not group_id: return diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_limit.py b/zhenxun/builtin_plugins/hooks/auth/auth_limit.py index e8201001..e16eb0f2 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_limit.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_limit.py @@ -19,9 +19,10 @@ from zhenxun.utils.limiters import CountLimiter, FreqLimiter, UserBlockLimiter from zhenxun.utils.manager.priority_manager import PriorityLifecycle from zhenxun.utils.message import MessageUtils from zhenxun.utils.time_utils import TimeUtils -from zhenxun.utils.utils import get_entity_ids +from zhenxun.utils.utils import EntityIDs, get_entity_ids from .config import LOGGER_COMMAND, WARNING_THRESHOLD +from .context import PermissionContext from .exception import SkipPluginException driver = nonebot.get_driver() @@ -31,7 +32,7 @@ _LIMIT_NOTICE_LIMITER = FreqLimiter(_LIMIT_NOTICE_CD) _LIMIT_NOTICE_TASKS: set[asyncio.Task] = set() -@PriorityLifecycle.on_startup(priority=5) +@PriorityLifecycle.on_startup(priority=7) async def _(): """初始化限制""" await LimitManager.init_limit() @@ -117,6 +118,7 @@ class LimitManager: cls.cd_limit = {} cls.block_limit = {} cls.count_limit = {} + cls.module_limit_cache.clear() # 添加新数据 for limit in limit_list: cls.add_limit(limit) @@ -137,22 +139,22 @@ class LimitManager: """ if limit.module not in cls.add_module: cls.add_module.append(limit.module) - if limit.limit_type == PluginLimitType.BLOCK: - cls.block_limit[limit.module] = Limit( - limit=limit, limiter=UserBlockLimiter() - ) - elif limit.limit_type == PluginLimitType.CD: - cd_value = int(limit.cd or 0) - cls.cd_limit[limit.module] = Limit( - limit=limit, limiter=FreqLimiter(cd_value) - ) - elif limit.limit_type == PluginLimitType.COUNT: - max_count = int(limit.max_count or 0) - if max_count <= 0: - return - cls.count_limit[limit.module] = Limit( - limit=limit, limiter=CountLimiter(max_count) - ) + if limit.limit_type == PluginLimitType.BLOCK: + cls.block_limit[limit.module] = Limit( + limit=limit, limiter=UserBlockLimiter() + ) + elif limit.limit_type == PluginLimitType.CD: + cd_value = int(limit.cd or 0) + cls.cd_limit[limit.module] = Limit( + limit=limit, limiter=FreqLimiter(cd_value) + ) + elif limit.limit_type == PluginLimitType.COUNT: + max_count = int(limit.max_count or 0) + if max_count <= 0: + return + cls.count_limit[limit.module] = Limit( + limit=limit, limiter=CountLimiter(max_count) + ) @classmethod def unblock( @@ -322,14 +324,23 @@ class LimitManager: limiter.increase(key_type) -async def auth_limit(plugin: PluginInfo, session: Uninfo): +async def auth_limit( + plugin: PluginInfo, + session: Uninfo, + *, + context: PermissionContext | None = None, + entity: EntityIDs | None = None, +): """插件限制 参数: plugin: PluginInfo session: Uninfo """ - entity = get_entity_ids(session) + if context is not None: + entity = context.entity + if entity is None: + entity = get_entity_ids(session) try: await asyncio.wait_for( LimitManager.check( diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_plugin.py b/zhenxun/builtin_plugins/hooks/auth/auth_plugin.py index 77841967..13b42fce 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_plugin.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_plugin.py @@ -10,6 +10,7 @@ from zhenxun.services.log import logger from zhenxun.utils.enum import BlockType from .config import LOGGER_COMMAND, WARNING_THRESHOLD +from .context import PermissionContext from .exception import IsSuperuserException, SkipPluginException from .utils import freq, is_poke @@ -107,11 +108,16 @@ class GroupCheck: class PluginCheck: def __init__( - self, group: GroupConsole | GroupSnapshot | None, session: Uninfo, is_poke: bool + self, + group: GroupConsole | GroupSnapshot | None, + session: Uninfo, + is_poke: bool, + user_id: str | None, ): self.session = session self.is_poke = is_poke self.group_data = group + self.user_id = user_id or session.user.id self.group_id = None if group: self.group_id = group.group_id @@ -126,13 +132,11 @@ class PluginCheck: IgnoredException: 忽略插件 """ if plugin.block_type == BlockType.PRIVATE: - should_tip = freq.is_send_limit_message( - plugin, self.session.user.id, self.is_poke - ) + should_tip = freq.is_send_limit_message(plugin, self.user_id, self.is_poke) raise SkipPluginException( f"{plugin.name}({plugin.module}) 该插件在私聊中已被禁用...", tip_message="该功能在私聊中已被禁用..." if should_tip else None, - tip_check_tag=self.session.user.id if should_tip else None, + tip_check_tag=self.user_id if should_tip else None, tip_background=should_tip, ) @@ -153,7 +157,7 @@ class PluginCheck: if self.group_data and self.group_data.is_super: raise IsSuperuserException() - sid = self.group_id or self.session.user.id + sid = self.group_id or self.user_id should_tip = freq.is_send_limit_message(plugin, sid, self.is_poke) raise SkipPluginException( f"{plugin.name}({plugin.module}) 全局未开启此功能...", @@ -176,7 +180,9 @@ async def auth_plugin( session: Uninfo, event: Event, *, + context: PermissionContext | None = None, skip_group_block: bool = False, + user_id: str | None = None, ): """插件状态 @@ -187,8 +193,11 @@ async def auth_plugin( """ start_time = time.time() try: + if context is not None: + group = context.group or group + user_id = context.user_id is_poke_event = is_poke(event) - user_check = PluginCheck(group, session, is_poke_event) + user_check = PluginCheck(group, session, is_poke_event, user_id) if group: block_set, super_block_set = _get_group_block_sets(group) diff --git a/zhenxun/builtin_plugins/hooks/auth/bot_filter.py b/zhenxun/builtin_plugins/hooks/auth/bot_filter.py index 04e47372..f6b5da1c 100644 --- a/zhenxun/builtin_plugins/hooks/auth/bot_filter.py +++ b/zhenxun/builtin_plugins/hooks/auth/bot_filter.py @@ -3,6 +3,7 @@ from nonebot_plugin_uninfo import Uninfo from zhenxun.configs.config import Config +from .context import PermissionContext from .exception import SkipPluginException Config.add_plugin_config( @@ -15,7 +16,12 @@ Config.add_plugin_config( ) -def bot_filter(session: Uninfo): +def bot_filter( + session: Uninfo, + *, + context: PermissionContext | None = None, + user_id: str | None = None, +): """过滤bot调用bot 参数: @@ -26,10 +32,13 @@ def bot_filter(session: Uninfo): """ if not Config.get_config("hook", "FILTER_BOT"): return + if context is not None: + user_id = context.user_id bot_ids = list(nonebot.get_bots().keys()) - if session.user.id == session.self_id: + checked_user_id = user_id or session.user.id + if checked_user_id == session.self_id: return - if session.user.id in bot_ids: + if checked_user_id in bot_ids: raise SkipPluginException( - f"bot:{session.self_id} 尝试调用 bot:{session.user.id}" + f"bot:{session.self_id} 尝试调用 bot:{checked_user_id}" ) diff --git a/zhenxun/builtin_plugins/hooks/auth/context.py b/zhenxun/builtin_plugins/hooks/auth/context.py new file mode 100644 index 00000000..3e270ffb --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth/context.py @@ -0,0 +1,321 @@ +from __future__ import annotations + +import asyncio +import contextlib +from dataclasses import dataclass, field +from typing import Any + +from nonebot.adapters import Bot, Event +from nonebot_plugin_alconna import UniMsg +from nonebot_plugin_uninfo import Uninfo + +from zhenxun.services.cache.cache_containers import CacheDict +from zhenxun.utils.platform import PlatformUtils +from zhenxun.utils.utils import EntityIDs, get_entity_ids + +AUTH_EVENT_CACHE_TTL = 5 + +STATE_EVENT_CONTEXT = "_zx_event_context" +STATE_PERMISSION_CONTEXT = "_zx_permission_context" +STATE_ENTITY = "_zx_entity" +STATE_EVENT_CACHE = "_zx_event_cache" +STATE_PLAIN_TEXT = "_zx_plain_text" +STATE_ROUTE_MODULES = "_zx_route_modules" +STATE_IS_SUPERUSER = "_zx_is_superuser" +STATE_PERMISSION_SIDE_EFFECTS = "_zx_permission_side_effects" +EVENT_CACHE_PERMISSION_SIDE_EFFECTS = "permission_side_effects" + +EVENT_CACHE = ( + CacheDict("AUTH_EVENT_CACHE", expire=AUTH_EVENT_CACHE_TTL) + if AUTH_EVENT_CACHE_TTL > 0 + else None +) + + +@dataclass +class EventContext: + bot_id: str + platform: str + event_type: str + message_id: str | int | None + entity: EntityIDs + plain_text: str = "" + route_modules: set[str] = field(default_factory=set) + route_modules_loaded: bool = False + is_superuser: bool = False + event_cache: dict[str, Any] | None = None + + @property + def user_id(self) -> str: + return self.entity.user_id + + @property + def group_id(self) -> str | None: + return self.entity.group_id + + @property + def channel_id(self) -> str | None: + return self.entity.channel_id + + +@dataclass +class PermissionSideEffectCache: + auth_results: dict[str, tuple[bool, str | None]] = field(default_factory=dict) + module_locks: dict[str, asyncio.Lock] = field(default_factory=dict) + + def lock_for(self, module: str) -> asyncio.Lock: + lock = self.module_locks.get(module) + if lock is None: + lock = asyncio.Lock() + self.module_locks[module] = lock + return lock + + +@dataclass +class PermissionContext: + event: EventContext + module: str + plugin: Any = None + user: Any = None + group: Any = None + bot_data: Any = None + admin_levels: Any = None + + @property + def entity(self) -> EntityIDs: + return self.event.entity + + @property + def user_id(self) -> str: + return self.event.user_id + + @property + def group_id(self) -> str | None: + return self.event.group_id + + @property + def channel_id(self) -> str | None: + return self.event.channel_id + + @property + def plain_text(self) -> str: + return self.event.plain_text + + @property + def is_superuser(self) -> bool: + return self.event.is_superuser + + +def resolve_actor_user_id(event: Event, fallback_user_id: str | None) -> str: + """优先使用事件发起者 ID,避免 notice 场景 session.user 指向 bot 自身。""" + event_user_id = getattr(event, "user_id", None) + if event_user_id is None: + return fallback_user_id or "" + resolved = str(event_user_id) + return resolved or fallback_user_id or "" + + +def resolve_event_group_id(event: Event, fallback_group_id: str | None) -> str | None: + """notice 场景 session.group 可能缺失,回退到事件上的 group_id。""" + event_group_id = getattr(event, "group_id", None) + if event_group_id is None: + return fallback_group_id + resolved = str(event_group_id) + return resolved or fallback_group_id + + +def resolve_event_channel_id( + event: Event, fallback_channel_id: str | None +) -> str | None: + """频道场景回退到事件上的 channel_id。""" + event_channel_id = getattr(event, "channel_id", None) + if event_channel_id is None: + return fallback_channel_id + resolved = str(event_channel_id) + return resolved or fallback_channel_id + + +def resolve_entity_ids(event: Event, session: Uninfo) -> EntityIDs: + entity = get_entity_ids(session) + entity.user_id = resolve_actor_user_id(event, entity.user_id) + entity.group_id = resolve_event_group_id(event, entity.group_id) + entity.channel_id = resolve_event_channel_id(event, entity.channel_id) + return entity + + +def extract_plain_text(message: UniMsg | None, event: Event) -> str: + if message is not None: + with contextlib.suppress(Exception): + return message.extract_plain_text() + with contextlib.suppress(Exception): + plain = event.get_plaintext() + if plain: + return plain.strip() + return "" + + +def _event_message_id(event: Event) -> str | int | None: + msg_id = getattr(event, "message_id", None) + if msg_id is None: + msg_id = getattr(event, "id", None) + return msg_id + + +def event_cache_key( + event: Event, + *, + bot_id: str, + platform: str, + entity: EntityIDs, +) -> str: + msg_id = _event_message_id(event) + if msg_id is None: + 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}" + + +def get_event_cache( + event: Event, + *, + bot_id: str, + platform: str, + 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) + try: + return EVENT_CACHE[key] + except KeyError: + cache: dict[str, Any] = {} + EVENT_CACHE[key] = cache + return cache + + +def _sync_context_state(state: dict[str, Any], context: EventContext) -> None: + state[STATE_EVENT_CONTEXT] = context + state[STATE_ENTITY] = context.entity + state[STATE_EVENT_CACHE] = context.event_cache + state[STATE_PLAIN_TEXT] = context.plain_text + state[STATE_ROUTE_MODULES] = context.route_modules + state[STATE_IS_SUPERUSER] = context.is_superuser + get_permission_side_effect_cache(state=state, event_cache=context.event_cache) + + +def get_permission_side_effect_cache( + *, + state: dict[str, Any] | None = None, + event_cache: dict[str, Any] | None = None, +) -> PermissionSideEffectCache: + side_effects = None + if state is not None: + side_effects = state.get(STATE_PERMISSION_SIDE_EFFECTS) + if ( + not isinstance(side_effects, PermissionSideEffectCache) + and event_cache is not None + ): + side_effects = event_cache.get(EVENT_CACHE_PERMISSION_SIDE_EFFECTS) + if not isinstance(side_effects, PermissionSideEffectCache): + side_effects = PermissionSideEffectCache() + if state is not None: + state[STATE_PERMISSION_SIDE_EFFECTS] = side_effects + if event_cache is not None: + event_cache[EVENT_CACHE_PERMISSION_SIDE_EFFECTS] = side_effects + return side_effects + + +def get_event_context(state: dict[str, Any] | None) -> EventContext | None: + if state is None: + return None + context = state.get(STATE_EVENT_CONTEXT) + return context if isinstance(context, EventContext) else None + + +def get_or_create_event_context( + bot: Bot, + event: Event, + session: Uninfo, + state: dict[str, Any], + *, + message: UniMsg | None = None, +) -> EventContext: + context = get_event_context(state) + if context is not None: + _sync_context_state(state, context) + return context + + entity = state.get(STATE_ENTITY) + if not isinstance(entity, EntityIDs): + entity = resolve_entity_ids(event, session) + + platform = PlatformUtils.get_platform(session) + bot_id = str(bot.self_id) + event_cache = state.get(STATE_EVENT_CACHE) + if not isinstance(event_cache, dict): + event_cache = get_event_cache( + event, + bot_id=bot_id, + platform=platform, + entity=entity, + ) + + text = state.get(STATE_PLAIN_TEXT) + if not isinstance(text, str): + cached_text = event_cache.get("plain_text") if event_cache is not None else None + text = ( + cached_text + if isinstance(cached_text, str) + else extract_plain_text(message, event) + ) + if event_cache is not None: + event_cache["plain_text"] = text + + route_modules_loaded = STATE_ROUTE_MODULES in state + route_modules = state.get(STATE_ROUTE_MODULES) + if not isinstance(route_modules, set): + cached_routes = ( + event_cache.get("route_modules") if event_cache is not None else None + ) + route_modules = cached_routes if isinstance(cached_routes, set) else set() + route_modules_loaded = isinstance(cached_routes, set) + + is_superuser = state.get(STATE_IS_SUPERUSER) + if not isinstance(is_superuser, bool): + is_superuser = entity.user_id in bot.config.superusers + + context = EventContext( + bot_id=bot_id, + platform=platform, + event_type=event.get_type(), + message_id=_event_message_id(event), + entity=entity, + plain_text=text, + route_modules=route_modules, + route_modules_loaded=route_modules_loaded, + is_superuser=is_superuser, + event_cache=event_cache, + ) + _sync_context_state(state, context) + return context + + +def set_route_modules( + state: dict[str, Any] | None, + context: EventContext, + route_modules: set[str], +) -> None: + context.route_modules = route_modules + context.route_modules_loaded = True + if context.event_cache is not None: + context.event_cache["route_modules"] = route_modules + if state is not None: + _sync_context_state(state, context) + + +def store_permission_context( + state: dict[str, Any] | None, context: PermissionContext +) -> None: + if state is not None: + state[STATE_PERMISSION_CONTEXT] = context diff --git a/zhenxun/builtin_plugins/hooks/auth_checker.py b/zhenxun/builtin_plugins/hooks/auth_checker.py index 203a37d5..558ac441 100644 --- a/zhenxun/builtin_plugins/hooks/auth_checker.py +++ b/zhenxun/builtin_plugins/hooks/auth_checker.py @@ -1,7 +1,7 @@ import asyncio from collections.abc import Awaitable, Callable import contextlib -import os +import importlib import re import time from typing import cast @@ -11,7 +11,6 @@ from nonebot.adapters import Bot, Event from nonebot.exception import IgnoredException from nonebot.matcher import Matcher import nonebot.message as nb_message -from nonebot_plugin_alconna import UniMsg from nonebot_plugin_uninfo import Uninfo from zhenxun.configs.utils import PluginExtraData @@ -33,7 +32,6 @@ from zhenxun.services.message_load import is_overloaded from zhenxun.utils.enum import BlockType, GoldHandle, PluginType from zhenxun.utils.exception import InsufficientGold from zhenxun.utils.platform import PlatformUtils -from zhenxun.utils.utils import get_entity_ids from .auth.auth_admin import auth_admin from .auth.auth_ban import auth_ban @@ -44,6 +42,16 @@ from .auth.auth_limit import LimitManager, auth_limit from .auth.auth_plugin import auth_plugin from .auth.bot_filter import bot_filter from .auth.config import LOGGER_COMMAND, WARNING_THRESHOLD +from .auth.context import ( + EVENT_CACHE, + STATE_PLAIN_TEXT, + EventContext, + PermissionContext, + get_event_context, + get_permission_side_effect_cache, + set_route_modules, + store_permission_context, +) from .auth.exception import ( IsSuperuserException, PermissionExemption, @@ -53,7 +61,6 @@ from .auth.utils import send_message AUTH_HOOKS_CONCURRENCY_LIMIT = 5 AUTH_DB_CONCURRENCY_LIMIT = 6 -AUTH_EVENT_CACHE_TTL = 5 # 增加到5秒,减少缓存抖动 # 超时设置(秒) @@ -74,13 +81,6 @@ CIRCUIT_RESET_TIME = 300 # 5分钟 HOOKS_CONCURRENCY_LIMIT = AUTH_HOOKS_CONCURRENCY_LIMIT DB_CONCURRENCY_LIMIT = AUTH_DB_CONCURRENCY_LIMIT -EVENT_CACHE_TTL = AUTH_EVENT_CACHE_TTL -EVENT_CACHE = ( - CacheDict("AUTH_EVENT_CACHE", expire=EVENT_CACHE_TTL) - if EVENT_CACHE_TTL > 0 - else None -) - # 路由索引缓存 _ROUTE_INDEX_LOCK = asyncio.Lock() _ROUTE_INDEX_READY = False @@ -91,21 +91,17 @@ MATCHER_ROUTE_PREFILTER_TTL = 2 PREFILTER_STATS_LOG_INTERVAL = 10.0 CACHE_SWEEP_INTERVAL = 1.0 -CPU_COUNT = os.cpu_count() or 4 -COMMAND_MATCHER_CONCURRENCY = max(8, min(48, CPU_COUNT * 4)) -HEAVY_COMMAND_CONCURRENCY = max(1, min(3, CPU_COUNT // 2)) -HEAVY_COMMAND_MODULES = frozenset({"shop", "sign_in"}) - # 全局信号量与计数器 HOOKS_ACTIVE_COUNT = 0 HOOKS_SEMAPHORE = asyncio.Semaphore(HOOKS_CONCURRENCY_LIMIT) -COMMAND_MATCHER_SEMAPHORE = asyncio.Semaphore(COMMAND_MATCHER_CONCURRENCY) -HEAVY_COMMAND_SEMAPHORE = asyncio.Semaphore(HEAVY_COMMAND_CONCURRENCY) DB_SEMAPHORE = asyncio.Semaphore(DB_CONCURRENCY_LIMIT) DB_ACTIVE_COUNT = 0 _CHECK_MATCHER_PATCHED = False _ORIGINAL_CHECK_AND_RUN_MATCHER: Callable[..., Awaitable[None]] | None = None +_HANDLE_EVENT_PATCHED = False +_ORIGINAL_HANDLE_EVENT: Callable[..., Awaitable[None]] | None = None +_ORIGINAL_ADAPTER_HANDLE_EVENTS: dict[object, object] = {} _MATCHER_COMMAND_TYPE_CACHE: dict[type[Matcher], bool] = {} _MATCHER_COMMAND_LITERAL_CACHE: dict[type[Matcher], tuple[str, ...] | None] = {} _MATCHER_ALCONNA_SHORTCUT_CACHE: dict[type[Matcher], bool] = {} @@ -115,6 +111,10 @@ _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, @@ -163,33 +163,6 @@ def _debug_log(message: str, *args, **kwargs) -> None: logger.debug(message, *args, **kwargs) -def _event_cache_key(event: Event, session: Uninfo, entity) -> str: - msg_id = getattr(event, "message_id", None) - if msg_id is None: - msg_id = getattr(event, "id", None) - if msg_id is None: - msg_id = id(event) - platform = PlatformUtils.get_platform(session) - group_id = entity.group_id or "" - channel_id = entity.channel_id or "" - return ( - f"{platform}:{session.self_id}:{entity.user_id}:" - f"{group_id}:{channel_id}:{msg_id}" - ) - - -def _get_event_cache(event: Event, session: Uninfo, entity): - if not EVENT_CACHE: - return None - key = _event_cache_key(event, session, entity) - try: - return EVENT_CACHE[key] - except KeyError: - cache = {} - EVENT_CACHE[key] = cache - return cache - - def _normalize_command(command: str) -> str: text = command.strip() if not text: @@ -488,6 +461,9 @@ def _event_plain_text(event: Event) -> str: def _state_plain_text(state: dict | None) -> str: if state is None: return "" + context = get_event_context(state) + if context is not None: + return context.plain_text.strip() text = state.get("_zx_plain_text") if isinstance(text, str): return text.strip() @@ -496,6 +472,9 @@ def _state_plain_text(state: dict | None) -> str: def _get_route_modules_for_event(event: Event, state: dict | None = None) -> set[str]: if state is not None: + context = get_event_context(state) + if context is not None and context.route_modules_loaded: + return context.route_modules route_modules = state.get("_zx_route_modules") if isinstance(route_modules, set): return route_modules @@ -506,15 +485,67 @@ def _get_route_modules_for_event(event: Event, state: dict | None = None) -> set route_modules = _match_route_modules(_event_plain_text(event)) _CHECK_MATCHER_ROUTE_CACHE[key] = route_modules if state is not None: - state["_zx_route_modules"] = route_modules + context = get_event_context(state) + if context is not None: + set_route_modules(state, context, route_modules) + else: + state["_zx_route_modules"] = route_modules return route_modules -def _record_prefilter_stats(skipped: bool, reason: str | None) -> None: +def _prepare_handle_event_state(event: Event, state: dict) -> None: + get_permission_side_effect_cache(state=state) + if event.get_type() != "message": + return + if _state_plain_text(state): + return + text = _event_plain_text(event) + if text: + state[STATE_PLAIN_TEXT] = text + + +def _build_matcher_state(base_state: dict) -> dict: + get_permission_side_effect_cache(state=base_state) + matcher_state = base_state.copy() + get_permission_side_effect_cache(state=matcher_state) + return matcher_state + + +async def _run_selected_matcher( + matcher: type[Matcher], + bot: Bot, + event: Event, + state: dict, + stack, + dependency_cache, +) -> None: + await nb_message.check_and_run_matcher( + matcher, + bot, + event, + state, + stack, + dependency_cache, + ) + + +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": @@ -537,6 +568,10 @@ def _record_prefilter_stats(skipped: bool, reason: str | None) -> None: "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']} " @@ -643,15 +678,6 @@ def _matcher_has_alconna_shortcuts(matcher_cls: type[Matcher]) -> bool: return has_shortcuts -def _is_heavy_command_module(module: str) -> bool: - normalized = module.strip().lower() - if not normalized: - return False - if normalized in HEAVY_COMMAND_MODULES: - return True - return any(normalized.endswith(f".{name}") for name in HEAVY_COMMAND_MODULES) - - async def _check_matcher_prefilter( matcher_cls: type[Matcher], event: Event, state: dict | None = None ) -> tuple[bool, str | None]: @@ -686,6 +712,17 @@ async def _check_matcher_prefilter( if not module: return False, None + command_matched = False + matcher_commands = _extract_matcher_command_literals(matcher_cls) + if matcher_commands: + for command in matcher_commands: + if _command_matches(text, command): + command_matched = True + break + else: + if not _matcher_has_alconna_shortcuts(matcher_cls): + return True, "command_miss" + ai_route_modules = _collect_ai_route_modules(event, state) ai_route_heads = _collect_ai_route_heads(event, state) if ai_route_modules and module not in ai_route_modules: @@ -696,25 +733,85 @@ async def _check_matcher_prefilter( await _ensure_route_index() if module not in _ROUTE_MODULES_WITH_COMMANDS: - matcher_commands = _extract_matcher_command_literals(matcher_cls) - if matcher_commands: - for command in matcher_commands: - if _command_matches(text, command): - return False, None - if _matcher_has_alconna_shortcuts(matcher_cls): - return False, None - return True, "command_miss" return False, None route_modules = _get_route_modules_for_event(event, state) if module not in route_modules: + if command_matched: + return False, None if _matcher_has_alconna_shortcuts(matcher_cls): return False, None return True, "route_miss" return False, None -_MATCHER_SEMAPHORE_TIMEOUT = 8.0 +def _check_matcher_prefilter_before_task( + matcher_cls: type[Matcher], event: Event, state: dict | None = None +) -> tuple[bool, str | None]: + """Conservative selector before creating matcher task. + + This mirrors the async matcher prefilter but never performs IO or route-index + rebuild. If anything is uncertain, let the existing check_and_run_matcher + patch handle it inside the task. + """ + event_type = event.get_type() + matcher_type = getattr(matcher_cls, "type", "") or "" + if isinstance(matcher_type, str) and matcher_type and matcher_type != event_type: + return True, "type_miss" + + if event_type != "message": + return False, None + + if getattr(matcher_cls, "temp", False): + return False, None + + if not _is_command_matcher_class(matcher_cls): + return False, None + + text = _state_plain_text(state) + if not text: + text = _event_plain_text(event) + if state is not None and text: + state["_zx_plain_text"] = text + if not text: + return True, "empty_text" + + module = _matcher_module_name(matcher_cls) + if not module: + return False, None + + command_matched = False + matcher_commands = _extract_matcher_command_literals(matcher_cls) + has_alconna_shortcuts = _matcher_has_alconna_shortcuts(matcher_cls) + if matcher_commands: + for command in matcher_commands: + if _command_matches(text, command): + command_matched = True + break + else: + if not has_alconna_shortcuts: + return True, "command_miss" + + ai_route_modules = _collect_ai_route_modules(event, state) + ai_route_heads = _collect_ai_route_heads(event, state) + if ai_route_modules and module not in ai_route_modules: + if not _matcher_matches_ai_route_heads(matcher_cls, ai_route_heads): + return True, "route_miss" + + if not _ROUTE_INDEX_READY: + return False, None + + if module not in _ROUTE_MODULES_WITH_COMMANDS: + return False, None + + route_modules = _get_route_modules_for_event(event, state) + if module not in route_modules: + if command_matched or has_alconna_shortcuts: + return False, None + return True, "route_miss" + return False, None + + _MAX_MATCHER_CACHE = 512 @@ -729,7 +826,7 @@ async def _patched_check_and_run_matcher( skip, reason = await _check_matcher_prefilter( Matcher, event, state if isinstance(state, dict) else None ) - _record_prefilter_stats(skip, reason) + _record_prefilter_stats(skip, reason, "inside_task") if skip: return @@ -744,28 +841,6 @@ async def _patched_check_and_run_matcher( "stack": stack, "dependency_cache": dependency_cache, } - if _is_command_matcher_class(Matcher): - module = _matcher_module_name(Matcher) - sem = ( - HEAVY_COMMAND_SEMAPHORE - if _is_heavy_command_module(module) - else COMMAND_MATCHER_SEMAPHORE - ) - try: - await asyncio.wait_for(sem.acquire(), timeout=_MATCHER_SEMAPHORE_TIMEOUT) - except asyncio.TimeoutError: - logger.warning( - f"matcher semaphore acquire timeout for {module}, " - "executing without concurrency limit", - LOGGER_COMMAND, - ) - await original(**kwargs) - return - try: - await original(**kwargs) - finally: - sem.release() - return await original(**kwargs) @@ -788,27 +863,137 @@ def _uninstall_matcher_prefilter() -> None: _ORIGINAL_CHECK_AND_RUN_MATCHER = None -def _get_message_text( - message: UniMsg | None, - event_cache: dict | None, - event: Event | None = None, -) -> str: - if event_cache is not None: - cached = event_cache.get("plain_text") - if isinstance(cached, str): - return cached +async def _patched_handle_event(bot: Bot, event: Event) -> None: + show_log = True + escape_tag = getattr(nb_message, "escape_tag") + logger_ = getattr(nb_message, "logger") + no_log_exception = getattr(nb_message, "NoLogException") - text = "" - if message is not None: - with contextlib.suppress(Exception): - text = message.extract_plain_text() - if not text and event is not None: - with contextlib.suppress(Exception): - text = (event.get_plaintext() or "").strip() + log_msg = f"{escape_tag(bot.type)} {escape_tag(bot.self_id)} | " + try: + log_msg += event.get_log_string() + except no_log_exception: + show_log = False + if show_log: + logger_.opt(colors=True).success(log_msg) - if event_cache is not None: - event_cache["plain_text"] = text - return text + state = {} + dependency_cache = {} + async_exit_stack = getattr(nb_message, "AsyncExitStack") + apply_event_preprocessors = getattr(nb_message, "_apply_event_preprocessors") + apply_event_postprocessors = getattr(nb_message, "_apply_event_postprocessors") + trie_rule = getattr(nb_message, "TrieRule") + matchers = getattr(nb_message, "matchers") + catch = getattr(nb_message, "catch") + stop_propagation = getattr(nb_message, "StopPropagation") + handle_exception = getattr(nb_message, "_handle_exception") + anyio_mod = getattr(nb_message, "anyio") + run_coro_with_shield = getattr(nb_message, "run_coro_with_shield") + + async with async_exit_stack() as stack: + if not await apply_event_preprocessors( + bot=bot, + event=event, + state=state, + stack=stack, + dependency_cache=dependency_cache, + ): + return + + try: + trie_rule.get_value(bot, event, state) + except Exception as e: + logger_.opt(colors=True, exception=e).warning( + "Error while parsing command for event" + ) + _prepare_handle_event_state(event, state) + + break_flag = False + + def _handle_stop_propagation(_exc_group) -> None: + nonlocal break_flag + break_flag = True + logger_.debug("Stop event propagation") + + for priority in sorted(matchers.keys()): + if break_flag: + break + + if show_log: + logger_.debug(f"Checking for matchers in priority {priority}...") + + if not (priority_matchers := matchers[priority]): + continue + + with catch( + { + stop_propagation: _handle_stop_propagation, + Exception: handle_exception( + "Error when checking Matcher." + ), + } + ): + async with anyio_mod.create_task_group() as tg: + for matcher in priority_matchers: + skip, reason = _check_matcher_prefilter_before_task( + matcher, + event, + state, + ) + _record_prefilter_stats(skip, reason, "before_task") + if skip: + continue + matcher_state = _build_matcher_state(state) + tg.start_soon( + run_coro_with_shield, + _run_selected_matcher( + matcher, + bot, + event, + matcher_state, + stack, + dependency_cache, + ), + ) + + if show_log: + logger_.debug("Checking for matchers completed") + + await apply_event_postprocessors(bot, event, state, stack, dependency_cache) + + +def _install_handle_event_selector() -> None: + global _HANDLE_EVENT_PATCHED, _ORIGINAL_HANDLE_EVENT + if _HANDLE_EVENT_PATCHED: + return + _ORIGINAL_HANDLE_EVENT = nb_message.handle_event + nb_message.handle_event = _patched_handle_event # type: ignore[assignment] + for module_name in ( + "nonebot.adapters.onebot.v11.bot", + "nonebot.adapters.onebot.v12.bot", + "onebug.mixin.process", + ): + with contextlib.suppress(Exception): + module = importlib.import_module(module_name) + current = getattr(module, "handle_event", None) + if current is not None: + _ORIGINAL_ADAPTER_HANDLE_EVENTS[module] = current + setattr(module, "handle_event", _patched_handle_event) + _HANDLE_EVENT_PATCHED = True + + +def _uninstall_handle_event_selector() -> None: + global _HANDLE_EVENT_PATCHED, _ORIGINAL_HANDLE_EVENT + if not _HANDLE_EVENT_PATCHED: + return + if _ORIGINAL_HANDLE_EVENT is not None: + nb_message.handle_event = _ORIGINAL_HANDLE_EVENT # type: ignore[assignment] + for module, original in list(_ORIGINAL_ADAPTER_HANDLE_EVENTS.items()): + with contextlib.suppress(Exception): + setattr(module, "handle_event", original) + _ORIGINAL_ADAPTER_HANDLE_EVENTS.clear() + _HANDLE_EVENT_PATCHED = False + _ORIGINAL_HANDLE_EVENT = None async def _get_route_context(text: str, event_cache: dict | None) -> set[str]: @@ -843,12 +1028,14 @@ async def start_auth_runtime_tasks() -> None: global _CACHE_SWEEP_TASK await _ensure_route_index() _install_matcher_prefilter() + _install_handle_event_selector() if _CACHE_SWEEP_TASK is None or _CACHE_SWEEP_TASK.done(): _CACHE_SWEEP_TASK = asyncio.create_task(_cache_sweep_loop()) async def stop_auth_runtime_tasks() -> None: global _CACHE_SWEEP_TASK + _uninstall_handle_event_selector() _uninstall_matcher_prefilter() task = _CACHE_SWEEP_TASK _CACHE_SWEEP_TASK = None @@ -873,6 +1060,12 @@ async def _has_limits_cached(module: str, event_cache: dict | None) -> bool: @contextlib.asynccontextmanager async def _db_section(): global DB_ACTIVE_COUNT + if DB_SEMAPHORE.locked(): + logger.warning( + "db semaphore saturated, allowing permission check to continue", + LOGGER_COMMAND, + ) + raise PermissionExemption("db semaphore saturated, allow pass") await DB_SEMAPHORE.acquire() DB_ACTIVE_COUNT += 1 try: @@ -918,7 +1111,9 @@ def _group_has_plugin_block(group, module: str) -> bool: ) -def _needs_auth_plugin(plugin: PluginInfo, group, entity) -> bool: +def _needs_auth_plugin(plugin: PluginInfo, context: PermissionContext) -> bool: + group = context.group + entity = context.entity if plugin.block_type == BlockType.ALL and not plugin.status: if group and getattr(group, "is_super", False): return False @@ -953,11 +1148,11 @@ async def _get_bot_data_cached( async def _get_admin_levels_cached( - session: Uninfo, entity, event_cache + entity, event_cache ) -> tuple[tuple[LevelUserSnapshot | None, LevelUserSnapshot | None] | None, bool]: if event_cache is not None and "admin_levels" in event_cache: return event_cache.get("admin_levels"), event_cache.get("admin_timeout", False) - levels = await LevelUserMemoryCache.get_levels(session.user.id, entity.group_id) + levels = await LevelUserMemoryCache.get_levels(entity.user_id, entity.group_id) if event_cache is not None: event_cache["admin_levels"] = levels event_cache["admin_timeout"] = False @@ -1074,12 +1269,18 @@ async def get_plugin_and_user( if user_id in user_cache: user = user_cache[user_id] else: - async with _db_section(): - user = await _fetch_user_readonly(user_dao, user_id) + try: + async with _db_section(): + user = await _fetch_user_readonly(user_dao, user_id) + except PermissionExemption: + user = None user_cache[user_id] = user else: - async with _db_section(): - user = await _fetch_user_readonly(user_dao, user_id) + try: + async with _db_section(): + user = await _fetch_user_readonly(user_dao, user_id) + except PermissionExemption: + user = None return plugin, user @@ -1089,7 +1290,7 @@ async def get_plugin_cost( plugin: PluginInfo, session: Uninfo, *, - is_superuser: bool = False, + context: PermissionContext | None = None, ) -> int: """获取插件费用 @@ -1106,7 +1307,10 @@ async def get_plugin_cost( 返回: int: 调用插件金币费用 """ - cost_gold = await with_timeout(auth_cost(user, plugin, session), name="auth_cost") + cost_gold = await with_timeout( + auth_cost(user, plugin, session, context=context), name="auth_cost" + ) + is_superuser = context.is_superuser if context is not None else False if is_superuser: if plugin.plugin_type == PluginType.SUPERUSER: raise IsSuperuserException() @@ -1124,7 +1328,7 @@ async def reduce_gold(user_id: str, module: str, cost_gold: int, session: Uninfo cost_gold: 消耗金币 session: Uninfo """ - user_dao = DataAccess(UserConsole) + should_clear_cache = False try: await with_timeout( UserConsole.reduce_gold( @@ -1141,14 +1345,16 @@ async def reduce_gold(user_id: str, module: str, cost_gold: int, session: Uninfo u.gold = 0 await u.save(update_fields=["gold"]) except asyncio.TimeoutError: + should_clear_cache = True logger.error( f"扣除金币超时,用户: {user_id}, 金币: {cost_gold}", LOGGER_COMMAND, session=session, ) - # 清除缓存,使下次查询时从数据库获取最新数据 - await user_dao.clear_cache(user_id=user_id) + # 正常写入路径由 UserConsole.save() 统一失效缓存;超时状态不确定时兜底清理。 + if should_clear_cache: + await DataAccess(UserConsole).clear_cache(user_id=user_id) logger.debug(f"调用功能花费金币: {cost_gold}", LOGGER_COMMAND, session=session) @@ -1174,16 +1380,15 @@ async def time_hook(coro, name, recorder: HookTraceRecorder | None = None): async def _enter_hooks_section(): - """尝试获取全局信号量并更新计数器,超时则抛出 PermissionExemption。""" + """尝试获取全局信号量并更新计数器,饱和时快速放行。""" global HOOKS_ACTIVE_COUNT - try: - await asyncio.wait_for(HOOKS_SEMAPHORE.acquire(), timeout=TIMEOUT_SECONDS) - except asyncio.TimeoutError: + if HOOKS_SEMAPHORE.locked(): logger.warning( - "hooks semaphore acquire timeout, allowing pass", + "hooks semaphore saturated, allowing pass", LOGGER_COMMAND, ) - raise PermissionExemption("hooks semaphore timeout, allow pass") + raise PermissionExemption("hooks semaphore saturated, allow pass") + await HOOKS_SEMAPHORE.acquire() HOOKS_ACTIVE_COUNT += 1 @@ -1197,14 +1402,7 @@ async def _leave_hooks_section(): async def route_precheck( matcher: Matcher, - event: Event, - session: Uninfo, - message: UniMsg | None, - *, - entity=None, - event_cache: dict | None = None, - text: str | None = None, - route_modules: set[str] | None = None, + context: EventContext, ) -> bool: module = matcher.plugin_name or "" if not module: @@ -1213,19 +1411,20 @@ async def route_precheck( return False if not _is_command_matcher_class(type(matcher)): return False - if entity is None: - entity = get_entity_ids(session) - if event_cache is None: - event_cache = _get_event_cache(event, session, entity) - if text is None: - text = _get_message_text(message, event_cache, event) + + route_modules = context.route_modules if context.route_modules_loaded else None if route_modules is None: - route_modules = await _get_route_context(text, event_cache) + 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 event_cache is not None: - event_cache["route_skip"] = True + if context.event_cache is not None: + context.event_cache["route_skip"] = True return True return False @@ -1235,14 +1434,10 @@ async def auth( event: Event, bot: Bot, session: Uninfo, - message: UniMsg | None, *, + context: EventContext, skip_ban: bool = False, - entity=None, - event_cache: dict | None = None, - text: str | None = None, - route_modules: set[str] | None = None, - is_superuser: bool = False, + state: dict | None = None, ): """权限检查 @@ -1251,20 +1446,28 @@ async def auth( event: Event bot: bot session: Uninfo - message: UniMsg + context: EventContext """ start_time = time.time() cost_gold = 0 ignore_flag = False - if entity is None: - entity = get_entity_ids(session) + entity = context.entity + event_cache = context.event_cache + text = context.plain_text + is_superuser = context.is_superuser + route_modules = context.route_modules if context.route_modules_loaded else None module = matcher.plugin_name or "" is_command_matcher = _is_command_matcher_class(type(matcher)) - if event_cache is None: - event_cache = _get_event_cache(event, session, entity) auth_allowed = None auth_result_cache = None admin_checked_pre = False + permission_context: PermissionContext | None = None + side_effect_cache = get_permission_side_effect_cache( + state=state, + event_cache=event_cache, + ) + side_effect_lock = None + entered_side_effect_lock = False # 仅在慢请求时记录 hook 明细,避免热路径高频构造字符串 hook_recorder = HookTraceRecorder(start_time) @@ -1278,14 +1481,17 @@ async def auth( auth_allowed = True return - if event_cache is not None: - auth_result_cache = event_cache.setdefault("auth_result", {}) - cached_result = auth_result_cache.get(module) - if cached_result is not None: - allowed, reason = cached_result - if not allowed: - raise SkipPluginException(reason or "auth cached skip") - return + side_effect_lock = side_effect_cache.lock_for(module) + await side_effect_lock.acquire() + entered_side_effect_lock = True + + auth_result_cache = side_effect_cache.auth_results + cached_result = auth_result_cache.get(module) + if cached_result is not None: + allowed, reason = cached_result + if not allowed: + raise SkipPluginException(reason or "auth cached skip") + return if _is_hidden_plugin(matcher): auth_allowed = True @@ -1293,10 +1499,9 @@ async def auth( if event_cache is not None and event_cache.get("ban_state") is True: raise SkipPluginException("user or group banned (cached)") - if text is None: - text = _get_message_text(message, event_cache, event) if route_modules is None: route_modules = await _get_route_context(text, event_cache) + set_route_modules(state, context, route_modules) route_skip_checks = ( is_command_matcher and module in _ROUTE_MODULES_WITH_COMMANDS @@ -1310,7 +1515,7 @@ async def auth( auth_allowed = True return - platform = PlatformUtils.get_platform(session) + platform = context.platform # 获取插件和用户数据 plugin_user_start = time.time() try: @@ -1336,6 +1541,14 @@ async def auth( auth_allowed = True return + permission_context = PermissionContext( + event=context, + module=module, + plugin=plugin, + user=user, + ) + store_permission_context(state, permission_context) + if not route_skip_checks and _needs_admin_check(plugin): if plugin.plugin_type in { PluginType.SUPERUSER, @@ -1356,13 +1569,18 @@ async def auth( admin_timeout = False if event_cache is not None: admin_levels, admin_timeout = await _get_admin_levels_cached( - session, entity, event_cache + entity, event_cache ) + permission_context.admin_levels = admin_levels if admin_timeout: hook_recorder.set("auth_admin", "timeout") else: admin_start = time.time() - await auth_admin(plugin, session, cached_levels=admin_levels) + await auth_admin( + plugin, + session, + context=permission_context, + ) hook_recorder.set( "auth_admin", f"{time.time() - admin_start:.3f}s(pre)" ) @@ -1386,8 +1604,7 @@ async def auth( matcher, session, plugin, - entity=entity, - is_superuser=is_superuser, + context=permission_context, ) hook_recorder.set("auth_ban", f"{time.time() - ban_start:.3f}s") if event_cache is not None: @@ -1407,7 +1624,7 @@ async def auth( user, plugin, session, - is_superuser=is_superuser, + context=permission_context, ), name="get_plugin_cost", ) @@ -1421,7 +1638,7 @@ async def auth( hook_recorder.set("cost_gold", "skipped") # 执行 bot_filter - bot_filter(session) + bot_filter(session, context=permission_context) group = await _get_group_cached(entity, event_cache) @@ -1439,13 +1656,23 @@ async def auth( and not route_skip_checks ): admin_levels, admin_timeout = await _get_admin_levels_cached( - session, entity, event_cache + entity, event_cache ) + permission_context.group = group + permission_context.bot_data = bot_data + if admin_levels is not None: + permission_context.admin_levels = admin_levels + store_permission_context(state, permission_context) + # 并行执行所有 hook 检查,并记录执行时间 hooks_start = time.time() allow_sleep_bypass = _is_bot_wake_command(module, text) + # 先进入 hooks 并行检查区域;饱和时快速放行,避免创建并积压协程。 + await _enter_hooks_section() + entered_hooks = True + # 创建所有 hook 任务 hook_tasks = [] if event_cache is None: @@ -1455,6 +1682,7 @@ async def auth( plugin, bot.self_id, allow_sleep_bypass=allow_sleep_bypass, + context=permission_context, ), "auth_bot", hook_recorder, @@ -1472,6 +1700,7 @@ async def auth( bot_data=bot_data, skip_fetch=True, allow_sleep_bypass=allow_sleep_bypass, + context=permission_context, ), "auth_bot", hook_recorder, @@ -1483,7 +1712,13 @@ async def auth( else: hook_tasks.append( time_hook( - auth_group(plugin, group, text, entity.group_id), + auth_group( + plugin, + group, + text, + entity.group_id, + context=permission_context, + ), "auth_group", hook_recorder, ) @@ -1492,7 +1727,11 @@ async def auth( if not route_skip_checks and plugin.admin_level and not admin_checked_pre: if event_cache is None: hook_tasks.append( - time_hook(auth_admin(plugin, session), "auth_admin", hook_recorder) + time_hook( + auth_admin(plugin, session, context=permission_context), + "auth_admin", + hook_recorder, + ) ) else: if admin_timeout: @@ -1500,7 +1739,11 @@ async def auth( else: hook_tasks.append( time_hook( - auth_admin(plugin, session, cached_levels=admin_levels), + auth_admin( + plugin, + session, + context=permission_context, + ), "auth_admin", hook_recorder, ) @@ -1510,7 +1753,7 @@ async def auth( if is_superuser: hook_recorder.set("auth_plugin", "superuser") - elif not route_skip_checks and _needs_auth_plugin(plugin, group, entity): + elif not route_skip_checks and _needs_auth_plugin(plugin, permission_context): hook_tasks.append( time_hook( auth_plugin( @@ -1518,6 +1761,7 @@ async def auth( group, session, event, + context=permission_context, skip_group_block=is_superuser, ), "auth_plugin", @@ -1531,18 +1775,17 @@ async def auth( has_limits = await _has_limits_cached(module, event_cache) if has_limits: hook_tasks.append( - time_hook(auth_limit(plugin, session), "auth_limit", hook_recorder) + time_hook( + auth_limit(plugin, session, context=permission_context), + "auth_limit", + hook_recorder, + ) ) else: hook_recorder.set("auth_limit", "skipped") else: hook_recorder.set("auth_limit", "skipped") - if hook_tasks: - # 进入 hooks 并行检查区域(会在高并发时排队) - await _enter_hooks_section() - entered_hooks = True - # 使用 gather 并行执行所有 hook,但添加总体超时控制 try: await with_timeout( @@ -1599,6 +1842,9 @@ async def auth( ) if auth_result_cache is not None and auth_allowed is not None: auth_result_cache[module] = (auth_allowed, None) + if entered_side_effect_lock and side_effect_lock is not None: + with contextlib.suppress(Exception): + side_effect_lock.release() # 扣除金币 if not ignore_flag and cost_gold > 0: gold_start = time.time() diff --git a/zhenxun/builtin_plugins/hooks/auth_hook.py b/zhenxun/builtin_plugins/hooks/auth_hook.py index 8a9b1337..94c08da4 100644 --- a/zhenxun/builtin_plugins/hooks/auth_hook.py +++ b/zhenxun/builtin_plugins/hooks/auth_hook.py @@ -1,4 +1,3 @@ -import contextlib import time from nonebot import get_driver @@ -12,14 +11,20 @@ from nonebot_plugin_uninfo import Uninfo from zhenxun.services.cache.runtime_cache import is_cache_ready from zhenxun.services.log import logger -from zhenxun.services.message_load import is_overloaded +from zhenxun.services.message_load import is_overloaded, mark_activity from zhenxun.services.runtime_bootstrap import register_runtime_bootstrap -from zhenxun.utils.utils import get_entity_ids from .auth.config import LOGGER_COMMAND +from .auth.context import ( + get_event_context, + get_or_create_event_context, + resolve_actor_user_id, + resolve_event_channel_id, + resolve_event_group_id, + set_route_modules, +) from .auth_checker import ( LimitManager, - _get_event_cache, _get_route_context, auth, route_precheck, @@ -41,17 +46,6 @@ async def _mark_bot_connected(bot: Bot): _BOT_CONNECT_TS = time.time() -def _extract_plain_text(message: UniMsg | None, event: Event) -> str: - if message is not None: - with contextlib.suppress(Exception): - return message.extract_plain_text() - with contextlib.suppress(Exception): - plain = event.get_plaintext() - if plain: - return plain.strip() - return "" - - @driver.on_startup async def _start_auth_runtime_tasks(): await start_auth_runtime_tasks() @@ -72,37 +66,9 @@ def _skip_auth_for_plugin(matcher: Matcher) -> bool: return "chat_history" in module_name -def _resolve_actor_user_id(event: Event, fallback_user_id: str) -> str: - """优先使用事件发起者ID,避免 notice 场景 session.user 指向 bot 自身。""" - event_user_id = getattr(event, "user_id", None) - if event_user_id is None: - return fallback_user_id - event_user_id = str(event_user_id) - return event_user_id or fallback_user_id - - -def _resolve_event_group_id(event: Event, fallback_group_id: str | None) -> str | None: - """notice 场景 session.group 可能缺失,回退到事件上的 group_id。""" - event_group_id = getattr(event, "group_id", None) - if event_group_id is None: - return fallback_group_id - resolved = str(event_group_id) - return resolved or fallback_group_id - - -def _resolve_event_channel_id( - event: Event, fallback_channel_id: str | None -) -> str | None: - """频道场景回退到事件上的 channel_id。""" - event_channel_id = getattr(event, "channel_id", None) - if event_channel_id is None: - return fallback_channel_id - resolved = str(event_channel_id) - return resolved or fallback_channel_id - - @event_preprocessor async def _drop_message_before_cache_ready(event: Event): + mark_activity() if event.get_type() != "message": return if not is_cache_ready(): @@ -130,46 +96,22 @@ async def _auth_preprocessor( return start_time = time.time() - entity = state.get("_zx_entity") - if entity is None: - entity = get_entity_ids(session) - entity.user_id = _resolve_actor_user_id(event, entity.user_id) - entity.group_id = _resolve_event_group_id(event, entity.group_id) - entity.channel_id = _resolve_event_channel_id(event, entity.channel_id) - state["_zx_entity"] = entity - - event_cache = state.get("_zx_event_cache") - if event_cache is None: - event_cache = _get_event_cache(event, session, entity) - state["_zx_event_cache"] = event_cache - - text = state.get("_zx_plain_text") - if text is None: - text = _extract_plain_text(message, event) - state["_zx_plain_text"] = text - if event_cache is not None: - event_cache["plain_text"] = text - - route_modules = state.get("_zx_route_modules") - if route_modules is None: - route_modules = await _get_route_context(text, event_cache) - state["_zx_route_modules"] = route_modules - - is_superuser = state.get("_zx_is_superuser") - if is_superuser is None: - is_superuser = entity.user_id in bot.config.superusers - state["_zx_is_superuser"] = is_superuser - - if await route_precheck( - matcher, + event_context = get_or_create_event_context( + bot, event, session, - message, - entity=entity, - event_cache=event_cache, - text=text, - route_modules=route_modules, - ): + state, + message=message, + ) + + if not event_context.route_modules_loaded: + route_modules = await _get_route_context( + event_context.plain_text, + event_context.event_cache, + ) + set_route_modules(state, event_context, route_modules) + + if await route_precheck(matcher, event_context): return try: @@ -178,13 +120,9 @@ async def _auth_preprocessor( event, bot, session, - message, + context=event_context, skip_ban=False, - entity=entity, - event_cache=event_cache, - text=text, - route_modules=route_modules, - is_superuser=is_superuser, + state=state, ) except IgnoredException: raise @@ -203,16 +141,27 @@ async def _auth_preprocessor( @run_postprocessor -async def _unblock_after_matcher(matcher: Matcher, session: Uninfo, event: Event): - user_id = _resolve_actor_user_id(event, session.user.id) - group_id = _resolve_event_group_id(event, None) - channel_id = _resolve_event_channel_id(event, None) - if session.group: - if session.group.parent: - group_id = session.group.parent.id - channel_id = session.group.id - else: - group_id = session.group.id +async def _unblock_after_matcher( + matcher: Matcher, + session: Uninfo, + event: Event, + state: T_State, +): + context = get_event_context(state) + if context is not None: + user_id = context.user_id + group_id = context.group_id + channel_id = context.channel_id + else: + user_id = resolve_actor_user_id(event, session.user.id) + group_id = resolve_event_group_id(event, None) + channel_id = resolve_event_channel_id(event, None) + if session.group: + if session.group.parent: + group_id = session.group.parent.id + channel_id = session.group.id + else: + group_id = session.group.id if user_id and matcher.plugin: module = matcher.plugin.name LimitManager.unblock(module, user_id, group_id, channel_id) diff --git a/zhenxun/builtin_plugins/hooks/chkdsk_hook.py b/zhenxun/builtin_plugins/hooks/chkdsk_hook.py index acfe07ee..c2f0ab43 100644 --- a/zhenxun/builtin_plugins/hooks/chkdsk_hook.py +++ b/zhenxun/builtin_plugins/hooks/chkdsk_hook.py @@ -16,13 +16,7 @@ from zhenxun.services.log import logger from zhenxun.utils.enum import PluginType from zhenxun.utils.message import MessageUtils -malicious_check_time = Config.get_config("hook", "MALICIOUS_CHECK_TIME") -malicious_ban_count = Config.get_config("hook", "MALICIOUS_BAN_COUNT") - -if not malicious_check_time: - raise ValueError("模块: [hook], 配置项: [MALICIOUS_CHECK_TIME] 为空或小于0") -if not malicious_ban_count: - raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_COUNT] 为空或小于0") +from .auth.context import resolve_actor_user_id, resolve_event_group_id class BanCheckLimiter: @@ -36,6 +30,10 @@ class BanCheckLimiter: self.default_check_time = default_check_time self.default_count = default_count + def configure(self, check_time: float, count: int) -> None: + self.default_check_time = check_time + self.default_count = count + def add(self, key: str | float): if self.mint[key] == 1: self.mtime[key] = time.time() @@ -59,11 +57,22 @@ class BanCheckLimiter: _blmt = BanCheckLimiter( - malicious_check_time, - malicious_ban_count, + 5, + 4, ) +def _get_positive_config(key: str, cast_type: type[int] | type[float]) -> int | float: + value = Config.get_config("hook", key) + try: + parsed_value = cast_type(value) + except (TypeError, ValueError) as e: + raise ValueError(f"模块: [hook], 配置项: [{key}] 不是有效数字") from e + if parsed_value <= 0: + raise ValueError(f"模块: [hook], 配置项: [{key}] 为空或小于0") + return parsed_value + + # 恶意触发命令检测 @run_preprocessor async def _( @@ -88,11 +97,12 @@ async def _( else: return - user_id = session.id1 - group_id = session.id3 or session.id2 - malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME") - if not malicious_ban_time: - raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_TIME] 为空或小于0") + user_id = resolve_actor_user_id(event, session.id1) + group_id = resolve_event_group_id(event, session.id3 or session.id2) + malicious_check_time = float(_get_positive_config("MALICIOUS_CHECK_TIME", float)) + malicious_ban_count = int(_get_positive_config("MALICIOUS_BAN_COUNT", int)) + malicious_ban_time = int(_get_positive_config("MALICIOUS_BAN_TIME", int)) + _blmt.configure(malicious_check_time, malicious_ban_count) if user_id and module: if _blmt.check(f"{user_id}__{module}"): await BanConsole.ban( diff --git a/zhenxun/builtin_plugins/init/__init_cache.py b/zhenxun/builtin_plugins/init/__init_cache.py index 5608d00e..8938c0a0 100644 --- a/zhenxun/builtin_plugins/init/__init_cache.py +++ b/zhenxun/builtin_plugins/init/__init_cache.py @@ -29,7 +29,6 @@ def register_cache_types(): GroupPluginSetting, key_format="{group_id}_{plugin_name}_{key}", ) - CacheRegistry.register(CacheType.GROUP_PLUGIN_SETTINGS_VIEW, dict) CacheRegistry.register( CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}" ) diff --git a/zhenxun/builtin_plugins/init/init_plugin.py b/zhenxun/builtin_plugins/init/init_plugin.py index 4b912dfd..8791e598 100644 --- a/zhenxun/builtin_plugins/init/init_plugin.py +++ b/zhenxun/builtin_plugins/init/init_plugin.py @@ -88,7 +88,7 @@ async def _handle_setting( ) -@PriorityLifecycle.on_startup(priority=5) +@PriorityLifecycle.on_startup(priority=4) async def _(): """ 初始化插件数据配置 diff --git a/zhenxun/builtin_plugins/plugin_store/data_source.py b/zhenxun/builtin_plugins/plugin_store/data_source.py index 999d2fc0..27f05e4c 100644 --- a/zhenxun/builtin_plugins/plugin_store/data_source.py +++ b/zhenxun/builtin_plugins/plugin_store/data_source.py @@ -3,12 +3,12 @@ from pathlib import Path import random import shutil -from aiocache import cached import ujson as json from zhenxun.builtin_plugins.plugin_store.models import StorePluginInfo from zhenxun.configs.path_config import TEMP_PATH from zhenxun.models.plugin_info import PluginInfo +from zhenxun.services.cache.bounded_ttl import BoundedTTLCache from zhenxun.services.log import logger from zhenxun.services.plugin_init import PluginInitManager from zhenxun.utils.enum import PluginType @@ -26,6 +26,14 @@ from .config import ( ) from .exceptions import PluginStoreException +_PLUGIN_STORE_DATA_CACHE = BoundedTTLCache[ + str, tuple[list[StorePluginInfo], list[StorePluginInfo]] +]( + "PLUGIN_STORE_DATA", + ttl_seconds=60, + max_items=1, +) + def row_style(column: str, text: str) -> RowStyle: """被动技能文本风格 @@ -56,20 +64,9 @@ class StoreManager: relative_parts = [part for part in plugin_info.module_path.split(".") if part] relative_path = Path(*relative_parts) if relative_parts else Path(plugin_name) path = BASE_PATH.parent / relative_path - if plugin_info.is_dir: - return path - return path.parent / f"{plugin_name}.py" + return path if plugin_info.is_dir else path.parent / f"{plugin_name}.py" @classmethod - def _is_plugin_installed( - cls, plugin_info: StorePluginInfo, *, is_external: bool - ) -> bool: - return cls._resolve_local_plugin_path( - plugin_info, is_external=is_external - ).exists() - - @classmethod - @cached(60) async def get_data(cls) -> tuple[list[StorePluginInfo], list[StorePluginInfo]]: """获取插件信息数据 @@ -77,15 +74,22 @@ class StoreManager: tuple[list[StorePluginInfo], list[StorePluginInfo]]: 原生插件信息数据,第三方插件信息数据 """ + cache_key = "plugins_json" + if cached_data := await _PLUGIN_STORE_DATA_CACHE.get(cache_key): + return cached_data + plugins = await RepoFileManager.get_file_content( DEFAULT_GITHUB_URL, "plugins.json" ) extra_plugins = await RepoFileManager.get_file_content( EXTRA_GITHUB_URL, "plugins.json", "index" ) - return [StorePluginInfo(**plugin) for plugin in json.loads(plugins)], [ - StorePluginInfo(**plugin) for plugin in json.loads(extra_plugins) - ] + result = ( + [StorePluginInfo(**plugin) for plugin in json.loads(plugins)], + [StorePluginInfo(**plugin) for plugin in json.loads(extra_plugins)], + ) + await _PLUGIN_STORE_DATA_CACHE.set(cache_key, result) + return result @classmethod def version_check(cls, plugin_info: StorePluginInfo, suc_plugin: dict[str, str]): @@ -330,13 +334,12 @@ class StoreManager: source: 源 """ repo_type = RepoType.GITHUB if is_external else None - if source == "ali": + if ( + source != "ali" and source != "git" and plugin_info.ali_url + ) or source == "ali": repo_type = RepoType.ALIYUN elif source == "git": repo_type = RepoType.GITHUB - else: - if plugin_info.ali_url: - repo_type = RepoType.ALIYUN module_path = plugin_info.module_path is_dir = plugin_info.is_dir github_url = plugin_info.github_url @@ -380,7 +383,7 @@ class StoreManager: requirement_file = target_dir / requirement_path.path if requirement_file.exists(): is_install_req = True - await VirtualEnvPackageManager.install_requirement(requirement_file) + await VirtualEnvPackageManager.add_requirement(requirement_file) if not is_install_req: # 从仓库根目录查找文件 @@ -401,13 +404,13 @@ class StoreManager: f"开始安装插件 {module_path} 依赖文件: {requirement_path}", LOG_COMMAND, ) - await VirtualEnvPackageManager.install_requirement(requirement_path) + await VirtualEnvPackageManager.add_requirement(requirement_path) if requirements_path.exists(): logger.info( f"开始安装插件 {module_path} 依赖文件: {requirements_path}", LOG_COMMAND, ) - await VirtualEnvPackageManager.install_requirement(requirements_path) + await VirtualEnvPackageManager.add_requirement(requirements_path) @classmethod async def remove_plugin(cls, index_or_module: str) -> str: diff --git a/zhenxun/builtin_plugins/shop/_data_source.py b/zhenxun/builtin_plugins/shop/_data_source.py index d4716e9a..f8076c3f 100644 --- a/zhenxun/builtin_plugins/shop/_data_source.py +++ b/zhenxun/builtin_plugins/shop/_data_source.py @@ -19,9 +19,9 @@ from zhenxun.models.friend_user import FriendUser from zhenxun.models.goods_info import GoodsInfo from zhenxun.models.group_member_info import GroupInfoUser from zhenxun.models.user_console import UserConsole -from zhenxun.models.user_gold_log import UserGoldLog from zhenxun.models.user_props_log import UserPropsLog from zhenxun.services import avatar_service +from zhenxun.services.buffered_writers import append_user_gold_log from zhenxun.services.log import logger from zhenxun.ui.models import ImageCell, TextCell from zhenxun.utils.enum import GoldHandle, PropHandle @@ -480,10 +480,6 @@ class ShopManage: ).count() if goods.daily_limit and count >= goods.daily_limit: return "今天的购买已达限制了喔!" - await UserGoldLog.create(user_id=user_id, gold=price, handle=GoldHandle.BUY) - await UserPropsLog.create( - user_id=user_id, uuid=goods.uuid, gold=price, num=num, handle=PropHandle.BUY - ) logger.info( f"花费 {price} 金币购买 {goods.goods_name} ×{num} 成功!", "购买道具", @@ -494,6 +490,12 @@ class ShopManage: user.props[goods.uuid] = 0 user.props[goods.uuid] += num await user.save(update_fields=["gold", "props"]) + await append_user_gold_log( + user_id=user_id, gold=int(price), handle=GoldHandle.BUY + ) + await UserPropsLog.create( + user_id=user_id, uuid=goods.uuid, gold=price, num=num, handle=PropHandle.BUY + ) return f"花费 {price} 金币购买 {goods.goods_name} ×{num} 成功!" @classmethod diff --git a/zhenxun/builtin_plugins/sign_in/_data_source.py b/zhenxun/builtin_plugins/sign_in/_data_source.py index 9cec618b..4dbe3ff0 100644 --- a/zhenxun/builtin_plugins/sign_in/_data_source.py +++ b/zhenxun/builtin_plugins/sign_in/_data_source.py @@ -9,13 +9,17 @@ import pytz from zhenxun import ui from zhenxun.configs.path_config import IMAGE_PATH from zhenxun.models.friend_user import FriendUser +from zhenxun.models.goods_info import GoodsInfo from zhenxun.models.group_member_info import GroupInfoUser from zhenxun.models.sign_log import SignLog from zhenxun.models.sign_user import SignUser from zhenxun.models.user_console import UserConsole from zhenxun.services.avatar_service import avatar_service +from zhenxun.services.buffered_writers import append_user_gold_log from zhenxun.services.log import logger from zhenxun.ui.models import ImageCell, TextCell +from zhenxun.utils.enum import GoldHandle +from zhenxun.utils.exception import GoodsNotFound from zhenxun.utils.platform import PlatformUtils from ._random_event import random_event @@ -182,11 +186,32 @@ class SignManage: gift = random_event(float(user.impression)) if isinstance(gift, int): gold += gift - await UserConsole.add_gold(user.user_id, gold + gift, "sign_in", platform) + user_console = await UserConsole.get_user(user.user_id, platform) + user_console.gold += gold + await user_console.save(update_fields=["gold"]) + await append_user_gold_log( + user_id=user.user_id, + gold=gold, + handle=GoldHandle.GET, + source="sign_in", + ) gift = f"额外金币 +{gift}" else: - await UserConsole.add_gold(user.user_id, gold, "sign_in", platform) - await UserConsole.add_props_by_name(user.user_id, gift, 1, platform) + goods = await GoodsInfo.get_or_none(goods_name=gift) + if not goods: + raise GoodsNotFound("未找到商品...") + user_console = await UserConsole.get_user(user.user_id, platform) + user_console.gold += gold + if goods.uuid not in user_console.props: + user_console.props[goods.uuid] = 0 + user_console.props[goods.uuid] += 1 + await user_console.save(update_fields=["gold", "props"]) + await append_user_gold_log( + user_id=user.user_id, + gold=gold, + handle=GoldHandle.GET, + source="sign_in", + ) gift += " + 1" logger.info( f"签到成功. score: {user.impression:.2f} " diff --git a/zhenxun/builtin_plugins/statistics/statistics_hook.py b/zhenxun/builtin_plugins/statistics/statistics_hook.py index d2df77f6..64e9e3d7 100644 --- a/zhenxun/builtin_plugins/statistics/statistics_hook.py +++ b/zhenxun/builtin_plugins/statistics/statistics_hook.py @@ -1,5 +1,7 @@ +import asyncio from datetime import datetime +from nonebot import get_driver from nonebot.adapters import Bot, Event from nonebot.adapters.onebot.v11 import PokeNotifyEvent from nonebot.matcher import Matcher @@ -25,7 +27,35 @@ __plugin_meta__ = PluginMetadata( ).to_dict(), ) -TEMP_LIST = [] +STATS_BUFFER_FLUSH_SIZE = 5000 +STATS_BUFFER_MAX_RETAIN = 10000 +TEMP_LIST: list[Statistics] = [] +_STATS_FLUSH_LOCK = asyncio.Lock() +driver = get_driver() + + +async def _flush_statistics_buffer(reason: str) -> int: + async with _STATS_FLUSH_LOCK: + call_list = TEMP_LIST.copy() + TEMP_LIST.clear() + if not call_list: + return 0 + try: + await Statistics.bulk_create(call_list) + except Exception as e: + logger.error(f"{reason}批量添加调用记录失败", "定时任务", e=e) + retain_count = max(STATS_BUFFER_MAX_RETAIN - len(TEMP_LIST), 0) + if retain_count: + TEMP_LIST[:0] = call_list[-retain_count:] + return 0 + logger.debug(f"{reason}批量添加调用记录 {len(call_list)} 条", "定时任务") + return len(call_list) + + +async def _append_statistics(record: Statistics) -> None: + TEMP_LIST.append(record) + if len(TEMP_LIST) >= STATS_BUFFER_FLUSH_SIZE and not _STATS_FLUSH_LOCK.locked(): + await _flush_statistics_buffer("缓冲区触发") @run_postprocessor @@ -50,7 +80,7 @@ async def _( if plugin_type == PluginType.NORMAL: entity = get_entity_ids(session) logger.debug(f"提交调用记录: {matcher.plugin_name}...", session=session) - TEMP_LIST.append( + await _append_statistics( Statistics( user_id=entity.user_id, group_id=entity.group_id, @@ -66,10 +96,11 @@ async def _(): try: if should_pause_tasks(): return - call_list = TEMP_LIST.copy() - TEMP_LIST.clear() - if call_list: - await Statistics.bulk_create(call_list) - logger.debug(f"批量添加调用记录 {len(call_list)} 条", "定时任务") + await _flush_statistics_buffer("定时") except Exception as e: logger.error("定时批量添加调用记录", "定时任务", e=e) + + +@driver.on_shutdown +async def _flush_statistics_on_shutdown(): + await _flush_statistics_buffer("关闭") diff --git a/zhenxun/builtin_plugins/superuser/reload_setting.py b/zhenxun/builtin_plugins/superuser/reload_setting.py index 2bf821cb..6d37d903 100644 --- a/zhenxun/builtin_plugins/superuser/reload_setting.py +++ b/zhenxun/builtin_plugins/superuser/reload_setting.py @@ -1,3 +1,5 @@ +import contextlib + from nonebot.permission import SUPERUSER from nonebot.plugin import PluginMetadata from nonebot.rule import to_me @@ -11,8 +13,11 @@ from zhenxun.services.llm.config.providers import get_llm_config from zhenxun.services.llm.manager import clear_model_cache from zhenxun.services.log import logger from zhenxun.utils.enum import PluginType +from zhenxun.utils.manager.priority_manager import PriorityLifecycle from zhenxun.utils.message import MessageUtils +AUTO_RELOAD_JOB_ID = "zhenxun.reload_setting.auto_reload" + __plugin_meta__ = PluginMetadata( name="重载配置", description="重新加载config.yaml", @@ -53,22 +58,75 @@ _matcher = on_alconna( ) -@_matcher.handle() -async def _(session: EventSession, arparma: Arparma): +def _get_auto_reload_interval() -> int: + value = Config.get_config("reload_setting", "AUTO_RELOAD_TIME", 180) + try: + seconds = int(value) + except (TypeError, ValueError): + logger.warning( + f"AUTO_RELOAD_TIME 配置无效: {value!r},已使用默认值 180 秒", + "重载配置", + ) + return 180 + if seconds <= 0: + logger.warning( + f"AUTO_RELOAD_TIME 配置小于等于 0: {seconds},已使用默认值 180 秒", + "重载配置", + ) + return 180 + return seconds + + +def _reschedule_auto_reload_job() -> None: + seconds = _get_auto_reload_interval() + if scheduler.get_job(AUTO_RELOAD_JOB_ID): + scheduler.reschedule_job( + AUTO_RELOAD_JOB_ID, + trigger="interval", + seconds=seconds, + ) + else: + scheduler.add_job( + _auto_reload_config, + "interval", + seconds=seconds, + id=AUTO_RELOAD_JOB_ID, + replace_existing=True, + ) + logger.debug(f"自动重载配置任务间隔已设置为 {seconds} 秒", "重载配置") + + +async def _reload_plugin_limit_config() -> None: + from zhenxun.builtin_plugins.hooks.auth.auth_limit import LimitManager + from zhenxun.builtin_plugins.init.manager import manager + + manager.init() + await manager.load_to_db() + await LimitManager.update_limits() + + +async def _reload_runtime_config() -> None: Config.reload() get_llm_config.cache_clear() clear_model_cache() + await _reload_plugin_limit_config() + with contextlib.suppress(Exception): + _reschedule_auto_reload_job() + + +@PriorityLifecycle.on_startup(priority=1) +def _init_auto_reload_job() -> None: + _reschedule_auto_reload_job() + + +@_matcher.handle() +async def _(session: EventSession, arparma: Arparma): + await _reload_runtime_config() logger.debug("自动重载配置文件", arparma.header_result, session=session) await MessageUtils.build_message("重载完成!").send(reply_to=True) -@scheduler.scheduled_job( - "interval", - seconds=Config.get_config("reload_setting", "AUTO_RELOAD_TIME", 180), -) -async def _(): +async def _auto_reload_config() -> None: if Config.get_config("reload_setting", "AUTO_RELOAD"): - Config.reload() - get_llm_config.cache_clear() - clear_model_cache() + await _reload_runtime_config() logger.debug("已自动重载配置文件...") diff --git a/zhenxun/builtin_plugins/web_ui/__init__.py b/zhenxun/builtin_plugins/web_ui/__init__.py index 619d56bf..61286953 100644 --- a/zhenxun/builtin_plugins/web_ui/__init__.py +++ b/zhenxun/builtin_plugins/web_ui/__init__.py @@ -1,20 +1,17 @@ -import asyncio import secrets from fastapi import APIRouter, FastAPI import nonebot -from nonebot.log import default_filter, default_format from nonebot.plugin import PluginMetadata from zhenxun.configs.config import Config as gConfig from zhenxun.configs.utils import PluginExtraData, RegisterConfig -from zhenxun.services.log import logger, logger_ +from zhenxun.services.log import logger from zhenxun.utils.enum import PluginType from zhenxun.utils.manager.priority_manager import PriorityLifecycle from .api.configure import router as configure_router from .api.logs import router as ws_log_routes -from .api.logs.log_manager import LOG_STORAGE from .api.menu import router as menu_router from .api.tabs.dashboard import router as dashboard_router from .api.tabs.database import router as database_router @@ -95,25 +92,6 @@ WsApiRouter.include_router(chat_routes) @PriorityLifecycle.on_startup(priority=0) async def _(): try: - # 存储任务引用的列表,防止任务被垃圾回收 - _tasks = [] - - async def log_sink(message: str): - loop = None - if not loop: - try: - loop = asyncio.get_running_loop() - except Exception as e: - logger.warning("Web Ui log_sink", e=e) - if not loop: - loop = asyncio.new_event_loop() - # 存储任务引用到外部列表中 - _tasks.append(loop.create_task(LOG_STORAGE.add(message.rstrip("\n")))) - - logger_.add( - log_sink, colorize=True, filter=default_filter, format=default_format - ) - app: FastAPI = nonebot.get_app() app.include_router(BaseApiRouter) app.include_router(WsApiRouter) diff --git a/zhenxun/builtin_plugins/web_ui/api/logs/log_manager.py b/zhenxun/builtin_plugins/web_ui/api/logs/log_manager.py index 3938c525..4de877c8 100644 --- a/zhenxun/builtin_plugins/web_ui/api/logs/log_manager.py +++ b/zhenxun/builtin_plugins/web_ui/api/logs/log_manager.py @@ -1,33 +1,97 @@ import asyncio +from collections import deque from collections.abc import Awaitable, Callable -from typing import Generic, TypeVar +import contextlib -_T = TypeVar("_T") -LogListener = Callable[[_T], Awaitable[None]] +from nonebot.log import default_filter, default_format + +from zhenxun.services.log import logger_ + +LogListener = Callable[[str], Awaitable[None]] +DEFAULT_MAX_LOGS = 1000 +DEFAULT_MAX_LISTENERS = 16 -class LogStorage(Generic[_T]): +class LogStorage: """ 日志存储 """ - def __init__(self, rotation: float = 5 * 60): + def __init__( + self, + rotation: float = 5 * 60, + max_logs: int = DEFAULT_MAX_LOGS, + max_listeners: int = DEFAULT_MAX_LISTENERS, + ): self.count, self.rotation = 0, rotation + self.max_logs = max_logs + self.max_listeners = max_listeners self.logs: dict[int, str] = {} - self.listeners: set[LogListener[str]] = set() + self._order: deque[int] = deque() + self.listeners: set[LogListener] = set() async def add(self, log: str): seq = self.count = self.count + 1 self.logs[seq] = log + self._order.append(seq) + self._trim() asyncio.get_running_loop().call_later(self.rotation, self.remove, seq) - await asyncio.gather( - *(listener(log) for listener in self.listeners), - return_exceptions=True, - ) + listeners = tuple(self.listeners) + if listeners: + results = await asyncio.gather( + *(listener(log) for listener in listeners), + return_exceptions=True, + ) + for listener, result in zip(listeners, results, strict=False): + if isinstance(result, BaseException): + self.listeners.discard(listener) return seq + def add_listener(self, listener: LogListener) -> bool: + if len(self.listeners) >= self.max_listeners: + return False + self.listeners.add(listener) + return True + + def remove_listener(self, listener: LogListener) -> None: + self.listeners.discard(listener) + def remove(self, seq: int): - del self.logs[seq] + self.logs.pop(seq, None) + with contextlib.suppress(ValueError): + self._order.remove(seq) + + def _trim(self) -> None: + while self._order and self._order[0] not in self.logs: + self._order.popleft() + while len(self.logs) > self.max_logs and self._order: + self.logs.pop(self._order.popleft(), None) -LOG_STORAGE: LogStorage[str] = LogStorage[str]() +LOG_STORAGE = LogStorage() + +_LOG_SINK_ID: int | None = None + + +async def ensure_log_sink_started() -> None: + global _LOG_SINK_ID + if _LOG_SINK_ID is not None: + return + + async def log_sink(message: str) -> None: + await LOG_STORAGE.add(message.rstrip("\n")) + + _LOG_SINK_ID = logger_.add( + log_sink, + colorize=True, + filter=default_filter, + format=default_format, + ) + + +def stop_log_sink_if_idle() -> None: + global _LOG_SINK_ID + if LOG_STORAGE.listeners or _LOG_SINK_ID is None: + return + logger_.remove(_LOG_SINK_ID) + _LOG_SINK_ID = None diff --git a/zhenxun/builtin_plugins/web_ui/api/logs/logs.py b/zhenxun/builtin_plugins/web_ui/api/logs/logs.py index b7fc660c..13015e43 100644 --- a/zhenxun/builtin_plugins/web_ui/api/logs/logs.py +++ b/zhenxun/builtin_plugins/web_ui/api/logs/logs.py @@ -3,7 +3,7 @@ from loguru import logger from nonebot.utils import escape_tag from starlette.websockets import WebSocket, WebSocketDisconnect, WebSocketState -from .log_manager import LOG_STORAGE +from .log_manager import LOG_STORAGE, ensure_log_sink_started, stop_log_sink_if_idle router = APIRouter() @@ -11,11 +11,16 @@ router = APIRouter() @router.websocket("/logs") async def system_logs_realtime(websocket: WebSocket): await websocket.accept() + await ensure_log_sink_started() async def log_listener(log: str): await websocket.send_text(log) - LOG_STORAGE.listeners.add(log_listener) + if not LOG_STORAGE.add_listener(log_listener): + await websocket.send_text("日志连接数已达上限,请稍后再试。") + await websocket.close() + stop_log_sink_if_idle() + return try: while websocket.client_state == WebSocketState.CONNECTED: recv = await websocket.receive() @@ -26,4 +31,5 @@ async def system_logs_realtime(websocket: WebSocket): except WebSocketDisconnect: pass finally: - LOG_STORAGE.listeners.remove(log_listener) + LOG_STORAGE.remove_listener(log_listener) + stop_log_sink_if_idle() diff --git a/zhenxun/builtin_plugins/web_ui/api/tabs/main/__init__.py b/zhenxun/builtin_plugins/web_ui/api/tabs/main/__init__.py index e1963a5c..9756a470 100644 --- a/zhenxun/builtin_plugins/web_ui/api/tabs/main/__init__.py +++ b/zhenxun/builtin_plugins/web_ui/api/tabs/main/__init__.py @@ -34,6 +34,31 @@ run_time = time.time() ws_router = APIRouter() router = APIRouter(prefix="/main") +_SYSTEM_STATUS_CONNECTIONS: set[WebSocket] = set() +_SYSTEM_STATUS_STOPPING = False + + +async def _close_system_status_websocket(websocket: WebSocket) -> None: + with contextlib.suppress(Exception): + if websocket.client_state == WebSocketState.CONNECTED: + await asyncio.wait_for( + websocket.close(code=1001, reason="server shutdown"), + timeout=2, + ) + + +@driver.on_shutdown +async def _close_system_status_websockets() -> None: + global _SYSTEM_STATUS_STOPPING + _SYSTEM_STATUS_STOPPING = True + websockets = list(_SYSTEM_STATUS_CONNECTIONS) + if not websockets: + return + await asyncio.gather( + *(_close_system_status_websocket(websocket) for websocket in websockets), + return_exceptions=True, + ) + _SYSTEM_STATUS_CONNECTIONS.clear() @router.get( @@ -243,11 +268,39 @@ async def _(param: BotManageUpdateParam): @ws_router.websocket("/system_status") async def system_logs_realtime(websocket: WebSocket, sleep: int = 5): await websocket.accept() + _SYSTEM_STATUS_CONNECTIONS.add(websocket) logger.debug("ws system_status is connect") - with contextlib.suppress( - WebSocketDisconnect, ConnectionClosedError, ConnectionClosedOK - ): - while websocket.client_state == WebSocketState.CONNECTED: + + disconnect_event = asyncio.Event() + + async def _watch_disconnect() -> None: + try: + while websocket.client_state == WebSocketState.CONNECTED: + await websocket.receive() + except (WebSocketDisconnect, ConnectionClosedError, ConnectionClosedOK): + pass + except Exception as e: + logger.debug(f"ws system_status receive stopped: {type(e).__name__}") + finally: + disconnect_event.set() + + receive_task = asyncio.create_task(_watch_disconnect()) + try: + while ( + websocket.client_state == WebSocketState.CONNECTED + and not _SYSTEM_STATUS_STOPPING + ): system_status = await get_system_status() - await websocket.send_text(system_status.json()) - await asyncio.sleep(sleep) + await asyncio.wait_for(websocket.send_text(system_status.json()), timeout=5) + try: + await asyncio.wait_for(disconnect_event.wait(), timeout=max(sleep, 1)) + except TimeoutError: + pass + except (WebSocketDisconnect, ConnectionClosedError, ConnectionClosedOK): + pass + finally: + _SYSTEM_STATUS_CONNECTIONS.discard(websocket) + receive_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await receive_task + await _close_system_status_websocket(websocket) diff --git a/zhenxun/builtin_plugins/web_ui/api/tabs/plugin_manage/__init__.py b/zhenxun/builtin_plugins/web_ui/api/tabs/plugin_manage/__init__.py index 2fa21143..ba578f71 100644 --- a/zhenxun/builtin_plugins/web_ui/api/tabs/plugin_manage/__init__.py +++ b/zhenxun/builtin_plugins/web_ui/api/tabs/plugin_manage/__init__.py @@ -52,34 +52,20 @@ async def _( async def _() -> Result[PluginCount]: try: plugin_count = PluginCount() - plugin_count.normal = len( - await DbPluginInfo.get_plugins( - plugin_type=PluginType.NORMAL, - load_status=True, - filter_parent=False, - ) - ) - plugin_count.admin = len( - await DbPluginInfo.get_plugins( - plugin_type__in=[PluginType.ADMIN, PluginType.SUPER_AND_ADMIN], - load_status=True, - filter_parent=False, - ) - ) - plugin_count.superuser = len( - await DbPluginInfo.get_plugins( - plugin_type__in=[PluginType.SUPERUSER, PluginType.SUPER_AND_ADMIN], - load_status=True, - filter_parent=False, - ) - ) - plugin_count.other = len( - await DbPluginInfo.get_plugins( - plugin_type__in=[PluginType.HIDDEN, PluginType.DEPENDANT], - load_status=True, - filter_parent=False, - ) + plugins = await DbPluginInfo.get_plugins( + load_status=True, + filter_parent=False, ) + for plugin in plugins: + plugin_type = plugin.plugin_type + if plugin_type == PluginType.NORMAL: + plugin_count.normal += 1 + if plugin_type in {PluginType.ADMIN, PluginType.SUPER_AND_ADMIN}: + plugin_count.admin += 1 + if plugin_type in {PluginType.SUPERUSER, PluginType.SUPER_AND_ADMIN}: + plugin_count.superuser += 1 + if plugin_type in {PluginType.HIDDEN, PluginType.DEPENDANT}: + plugin_count.other += 1 return Result.ok(plugin_count, "拿到信息啦!") except Exception as e: logger.error(f"{router.prefix}/get_plugin_count 调用错误", "WebUi", e=e) diff --git a/zhenxun/builtin_plugins/web_ui/api/tabs/plugin_manage/data_source.py b/zhenxun/builtin_plugins/web_ui/api/tabs/plugin_manage/data_source.py index cc6458e4..fdae3c98 100644 --- a/zhenxun/builtin_plugins/web_ui/api/tabs/plugin_manage/data_source.py +++ b/zhenxun/builtin_plugins/web_ui/api/tabs/plugin_manage/data_source.py @@ -111,10 +111,19 @@ class ApiDataSource: other_update_fields = set() updated_count = 0 errors = [] + modules = [item.module for item in params.updates] + plugin_records = await DbPluginInfo.get_plugins( + module__in=modules, + load_status=None, + filter_parent=False, + ) + plugin_map = {plugin.module: plugin for plugin in plugin_records} for item in params.updates: try: - db_plugin = await DbPluginInfo.get(module=item.module) + db_plugin = plugin_map.get(item.module) + if db_plugin is None: + raise DoesNotExist() plugin_changed_other = False plugin_changed_block = False diff --git a/zhenxun/cli.py b/zhenxun/cli.py index 4e493755..5a908d89 100644 --- a/zhenxun/cli.py +++ b/zhenxun/cli.py @@ -8,12 +8,27 @@ from __future__ import annotations +import atexit import importlib.metadata +import os from pathlib import Path +import signal import subprocess import sys import time +GRACEFUL_SHUTDOWN_TIMEOUT = 15 +WORKER_POLL_INTERVAL = 0.1 +RESTART_POLL_INTERVAL = 0.5 +WORKER_SOFT_EXIT_TIMEOUT = 15.0 +WORKER_TERMINATE_TIMEOUT = 5.0 +WORKER_KILL_TIMEOUT = 5.0 + + +def _launcher_log(message: str) -> None: + sys.stderr.write(f"[zx launcher] {message}\n") + sys.stderr.flush() + def _print_version() -> None: try: @@ -94,32 +109,69 @@ def _run_worker() -> None: nonebot.logger.info(f"加载第三方插件目录: {ext}") nonebot.load_plugins(ext) - nonebot.run() + nonebot.run(timeout_graceful_shutdown=GRACEFUL_SHUTDOWN_TIMEOUT) def _build_worker_command() -> list[str]: return [sys.executable, "-m", "zhenxun.cli", "run-worker"] +def _get_worker_creationflags() -> int: + if os.name == "nt": + return getattr(subprocess, "CREATE_NEW_PROCESS_GROUP", 0) + return 0 + + def _wait_worker_exit(proc: subprocess.Popen, timeout_seconds: float) -> bool: deadline = time.monotonic() + timeout_seconds while time.monotonic() < deadline: if proc.poll() is not None: return True - time.sleep(0.1) + time.sleep(WORKER_POLL_INTERVAL) return proc.poll() is not None def _terminate_worker(proc: subprocess.Popen) -> None: if proc.poll() is not None: return - if _wait_worker_exit(proc, 8.0): - return - proc.terminate() - if _wait_worker_exit(proc, 5.0): + _launcher_log(f"stopping worker pid={proc.pid}") + if os.name == "nt": + ctrl_break_event = getattr(signal, "CTRL_BREAK_EVENT", None) + if ctrl_break_event is not None: + try: + _launcher_log(f"sending CTRL_BREAK_EVENT to worker pid={proc.pid}") + proc.send_signal(ctrl_break_event) + except Exception as e: + _launcher_log(f"failed to send CTRL_BREAK_EVENT: {e!r}") + else: + if _wait_worker_exit(proc, WORKER_SOFT_EXIT_TIMEOUT): + _launcher_log( + f"worker pid={proc.pid} exited after CTRL_BREAK_EVENT " + f"with code {proc.returncode}" + ) + return + _launcher_log( + f"worker pid={proc.pid} did not exit after " + f"{WORKER_SOFT_EXIT_TIMEOUT:.0f}s" + ) + if _wait_worker_exit(proc, 1.0): return + try: + _launcher_log(f"terminating worker pid={proc.pid}") + proc.terminate() + except Exception as e: + _launcher_log(f"failed to terminate worker: {e!r}") + else: + if _wait_worker_exit(proc, WORKER_TERMINATE_TIMEOUT): + _launcher_log( + f"worker pid={proc.pid} exited after terminate with code " + f"{proc.returncode}" + ) + return + _launcher_log(f"worker pid={proc.pid} did not exit after terminate timeout") + _launcher_log(f"killing worker pid={proc.pid}") proc.kill() - proc.wait(timeout=5) + proc.wait(timeout=WORKER_KILL_TIMEOUT) def _run_launcher() -> None: @@ -130,19 +182,83 @@ def _run_launcher() -> None: ) clear_launcher_restart_signal() - while True: - worker = subprocess.Popen(_build_worker_command(), cwd=str(cwd)) + current_worker: subprocess.Popen | None = None + stop_requested = False + stop_signal: int | None = None + + def _cleanup_current_worker() -> None: + if current_worker is not None: + _terminate_worker(current_worker) + + atexit.register(_cleanup_current_worker) + + def _handle_launcher_signal(signum, _frame) -> None: + nonlocal stop_requested, stop_signal + if stop_requested: + _launcher_log(f"received signal {signum} while stopping, exiting launcher") + raise SystemExit(128 + int(signum)) + stop_requested = True + stop_signal = int(signum) + _launcher_log(f"received signal {signum}, scheduling worker shutdown") + + handled_signals = [signal.SIGINT] + if hasattr(signal, "SIGTERM"): + handled_signals.append(signal.SIGTERM) + if hasattr(signal, "SIGBREAK"): + handled_signals.append(signal.SIGBREAK) + for sig in handled_signals: try: - return_code = worker.wait() + signal.signal(sig, _handle_launcher_signal) + except Exception: + pass + + while True: + if stop_requested: + raise SystemExit(128 + int(stop_signal or signal.SIGINT)) + worker_env = os.environ.copy() + worker_env["ZHENXUN_LAUNCHER_PID"] = str(os.getpid()) + worker = subprocess.Popen( + _build_worker_command(), + cwd=str(cwd), + creationflags=_get_worker_creationflags(), + env=worker_env, + ) + current_worker = worker + restart_requested = False + return_code: int | None = None + next_restart_check = 0.0 + try: + while True: + return_code = worker.poll() + if return_code is not None: + break + if stop_requested: + clear_launcher_restart_signal() + _terminate_worker(worker) + raise SystemExit(128 + int(stop_signal or signal.SIGINT)) + now = time.monotonic() + if now >= next_restart_check: + next_restart_check = now + RESTART_POLL_INTERVAL + if consume_launcher_restart_signal(): + restart_requested = True + _launcher_log( + "detected restart request, stopping current worker" + ) + _terminate_worker(worker) + return_code = worker.poll() + break + time.sleep(WORKER_POLL_INTERVAL) except KeyboardInterrupt: clear_launcher_restart_signal() _terminate_worker(worker) return + finally: + if current_worker is worker: + current_worker = None - should_restart = consume_launcher_restart_signal() - if should_restart: + if restart_requested or consume_launcher_restart_signal(): continue - raise SystemExit(return_code) + raise SystemExit(return_code if return_code is not None else 1) def main() -> None: diff --git a/zhenxun/configs/utils/__init__.py b/zhenxun/configs/utils/__init__.py index 15a746f2..5b219e80 100644 --- a/zhenxun/configs/utils/__init__.py +++ b/zhenxun/configs/utils/__init__.py @@ -4,12 +4,12 @@ from pathlib import Path from typing import Any, TypeVar import cattrs +from nonebot.log import logger as _nonebot_logger from pydantic import BaseModel, Field from ruamel.yaml import YAML from ruamel.yaml.scanner import ScannerError from zhenxun.configs.path_config import DATA_PATH -from zhenxun.services.log import logger from zhenxun.utils.pydantic_compat import ( _dump_pydantic_obj, _is_pydantic_type, @@ -38,6 +38,33 @@ _yaml.indent = 2 _yaml.allow_unicode = True T = TypeVar("T") +_MISSING = object() + + +class _ConfigLogger: + @staticmethod + def _emit(level: str, info: str, *, e: Exception | None = None) -> None: + logger = _nonebot_logger.opt(exception=e) if e else _nonebot_logger + getattr(logger, level)(info) + + @classmethod + def debug(cls, info: str, *_, e: Exception | None = None, **__) -> None: + cls._emit("debug", info, e=e) + + @classmethod + def info(cls, info: str, *_, e: Exception | None = None, **__) -> None: + cls._emit("info", info, e=e) + + @classmethod + def warning(cls, info: str, *_, e: Exception | None = None, **__) -> None: + cls._emit("warning", info, e=e) + + @classmethod + def error(cls, info: str, *_, e: Exception | None = None, **__) -> None: + cls._emit("error", info, e=e) + + +logger = _ConfigLogger() class NoSuchConfig(Exception): @@ -114,21 +141,83 @@ class ConfigsManager: self._simple_data: dict = {} self._simple_file = DATA_PATH / "config.yaml" self.add_module = [] - _yaml = YAML() if file: file.parent.mkdir(exist_ok=True, parents=True) self.file = file self.load_data() if self._simple_file.exists(): - try: - with self._simple_file.open(encoding="utf8") as f: - self._simple_data = _yaml.load(f) - except ScannerError as e: - raise ScannerError( - f"{e}\n**********************************************\n" - f"****** 可能为config.yaml配置文件填写不规范 ******\n" - f"**********************************************" - ) from e + self._load_simple_data(raise_on_error=True) + self._apply_simple_data(warn_unknown=False) + + def _load_simple_data(self, *, raise_on_error: bool = False) -> None: + if not self._simple_file.exists(): + self._simple_data = {} + return + try: + with self._simple_file.open(encoding="utf8") as f: + simple_data = _yaml.load(f) or {} + except ScannerError as e: + message = ( + f"{e}\n**********************************************\n" + f"****** 可能为config.yaml配置文件填写不规范 ******\n" + f"**********************************************" + ) + if raise_on_error: + raise ScannerError(message) from e + logger.warning(f"读取config.yaml失败,已跳过本次重载: {message}", e=e) + return + except Exception as e: + if raise_on_error: + raise RuntimeError(f"读取config.yaml失败: {e}") from e + logger.warning(f"读取config.yaml失败,已跳过本次重载: {e}", e=e) + return + if not isinstance(simple_data, dict): + message = "config.yaml 顶层必须为字典,已忽略当前内容。" + if raise_on_error: + raise ValueError(message) + logger.warning(message) + self._simple_data = {} + return + self._simple_data = simple_data + + @staticmethod + def _find_mapping_key(data: dict, key: str) -> str | None: + if key in data: + return key + upper_key = key.upper() + for raw_key in data: + if str(raw_key).upper() == upper_key: + return raw_key + return None + + def _get_simple_config_value(self, module: str, key: str) -> Any: + module_data = self._simple_data.get(module) + if not isinstance(module_data, dict): + return _MISSING + simple_key = self._find_mapping_key(module_data, key.upper()) + if simple_key is None: + return _MISSING + return module_data[simple_key] + + def _apply_simple_data(self, *, warn_unknown: bool) -> None: + for module, module_data in self._simple_data.items(): + if not isinstance(module_data, dict): + if warn_unknown: + logger.warning(f"配置组 {module} 不是字典,已跳过。") + continue + config_group = self._data.get(module) + if not config_group: + if warn_unknown: + logger.warning(f"未知配置组 {module},已跳过。") + continue + for raw_key, value in module_data.items(): + key = str(raw_key).upper() + config_key = self._find_mapping_key(config_group.configs, key) + if config_key is None: + if warn_unknown: + logger.warning(f"未知配置项 {module}.{raw_key},已跳过。") + continue + config_group.configs[config_key].value = value def set_name(self, module: str, name: str): """设置插件配置中文名出 @@ -223,7 +312,11 @@ class ConfigsManager: if module in self._data and (config := self._data[module].configs.get(key)): existing_value = config.value - processed_value = self._normalize_config_data(value, existing_value) + simple_value = self._get_simple_config_value(module, key) + if simple_value is _MISSING: + processed_value = self._normalize_config_data(value, existing_value) + else: + processed_value = self._normalize_config_data(value, simple_value) processed_default_value = self._normalize_config_data(default_value) self.add_module.append(f"{module}:{key}".lower()) @@ -231,7 +324,7 @@ class ConfigsManager: config.help = help config.arg_parser = arg_parser config.type = type - if _override: + if simple_value is not _MISSING or _override: config.value = processed_value config.default_value = processed_default_value else: @@ -371,12 +464,8 @@ class ConfigsManager: def reload(self): """重新加载配置文件""" - if self._simple_file.exists(): - with open(self._simple_file, encoding="utf8") as f: - self._simple_data = _yaml.load(f) - for key in self._simple_data.keys(): - for k in self._simple_data[key].keys(): - self._data[key].configs[k].value = self._simple_data[key][k] + self._load_simple_data() + self._apply_simple_data(warn_unknown=True) self.save() def load_data(self): diff --git a/zhenxun/models/group_member_info.py b/zhenxun/models/group_member_info.py index fd8f36c5..ee5f3584 100644 --- a/zhenxun/models/group_member_info.py +++ b/zhenxun/models/group_member_info.py @@ -2,6 +2,7 @@ from typing import ClassVar from tortoise import fields +from zhenxun.configs.config import BotConfig from zhenxun.services.db_context import Model @@ -51,4 +52,54 @@ class GroupInfoUser(Model): @classmethod async def _run_script(cls): - return ["ALTER TABLE group_info_users DROP COLUMN nickname;"] + db_type = (BotConfig.get_sql_type() or "").lower() + scripts = ["ALTER TABLE group_info_users DROP COLUMN nickname;"] + + if "postgres" in db_type: + scripts.extend( + [ + ( + "ALTER TABLE group_info_users ADD COLUMN IF NOT EXISTS " + "platform character varying(255);" + ), + ( + "ALTER TABLE group_info_users ALTER COLUMN user_id " + "TYPE character varying(255) USING user_id::character varying;" + ), + ( + "ALTER TABLE group_info_users ALTER COLUMN group_id " + "TYPE character varying(255) USING group_id::character varying;" + ), + ( + "ALTER TABLE group_info_users ALTER COLUMN platform " + "TYPE character varying(255) USING platform::character varying;" + ), + ] + ) + elif "mysql" in db_type: + scripts.extend( + [ + ( + "ALTER TABLE group_info_users ADD COLUMN " + "platform VARCHAR(255) NULL;" + ), + ( + "ALTER TABLE group_info_users MODIFY COLUMN " + "user_id VARCHAR(255) NOT NULL;" + ), + ( + "ALTER TABLE group_info_users MODIFY COLUMN " + "group_id VARCHAR(255) NOT NULL;" + ), + ( + "ALTER TABLE group_info_users MODIFY COLUMN " + "platform VARCHAR(255) NULL;" + ), + ] + ) + elif "sqlite" in db_type: + scripts.append( + "ALTER TABLE group_info_users ADD COLUMN platform VARCHAR(255);" + ) + + return scripts diff --git a/zhenxun/models/user_console.py b/zhenxun/models/user_console.py index f9dbcb11..50f49fe2 100644 --- a/zhenxun/models/user_console.py +++ b/zhenxun/models/user_console.py @@ -5,12 +5,11 @@ from tortoise import fields from tortoise.exceptions import IntegrityError from zhenxun.models.goods_info import GoodsInfo +from zhenxun.services.buffered_writers import append_user_gold_log from zhenxun.services.db_context import Model from zhenxun.utils.enum import CacheType, GoldHandle from zhenxun.utils.exception import GoodsNotFound, InsufficientGold -from .user_gold_log import UserGoldLog - class UserConsole(Model): id = fields.IntField(pk=True, generated=True, auto_increment=True) @@ -77,6 +76,17 @@ class UserConsole(Model): user, _ = await cls.get_or_create_user(user_id=user_id, platform=platform) return user + @classmethod + async def _get_user_for_write( + cls, user_id: str, platform: str | None = None + ) -> "UserConsole": + """获取写入用用户;已有用户不走 get_or_create,避免重复清理缓存。""" + user = await cls.get_or_none(user_id=user_id) + if user is not None: + return user + user, _ = await cls.get_or_create_user(user_id=user_id, platform=platform) + return user + @classmethod async def get_new_uid(cls) -> int: """获取最新uid @@ -103,10 +113,10 @@ class UserConsole(Model): source: 来源 platform: 平台. """ - user, _ = await cls.get_or_create_user(user_id=user_id, platform=platform) + user = await cls._get_user_for_write(user_id=user_id, platform=platform) user.gold += gold await user.save(update_fields=["gold"]) - await UserGoldLog.create( + await append_user_gold_log( user_id=user_id, gold=gold, handle=GoldHandle.GET, source=source ) @@ -131,12 +141,12 @@ class UserConsole(Model): 异常: InsufficientGold: 金币不足 """ - user, _ = await cls.get_or_create_user(user_id=user_id, platform=platform) + user = await cls._get_user_for_write(user_id=user_id, platform=platform) if user.gold < gold: raise InsufficientGold() user.gold -= gold await user.save(update_fields=["gold"]) - await UserGoldLog.create( + await append_user_gold_log( user_id=user_id, gold=gold, handle=handle, source=plugin_module ) @@ -152,7 +162,7 @@ class UserConsole(Model): num: 道具数量. platform: 平台. """ - user, _ = await cls.get_or_create_user(user_id=user_id, platform=platform) + user = await cls._get_user_for_write(user_id=user_id, platform=platform) if goods_uuid not in user.props: user.props[goods_uuid] = 0 user.props[goods_uuid] += num @@ -186,7 +196,7 @@ class UserConsole(Model): num: 道具数量. platform: 平台. """ - user, _ = await cls.get_or_create_user(user_id=user_id, platform=platform) + user = await cls._get_user_for_write(user_id=user_id, platform=platform) if goods_uuid not in user.props or user.props[goods_uuid] < num: raise GoodsNotFound("未找到商品或道具数量不足...") diff --git a/zhenxun/services/avatar_service.py b/zhenxun/services/avatar_service.py index 8e46ac67..6eeeb856 100644 --- a/zhenxun/services/avatar_service.py +++ b/zhenxun/services/avatar_service.py @@ -63,6 +63,11 @@ class AvatarService: identifier = str(identifier) return self.cache_path / platform / f"{identifier}.png" + def clear_memory_cache(self) -> int: + size = len(self._memory_cache) + self._memory_cache.clear() + return size + async def get_avatar_path( self, platform: str, identifier: str, force_refresh: bool = False ) -> Path | None: diff --git a/zhenxun/services/buffered_writers.py b/zhenxun/services/buffered_writers.py new file mode 100644 index 00000000..b2dd7842 --- /dev/null +++ b/zhenxun/services/buffered_writers.py @@ -0,0 +1,128 @@ +from __future__ import annotations + +import asyncio +from collections import deque +import contextlib +import time + +from zhenxun.models.user_gold_log import UserGoldLog +from zhenxun.services.log import logger +from zhenxun.utils.enum import GoldHandle +from zhenxun.utils.manager.priority_manager import PriorityLifecycle + +LOG_COMMAND = "BufferedWriters" + +_USER_GOLD_LOG_BUFFER_MAX_RETAIN = 10_000 +_USER_GOLD_LOG_FLUSH_TRIGGER_SIZE = 128 +_USER_GOLD_LOG_FLUSH_BATCH_SIZE = 500 +_USER_GOLD_LOG_FLUSH_INTERVAL_SECONDS = 60.0 +_USER_GOLD_LOG_DROP_LOG_INTERVAL_SECONDS = 10.0 + +_user_gold_log_buffer: deque[UserGoldLog] = deque() +_user_gold_log_buffer_lock = asyncio.Lock() +_user_gold_log_flush_lock = asyncio.Lock() +_user_gold_log_flush_task: asyncio.Task[None] | None = None +_user_gold_log_dropped = 0 +_user_gold_log_last_drop_log_at = 0.0 + + +def _ensure_user_gold_log_flush_task() -> None: + global _user_gold_log_flush_task + if _user_gold_log_flush_task is not None and not _user_gold_log_flush_task.done(): + return + _user_gold_log_flush_task = asyncio.create_task(_user_gold_log_flush_loop()) + + +def _record_user_gold_log_drop() -> None: + global _user_gold_log_dropped, _user_gold_log_last_drop_log_at + _user_gold_log_dropped += 1 + now = time.monotonic() + if now - _user_gold_log_last_drop_log_at < _USER_GOLD_LOG_DROP_LOG_INTERVAL_SECONDS: + return + _user_gold_log_last_drop_log_at = now + logger.warning( + "user_gold_log buffer full, dropped " + f"{_user_gold_log_dropped} records, backlog={len(_user_gold_log_buffer)}", + LOG_COMMAND, + ) + + +async def _user_gold_log_flush_loop() -> None: + while True: + await asyncio.sleep(_USER_GOLD_LOG_FLUSH_INTERVAL_SECONDS) + try: + await flush_user_gold_log_buffer("定时") + except asyncio.CancelledError: + raise + except Exception as exc: + logger.warning("定时批量写入金币流水失败", LOG_COMMAND, e=exc) + + +async def append_user_gold_log( + user_id: str, + gold: int, + handle: GoldHandle, + source: str | None = None, +) -> None: + _ensure_user_gold_log_flush_task() + record = UserGoldLog(user_id=user_id, gold=gold, handle=handle, source=source) + async with _user_gold_log_buffer_lock: + if len(_user_gold_log_buffer) >= _USER_GOLD_LOG_BUFFER_MAX_RETAIN: + _user_gold_log_buffer.popleft() + _record_user_gold_log_drop() + _user_gold_log_buffer.append(record) + should_flush = ( + len(_user_gold_log_buffer) >= _USER_GOLD_LOG_FLUSH_TRIGGER_SIZE + and not _user_gold_log_flush_lock.locked() + ) + if should_flush: + await flush_user_gold_log_buffer("缓冲区触发") + + +async def flush_user_gold_log_buffer(reason: str) -> int: + async with _user_gold_log_flush_lock: + written = 0 + while True: + batch: list[UserGoldLog] = [] + async with _user_gold_log_buffer_lock: + if not _user_gold_log_buffer: + break + while ( + _user_gold_log_buffer + and len(batch) < _USER_GOLD_LOG_FLUSH_BATCH_SIZE + ): + batch.append(_user_gold_log_buffer.popleft()) + if not batch: + break + try: + await UserGoldLog.bulk_create(batch, _USER_GOLD_LOG_FLUSH_BATCH_SIZE) + except Exception as exc: + async with _user_gold_log_buffer_lock: + retain_count = max( + _USER_GOLD_LOG_BUFFER_MAX_RETAIN - len(_user_gold_log_buffer), + 0, + ) + for record in reversed(batch[-retain_count:]): + _user_gold_log_buffer.appendleft(record) + logger.error(f"{reason}批量写入金币流水失败", LOG_COMMAND, e=exc) + return written + written += len(batch) + if written: + logger.debug(f"{reason}批量写入金币流水 {written} 条", LOG_COMMAND) + return written + + +async def stop_user_gold_log_buffer() -> int: + global _user_gold_log_flush_task + task = _user_gold_log_flush_task + _user_gold_log_flush_task = None + if task is not None: + task.cancel() + with contextlib.suppress(BaseException): + await task + return await flush_user_gold_log_buffer("关闭") + + +@PriorityLifecycle.on_shutdown(priority=90) +async def _flush_user_gold_log_buffer_on_shutdown() -> None: + await stop_user_gold_log_buffer() diff --git a/zhenxun/services/cache/__init__.py b/zhenxun/services/cache/__init__.py index 02ce516f..44d1b9c8 100644 --- a/zhenxun/services/cache/__init__.py +++ b/zhenxun/services/cache/__init__.py @@ -74,6 +74,7 @@ from .config import ( ) __all__ = [ + "BoundedTTLCache", "Cache", "CacheDict", "CacheManager", @@ -82,6 +83,7 @@ __all__ = [ ] from . import runtime_cache as _runtime_cache # noqa: F401 +from .bounded_ttl import BoundedTTLCache T = TypeVar("T") U = TypeVar("U") @@ -400,9 +402,10 @@ class CacheManager: """清除缓存 参数: - cache_type: 缓存类型,为None时清除所有缓存。 - 注意:受 aiocache 限制,无法按类型精确删除, - 指定 cache_type 时仅清除整个 backend(行为与不指定相同)。 + cache_type: 缓存类型。为 None 时清除整个 backend。 + 指定 cache_type 时不再退化为清除整个 backend,避免误删其他类型缓存。 + 需要刷新模型运行态缓存时,应调用对应 + RuntimeCache.refresh/upsert/remove。 返回: bool: 是否成功 @@ -413,11 +416,12 @@ class CacheManager: try: if cache_type: - logger.debug( - f"清除缓存类型 {cache_type}" - "(aiocache 不支持按前缀删除,清除整个 backend)", + logger.warning( + f"拒绝清除缓存类型 {cache_type}: " + "当前后端不支持可靠的按类型清理,已避免清除整个 backend", LOG_COMMAND, ) + return False await self.cache_backend.clear() # type: ignore return True except Exception as e: diff --git a/zhenxun/services/cache/bounded_ttl.py b/zhenxun/services/cache/bounded_ttl.py new file mode 100644 index 00000000..04dc9622 --- /dev/null +++ b/zhenxun/services/cache/bounded_ttl.py @@ -0,0 +1,218 @@ +from __future__ import annotations + +import asyncio +from collections import OrderedDict +from collections.abc import Callable +from dataclasses import dataclass +import sys +import time +from typing import Generic, TypeVar +import weakref + +K = TypeVar("K") +V = TypeVar("V") + + +def _default_sizeof(value: object) -> int: + if isinstance(value, bytes | bytearray | memoryview): + return len(value) + return 0 + + +@dataclass(frozen=True) +class BoundedTTLCacheStats: + name: str + items: int + max_items: int + total_bytes: int + max_total_bytes: int | None + hits: int + misses: int + sets: int + evictions: int + + def to_dict(self) -> dict[str, int | str | None]: + return { + "name": self.name, + "items": self.items, + "max_items": self.max_items, + "total_bytes": self.total_bytes, + "max_total_bytes": self.max_total_bytes, + "hits": self.hits, + "misses": self.misses, + "sets": self.sets, + "evictions": self.evictions, + } + + +class BoundedTTLCache(Generic[K, V]): + """Small async TTL/LRU cache with optional total-byte limit.""" + + _instances: weakref.WeakSet["BoundedTTLCache"] = weakref.WeakSet() + + def __init__( + self, + name: str, + ttl_seconds: float, + max_items: int, + max_total_bytes: int | None = None, + sizeof: Callable[[V], int] | None = None, + ) -> None: + self.name = name.upper() + self._ttl_seconds = max(ttl_seconds, 0.0) + self._max_items = max(max_items, 1) + self._max_total_bytes = ( + max_total_bytes + if isinstance(max_total_bytes, int) and max_total_bytes > 0 + else None + ) + self._sizeof = sizeof or _default_sizeof + self._cache: OrderedDict[K, tuple[float, V, int]] = OrderedDict() + self._total_bytes = 0 + self._hits = 0 + self._misses = 0 + self._sets = 0 + self._evictions = 0 + self._lock = asyncio.Lock() + self.__class__._instances.add(self) + + def _expire_at(self, now: float) -> float: + if self._ttl_seconds <= 0: + return sys.float_info.max + return now + self._ttl_seconds + + def _value_size(self, value: V) -> int: + try: + return max(0, int(self._sizeof(value))) + except Exception: + return 0 + + def _remove_key_nolock(self, key: K) -> bool: + item = self._cache.pop(key, None) + if item is None: + return False + self._total_bytes -= item[2] + if self._total_bytes < 0: + self._total_bytes = 0 + return True + + def _pop_oldest_nolock(self) -> bool: + if not self._cache: + return False + _, (_, _, size) = self._cache.popitem(last=False) + self._total_bytes -= size + if self._total_bytes < 0: + self._total_bytes = 0 + self._evictions += 1 + return True + + def _cleanup_nolock(self, now: float) -> None: + expired_keys = [ + key for key, (expire_at, _, _) in self._cache.items() if expire_at <= now + ] + for key in expired_keys: + if self._remove_key_nolock(key): + self._evictions += 1 + while len(self._cache) > self._max_items: + self._pop_oldest_nolock() + if self._max_total_bytes is not None: + while self._total_bytes > self._max_total_bytes and self._cache: + self._pop_oldest_nolock() + + async def get(self, key: K) -> V | None: + now = time.monotonic() + async with self._lock: + self._cleanup_nolock(now) + item = self._cache.get(key) + if item is None: + self._misses += 1 + return None + expire_at, value, _ = item + if expire_at <= now: + self._remove_key_nolock(key) + self._misses += 1 + return None + self._cache.move_to_end(key) + self._hits += 1 + return value + + async def set(self, key: K, value: V) -> bool: + value_size = self._value_size(value) + if self._max_total_bytes is not None and value_size > self._max_total_bytes: + return False + + now = time.monotonic() + async with self._lock: + self._remove_key_nolock(key) + self._cache[key] = (self._expire_at(now), value, value_size) + self._total_bytes += value_size + self._sets += 1 + self._cache.move_to_end(key) + self._cleanup_nolock(now) + return key in self._cache + + async def delete(self, key: K) -> bool: + async with self._lock: + return self._remove_key_nolock(key) + + async def clear(self) -> int: + async with self._lock: + size = len(self._cache) + self._cache.clear() + self._total_bytes = 0 + return size + + async def stats(self) -> BoundedTTLCacheStats: + now = time.monotonic() + async with self._lock: + self._cleanup_nolock(now) + return BoundedTTLCacheStats( + name=self.name, + items=len(self._cache), + max_items=self._max_items, + total_bytes=self._total_bytes, + max_total_bytes=self._max_total_bytes, + hits=self._hits, + misses=self._misses, + sets=self._sets, + evictions=self._evictions, + ) + + @classmethod + async def clear_all(cls) -> dict[str, int]: + result: dict[str, int] = {} + for cache in list(cls._instances): + size = await cache.clear() + if size: + result[cache.name] = result.get(cache.name, 0) + size + return result + + @classmethod + async def stats_all(cls) -> dict[str, dict[str, int | str | None]]: + result: dict[str, dict[str, int | str | None]] = {} + for cache in list(cls._instances): + stats = await cache.stats() + if not stats.items: + continue + if cache.name not in result: + result[cache.name] = stats.to_dict() + continue + current = result[cache.name] + for key in ( + "items", + "max_items", + "total_bytes", + "hits", + "misses", + "sets", + "evictions", + ): + current[key] = int(current.get(key) or 0) + int( + getattr(stats, key) or 0 + ) + current_max_bytes = current.get("max_total_bytes") + if current_max_bytes is not None or stats.max_total_bytes is not None: + current["max_total_bytes"] = int(current_max_bytes or 0) + int( + stats.max_total_bytes or 0 + ) + return result diff --git a/zhenxun/services/cache/cache_containers.py b/zhenxun/services/cache/cache_containers.py index e6829007..47e393b9 100644 --- a/zhenxun/services/cache/cache_containers.py +++ b/zhenxun/services/cache/cache_containers.py @@ -1,9 +1,12 @@ from dataclasses import dataclass import time from typing import Any, Generic, TypeVar +import weakref T = TypeVar("T") +DEFAULT_CACHE_MAX_ITEMS = 10000 + @dataclass class CacheData(Generic[T]): @@ -16,16 +19,21 @@ class CacheData(Generic[T]): class CacheDict(Generic[T]): """缓存字典类,提供类似普通字典的接口,数据只存储在内存中""" - def __init__(self, name: str, expire: int = 0): + _instances: weakref.WeakSet = weakref.WeakSet() + + def __init__(self, name: str, expire: int = 0, max_items: int | None = None): """初始化缓存字典 参数: name: 字典名称 expire: 过期时间(秒),默认为0表示永不过期 + max_items: 最大缓存项数,None 使用统一默认值,0 表示不限制 """ self.name = name.upper() self.expire = expire + self.max_items = DEFAULT_CACHE_MAX_ITEMS if max_items is None else max_items self._data: dict[str, CacheData[T]] = {} + self.__class__._instances.add(self) def expire_time(self, key: str) -> float: """获取字典项的过期时间""" @@ -62,6 +70,7 @@ class CacheDict(Generic[T]): """ expire_time = time.time() + self.expire if self.expire > 0 else 0 self._data[key] = CacheData(value=value, expire_time=expire_time) + self._enforce_limit() def __delitem__(self, key: str) -> None: """删除字典项 @@ -122,6 +131,7 @@ class CacheDict(Generic[T]): expire_time = time.time() + self.expire self._data[key] = CacheData(value=value, expire_time=expire_time) + self._enforce_limit() def pop(self, key: str, default: Any = None) -> T: """删除并返回字典项 @@ -146,6 +156,32 @@ class CacheDict(Generic[T]): """清空字典""" self._data.clear() + def stats(self) -> dict[str, int]: + """返回当前缓存条目统计。""" + self._clean_expired() + return {"items": len(self._data), "max_items": self.max_items} + + @classmethod + def stats_all(cls) -> dict[str, dict[str, int]]: + """返回所有 CacheDict 实例的条目统计。""" + result: dict[str, dict[str, int]] = {} + for cache in list(cls._instances): + stats = cache.stats() + if stats["items"]: + result[cache.name] = stats + return result + + @classmethod + def clear_all(cls) -> dict[str, int]: + """清空所有 CacheDict,返回各缓存清理的条目数。""" + result: dict[str, int] = {} + for cache in list(cls._instances): + size = len(cache._data) + if size: + cache.clear() + result[cache.name] = result.get(cache.name, 0) + size + return result + def keys(self) -> list[str]: """获取所有键 @@ -187,6 +223,12 @@ class CacheDict(Generic[T]): for key in expired_keys: del self._data[key] + def _enforce_limit(self) -> None: + if self.max_items <= 0: + return + while len(self._data) > self.max_items: + self._data.pop(next(iter(self._data))) + def __len__(self) -> int: """获取字典长度 @@ -211,17 +253,22 @@ class CacheDict(Generic[T]): class CacheList(Generic[T]): """缓存列表类,提供类似普通列表的接口,数据只存储在内存中""" - def __init__(self, name: str, expire: int = 0): + _instances: weakref.WeakSet = weakref.WeakSet() + + def __init__(self, name: str, expire: int = 0, max_items: int | None = None): """初始化缓存列表 参数: name: 列表名称 expire: 过期时间(秒),默认为0表示永不过期 + max_items: 最大缓存项数,None 使用统一默认值,0 表示不限制 """ self.name = name.upper() self.expire = expire + self.max_items = DEFAULT_CACHE_MAX_ITEMS if max_items is None else max_items self._data: list[CacheData[T]] = [] self._expire_time = 0 + self.__class__._instances.add(self) # 如果设置了过期时间,计算整个列表的过期时间 if self.expire > 0: @@ -303,6 +350,7 @@ class CacheList(Generic[T]): self.clear() self._data.append(CacheData(value=value)) + self._enforce_limit() # 更新过期时间 self._update_expire_time() @@ -318,6 +366,7 @@ class CacheList(Generic[T]): self.clear() self._data.extend([CacheData(value=v) for v in values]) + self._enforce_limit() # 更新过期时间 self._update_expire_time() @@ -334,6 +383,7 @@ class CacheList(Generic[T]): self.clear() self._data.insert(index, CacheData(value=value)) + self._enforce_limit() # 更新过期时间 self._update_expire_time() @@ -389,6 +439,32 @@ class CacheList(Generic[T]): # 重置过期时间 self._update_expire_time() + def stats(self) -> dict[str, int]: + """返回当前缓存条目统计。""" + if self._is_expired(): + self.clear() + return {"items": len(self._data), "max_items": self.max_items} + + @classmethod + def stats_all(cls) -> dict[str, dict[str, int]]: + """返回所有 CacheList 实例的条目统计。""" + result: dict[str, dict[str, int]] = {} + for cache in list(cls._instances): + stats = cache.stats() + if stats["items"]: + result[cache.name] = stats + return result + + @classmethod + def clear_all(cls) -> dict[str, int]: + result: dict[str, int] = {} + for cache in list(cls._instances): + size = len(cache._data) + if size: + cache.clear() + result[cache.name] = result.get(cache.name, 0) + size + return result + def index(self, value: T, start: int = 0, end: int | None = None) -> int: """查找值的索引 @@ -438,6 +514,12 @@ class CacheList(Generic[T]): """更新过期时间""" self._expire_time = time.time() + self.expire if self.expire > 0 else 0 + def _enforce_limit(self) -> None: + if self.max_items <= 0: + return + if len(self._data) > self.max_items: + del self._data[: len(self._data) - self.max_items] + def __str__(self) -> str: """字符串表示 diff --git a/zhenxun/services/cache/runtime_cache.py b/zhenxun/services/cache/runtime_cache.py index 77c2346c..e779a47c 100644 --- a/zhenxun/services/cache/runtime_cache.py +++ b/zhenxun/services/cache/runtime_cache.py @@ -10,7 +10,13 @@ import uuid from zhenxun.services.cache.config import CacheMode from zhenxun.services.log import logger -from zhenxun.utils.enum import LimitCheckType, LimitWatchType, PluginLimitType +from zhenxun.utils.enum import ( + BlockType, + LimitCheckType, + LimitWatchType, + PluginLimitType, + PluginType, +) from zhenxun.utils.manager.priority_manager import PriorityLifecycle if TYPE_CHECKING: @@ -18,24 +24,6 @@ if TYPE_CHECKING: LOG_COMMAND = "RuntimeCache" -PLUGININFO_MEM_REFRESH_INTERVAL = 1800 # 30分钟 - 插件信息很少变化 -BAN_MEM_REFRESH_INTERVAL = 60 -BAN_MEM_CLEAN_INTERVAL = 60 -BAN_MEM_CLEANUP_DB = True -BAN_MEM_NEGATIVE_TTL = 5 -BOT_MEM_REFRESH_INTERVAL = 300 # 5分钟 -BOT_MEM_NEGATIVE_TTL = 60 -GROUP_MEM_REFRESH_INTERVAL = 300 # 5分钟 - 群组信息很少变化 -GROUP_MEM_NEGATIVE_TTL = 60 -LEVEL_MEM_REFRESH_INTERVAL = 300 # 5分钟 - 用户等级很少变化 -LEVEL_MEM_NEGATIVE_TTL = 60 -TASK_MEM_REFRESH_INTERVAL = 900 -TASK_MEM_NEGATIVE_TTL = 60 -LIMIT_MEM_REFRESH_INTERVAL = 300 # 5分钟 -LIMIT_MEM_NEGATIVE_TTL = 30 -RUNTIME_CACHE_SYNC_ENABLED = True -RUNTIME_CACHE_SYNC_CHANNEL = "ZHENXUN_RUNTIME_CACHE_SYNC" - def _coerce_int(value, default: int) -> int: try: @@ -45,6 +33,27 @@ def _coerce_int(value, default: int) -> int: return value_int if value_int >= 0 else default +# RuntimeCache 以模型 save/delete 主动失效为主,周期 refresh 只做兜底。 +# 这些默认值避免低压力运行时频繁全量扫表。 +PLUGININFO_MEM_REFRESH_INTERVAL = 3600 # 60分钟 +BAN_MEM_REFRESH_INTERVAL = 300 +BAN_MEM_CLEAN_INTERVAL = 60 +BAN_MEM_CLEANUP_DB = True +BAN_MEM_NEGATIVE_TTL = 5 +BOT_MEM_REFRESH_INTERVAL = 900 # 15分钟 +BOT_MEM_NEGATIVE_TTL = 60 +GROUP_MEM_REFRESH_INTERVAL = 900 # 15分钟 +GROUP_MEM_NEGATIVE_TTL = 60 +LEVEL_MEM_REFRESH_INTERVAL = 900 # 15分钟 +LEVEL_MEM_NEGATIVE_TTL = 60 +TASK_MEM_REFRESH_INTERVAL = 1800 +TASK_MEM_NEGATIVE_TTL = 60 +LIMIT_MEM_REFRESH_INTERVAL = 900 # 15分钟 +LIMIT_MEM_NEGATIVE_TTL = 30 +RUNTIME_CACHE_SYNC_ENABLED = True +RUNTIME_CACHE_SYNC_CHANNEL = "ZHENXUN_RUNTIME_CACHE_SYNC" + + INSTANCE_ID = uuid.uuid4().hex _CACHE_READY_EVENT = asyncio.Event() @@ -89,6 +98,89 @@ def _parse_block_modules(value: str) -> frozenset[str]: return frozenset(items) +@dataclass(frozen=True) +class PluginInfoSnapshot: + id: int + module: str + module_path: str + name: str + status: bool + block_type: BlockType | None + load_status: bool + author: str | None + version: str | None + level: int + default_status: bool + limit_superuser: bool + menu_type: str + plugin_type: PluginType | None + cost_gold: int + admin_level: int | None + ignore_prompt: bool + is_delete: bool + parent: str | None + is_show: bool + ignore_statistics: bool + impression: float + + @classmethod + def from_model(cls, model) -> "PluginInfoSnapshot": + return cls( + id=int(getattr(model, "id", 0) or 0), + module=str(getattr(model, "module", "") or ""), + module_path=str(getattr(model, "module_path", "") or ""), + name=str(getattr(model, "name", "") or ""), + status=bool(getattr(model, "status", True)), + block_type=getattr(model, "block_type", None), + load_status=bool(getattr(model, "load_status", True)), + author=getattr(model, "author", None), + version=getattr(model, "version", None), + level=int(getattr(model, "level", 0) or 0), + default_status=bool(getattr(model, "default_status", True)), + limit_superuser=bool(getattr(model, "limit_superuser", False)), + menu_type=str(getattr(model, "menu_type", "") or ""), + plugin_type=getattr(model, "plugin_type", None), + cost_gold=int(getattr(model, "cost_gold", 0) or 0), + admin_level=getattr(model, "admin_level", None), + ignore_prompt=bool(getattr(model, "ignore_prompt", False)), + is_delete=bool(getattr(model, "is_delete", False)), + parent=getattr(model, "parent", None), + is_show=bool(getattr(model, "is_show", True)), + ignore_statistics=bool(getattr(model, "ignore_statistics", False)), + impression=float(getattr(model, "impression", 0) or 0), + ) + + def to_model(self): + from zhenxun.models.plugin_info import PluginInfo + + plugin = PluginInfo( + id=self.id, + module=self.module, + module_path=self.module_path, + name=self.name, + status=self.status, + block_type=self.block_type, + 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, + 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, + ) + plugin._saved_in_db = True + return plugin + + @dataclass(frozen=True) class BanEntry: user_id: str | None @@ -465,6 +557,20 @@ class RuntimeCacheSync: @classmethod async def stop(cls) -> None: + cls._ready = False + if cls._publish_tasks: + tasks = list(cls._publish_tasks) + try: + await asyncio.wait_for( + asyncio.gather(*tasks, return_exceptions=True), + timeout=1.0, + ) + except asyncio.TimeoutError: + for task in tasks: + if not task.done(): + task.cancel() + finally: + cls._publish_tasks.difference_update(tasks) if cls._task and not cls._task.done(): cls._task.cancel() cls._task = None @@ -480,7 +586,6 @@ class RuntimeCacheSync: except Exception: pass cls._redis = None - cls._ready = False @classmethod def publish_event(cls, cache_type: str, action: str, data: dict[str, Any]) -> None: @@ -559,25 +664,43 @@ class RuntimeCacheSync: class PluginInfoMemoryCache: _lock: ClassVar[asyncio.Lock] = asyncio.Lock() - _by_module: ClassVar[dict[str, "PluginInfo"]] = {} - _by_module_path: ClassVar[dict[str, "PluginInfo"]] = {} + _by_module: ClassVar[dict[str, PluginInfoSnapshot]] = {} + _by_module_path: ClassVar[dict[str, PluginInfoSnapshot]] = {} _loaded: ClassVar[bool] = False _refresh_task: ClassVar[asyncio.Task | None] = None _last_refresh: ClassVar[float] = 0.0 + @classmethod + def _to_model(cls, snapshot: PluginInfoSnapshot | None) -> "PluginInfo | None": + return snapshot.to_model() if snapshot else None + + @classmethod + def _store_snapshot(cls, snapshot: PluginInfoSnapshot) -> None: + if snapshot.module: + old = cls._by_module.get(snapshot.module) + if old and old.module_path != snapshot.module_path: + cls._by_module_path.pop(old.module_path, None) + cls._by_module[snapshot.module] = snapshot + if snapshot.module_path: + old = cls._by_module_path.get(snapshot.module_path) + if old and old.module != snapshot.module: + cls._by_module.pop(old.module, None) + cls._by_module_path[snapshot.module_path] = snapshot + @classmethod async def refresh(cls) -> None: from zhenxun.models.plugin_info import PluginInfo async with cls._lock: plugins = await PluginInfo.all() - by_module: dict[str, "PluginInfo"] = {} - by_module_path: dict[str, "PluginInfo"] = {} + by_module: dict[str, PluginInfoSnapshot] = {} + by_module_path: dict[str, PluginInfoSnapshot] = {} for plugin in plugins: - if plugin.module: - by_module[plugin.module] = plugin - if plugin.module_path: - by_module_path[plugin.module_path] = plugin + snapshot = PluginInfoSnapshot.from_model(plugin) + if snapshot.module: + 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 @@ -596,42 +719,42 @@ class PluginInfoMemoryCache: async def get_by_module(cls, module: str) -> "PluginInfo | None": if not cls._loaded: await cls.ensure_loaded() - return cls._by_module.get(module) + return cls._to_model(cls._by_module.get(module)) @classmethod async def get_all(cls) -> dict[str, "PluginInfo"]: if not cls._loaded: await cls.ensure_loaded() - return dict(cls._by_module) + return { + module: snapshot.to_model() for module, snapshot in cls._by_module.items() + } @classmethod def get_by_module_path(cls, module_path: str) -> "PluginInfo | None": - return cls._by_module_path.get(module_path) + return cls._to_model(cls._by_module_path.get(module_path)) @classmethod def set_plugin(cls, plugin) -> None: if not plugin: return - if plugin.module: - cls._by_module[plugin.module] = plugin - if getattr(plugin, "module_path", None): - cls._by_module_path[plugin.module_path] = plugin + 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: - cls._by_module.pop(module, None) + snapshot = cls._by_module.pop(module, None) + if snapshot and snapshot.module_path: + cls._by_module_path.pop(snapshot.module_path, None) @classmethod async def upsert_from_model(cls, plugin) -> None: if not plugin: return async with cls._lock: - if getattr(plugin, "module", None): - cls._by_module[plugin.module] = plugin - if getattr(plugin, "module_path", None): - cls._by_module_path[plugin.module_path] = plugin + snapshot = PluginInfoSnapshot.from_model(plugin) + cls._store_snapshot(snapshot) cls._loaded = True cls._last_refresh = time.time() @@ -643,9 +766,13 @@ class PluginInfoMemoryCache: return async with cls._lock: if module: - cls._by_module.pop(module, None) + snapshot = cls._by_module.pop(module, None) + if snapshot and snapshot.module_path: + cls._by_module_path.pop(snapshot.module_path, None) if module_path: - cls._by_module_path.pop(module_path, None) + snapshot = cls._by_module_path.pop(module_path, None) + if snapshot and snapshot.module: + cls._by_module.pop(snapshot.module, None) @classmethod async def _refresh_loop(cls, interval: int) -> None: diff --git a/zhenxun/services/data_access.py b/zhenxun/services/data_access.py index fa3bb6e5..ec0ef714 100644 --- a/zhenxun/services/data_access.py +++ b/zhenxun/services/data_access.py @@ -10,7 +10,11 @@ T = TypeVar("T", bound=Model) class DataAccess(Generic[T]): - """数据访问层,根据配置决定是否使用缓存 + """数据访问兼容层,根据配置保留单点缓存读取和清理能力 + + 新的高频运行态路径应优先使用 RuntimeCache 或 BoundedTTLCache。 + 这里不再把 filter/all/create/update_or_create 结果写入通用缓存, + create/update_or_create 只负责清理旧缓存,避免旧值残留。 使用示例: ```python @@ -395,34 +399,25 @@ class DataAccess(Generic[T]): return COMPOSITE_KEY_SEPARATOR.join(key_parts) - async def _cache_items(self, data_list: list[T]) -> None: - """将数据列表存入缓存 - - 参数: - data_list: 数据列表 - """ - if ( - not data_list - or not self.cache_type - or cache_config.cache_mode == CacheMode.NONE - ): + async def _invalidate_item_cache(self, item: T, action: str) -> None: + if not self.cache_type or cache_config.cache_mode == CacheMode.NONE: return try: - # 遍历数据列表,将每条数据存入缓存 - cached_count = 0 - for item in data_list: - cache_key = self._build_cache_key_for_item(item) - if cache_key is not None: - await self.cache.set(cache_key, item) - cached_count += 1 - self._cache_stats[self.cache_type]["sets"] += 1 + cache_key = self._build_cache_key_for_item(item) + if cache_key is None: + return + await self.cache.delete(cache_key) + self._cache_stats[self.cache_type]["deletes"] += 1 logger.debug( - f"{self.model_cls.__name__} 批量缓存: {cached_count}/{len(data_list)}项" + f"{self.model_cls.__name__} {action}: 已失效兼容缓存: {cache_key}" ) except Exception as e: - logger.error(f"{self.model_cls.__name__} 批量缓存失败", e=e) + logger.error( + f"{self.model_cls.__name__} {action}: 更新兼容缓存失败", + e=e, + ) async def filter(self, *args, **kwargs) -> list[T]: """筛选数据 @@ -441,9 +436,6 @@ class DataAccess(Generic[T]): f"{self.model_cls.__name__} filter: 查询结果数量: {len(data_list)}" ) - # 将数据存入缓存 - await self._cache_items(data_list) - return data_list async def all(self) -> list[T]: @@ -457,9 +449,6 @@ class DataAccess(Generic[T]): data_list = await self.model_cls.all() logger.debug(f"{self.model_cls.__name__} all: 查询结果数量: {len(data_list)}") - # 将数据存入缓存 - await self._cache_items(data_list) - return data_list async def count(self, *args, **kwargs) -> int: @@ -501,24 +490,7 @@ class DataAccess(Generic[T]): logger.debug(f"{self.model_cls.__name__} create: 创建数据, 参数: {kwargs}") data = await self.model_cls.create(**kwargs) - # 如果有缓存类型,将数据存入缓存 - if self.cache_type and cache_config.cache_mode != CacheMode.NONE: - try: - # 生成缓存键 - cache_key = self._build_cache_key_for_item(data) - if cache_key is not None: - # 存入缓存 - await self.cache.set(cache_key, data) - self._cache_stats[self.cache_type]["sets"] += 1 - logger.debug( - f"{self.model_cls.__name__} create: " - f"新创建的数据已存入缓存: {cache_key}" - ) - except Exception as e: - logger.error( - f"{self.model_cls.__name__} create: 存入缓存失败,参数: {kwargs}", - e=e, - ) + await self._invalidate_item_cache(data, "create") return data @@ -539,18 +511,7 @@ class DataAccess(Generic[T]): defaults=defaults, **kwargs ) - # 如果有缓存类型,将数据存入缓存 - if self.cache_type and cache_config.cache_mode != CacheMode.NONE: - try: - # 生成缓存键 - cache_key = self._build_cache_key_for_item(data) - if cache_key is not None: - # 存入缓存 - await self.cache.set(cache_key, data) - self._cache_stats[self.cache_type]["sets"] += 1 - logger.debug(f"更新或创建的数据已存入缓存: {cache_key}") - except Exception as e: - logger.error(f"存入缓存失败,参数: {kwargs}", e=e) + await self._invalidate_item_cache(data, "update_or_create") return data, created diff --git a/zhenxun/services/db_context/base_model.py b/zhenxun/services/db_context/base_model.py index 07c63ced..91c998a2 100644 --- a/zhenxun/services/db_context/base_model.py +++ b/zhenxun/services/db_context/base_model.py @@ -150,7 +150,7 @@ class Model(TortoiseModel): obj = await cls.filter(**kwargs).using_db(connection).get() result = (obj, False) - if cache_type := cls.get_cache_type(): + if result[1] and (cache_type := cls.get_cache_type()): await CacheRoot.invalidate_cache(cache_type, cls.get_cache_key(result[0])) return result diff --git a/zhenxun/services/group_settings_service.py b/zhenxun/services/group_settings_service.py index 49cb2f34..366f71d8 100644 --- a/zhenxun/services/group_settings_service.py +++ b/zhenxun/services/group_settings_service.py @@ -5,10 +5,9 @@ import ujson as json from zhenxun.configs.config import Config from zhenxun.models.group_plugin_setting import GroupPluginSetting -from zhenxun.services.cache import Cache +from zhenxun.services.cache import BoundedTTLCache from zhenxun.services.data_access import DataAccess from zhenxun.services.log import logger -from zhenxun.utils.enum import CacheType from zhenxun.utils.pydantic_compat import model_dump, model_validate, parse_as T = TypeVar("T", bound=BaseModel) @@ -22,7 +21,11 @@ class GroupSettingsService: def __init__(self): self.dao = DataAccess(GroupPluginSetting) - self._cache = Cache[dict[str, Any]](CacheType.GROUP_PLUGIN_SETTINGS_VIEW) + self._cache = BoundedTTLCache[str, dict[str, Any]]( + "GROUP_PLUGIN_SETTINGS_VIEW", + ttl_seconds=600, + max_items=10000, + ) @staticmethod def _build_cache_key(group_id: str, plugin_name: str) -> str: diff --git a/zhenxun/services/memory_governor.py b/zhenxun/services/memory_governor.py new file mode 100644 index 00000000..b186c676 --- /dev/null +++ b/zhenxun/services/memory_governor.py @@ -0,0 +1,307 @@ +from __future__ import annotations + +import asyncio +import contextlib +import gc +import inspect +import sys +import time +from typing import Any + +from aiocache import SimpleMemoryCache + +from zhenxun.services.cache import CacheRoot +from zhenxun.services.cache.bounded_ttl import BoundedTTLCache +from zhenxun.services.cache.cache_containers import CacheDict, CacheList +from zhenxun.services.log import logger +from zhenxun.services.message_load import idle_seconds, is_overloaded + +LOG_COMMAND = "MemoryGovernor" + +IDLE_CHECK_INTERVAL_SECONDS = 60 +IDLE_RECLAIM_SECONDS = 600 +RECLAIM_COOLDOWN_SECONDS = 3 * 60 * 60 +RECLAIM_TIMEOUT_SECONDS = 10 + +_task: asyncio.Task | None = None +_reclaim_lock = asyncio.Lock() +_last_reclaim_at = 0.0 + + +def _cooldown_left(now: float | None = None) -> float: + now = time.monotonic() if now is None else now + return max(0.0, _last_reclaim_at + RECLAIM_COOLDOWN_SECONDS - now) + + +async def start_memory_governor() -> None: + global _task + if _task is not None and not _task.done(): + return + if IDLE_CHECK_INTERVAL_SECONDS <= 0 or IDLE_RECLAIM_SECONDS <= 0: + logger.info("idle memory governor disabled", LOG_COMMAND) + return + _task = asyncio.create_task(_idle_reclaim_loop()) + + +async def stop_memory_governor() -> None: + global _task + task = _task + _task = None + if task is not None: + task.cancel() + with contextlib.suppress(BaseException): + await task + + +async def _idle_reclaim_loop() -> None: + while True: + await asyncio.sleep(IDLE_CHECK_INTERVAL_SECONDS) + if not await _should_reclaim(): + continue + if _reclaim_lock.locked(): + continue + async with _reclaim_lock: + if not await _should_reclaim(): + continue + try: + await asyncio.wait_for( + _run_reclaim(), + timeout=max(RECLAIM_TIMEOUT_SECONDS, 1), + ) + except asyncio.TimeoutError: + logger.warning("idle memory reclaim timed out", LOG_COMMAND) + except Exception as exc: + logger.warning("idle memory reclaim failed", LOG_COMMAND, e=exc) + + +async def _should_reclaim() -> bool: + if _cooldown_left() > 0: + return False + if idle_seconds() < IDLE_RECLAIM_SECONDS: + return False + if is_overloaded(): + return False + if await _has_active_auth_work(): + return False + return not await _has_active_render_work() + + +async def _has_active_auth_work() -> bool: + module = sys.modules.get("zhenxun.builtin_plugins.hooks.auth_checker") + if module is None: + return False + hooks_active = int(getattr(module, "HOOKS_ACTIVE_COUNT", 0) or 0) + db_active = int(getattr(module, "DB_ACTIVE_COUNT", 0) or 0) + return hooks_active > 0 or db_active > 0 + + +async def _has_active_render_work() -> bool: + module = sys.modules.get("zhenxun.services.renderer.engine") + if module is None: + return False + manager = getattr(module, "engine_manager", None) + engine = getattr(manager, "_instance", None) + if engine is None: + return False + try: + snapshot = await asyncio.wait_for(engine.get_runtime_snapshot(), timeout=1.0) + except Exception: + return True + if snapshot.get("active_renders", 0): + return True + if snapshot.get("htmlrender_active_tasks", 0): + return True + active_generation = snapshot.get("active_generation") + if isinstance(active_generation, dict) and active_generation.get( + "active_leases", 0 + ): + return True + retiring = snapshot.get("retiring_generations", []) + if isinstance(retiring, list): + return any( + isinstance(item, dict) and item.get("active_leases", 0) for item in retiring + ) + return False + + +async def _run_reclaim() -> None: + global _last_reclaim_at + start = time.monotonic() + before_rss = _get_total_rss() + cleared: dict[str, Any] = {} + cache_stats_before = { + "cache_dict": CacheDict.stats_all(), + "cache_list": CacheList.stats_all(), + "bounded_ttl": await BoundedTTLCache.stats_all(), + } + + cleared["statistics"] = await _flush_statistics_buffer() + cleared["user_gold_logs"] = await _flush_user_gold_log_buffer() + cleared["bounded_ttl_clear"] = await BoundedTTLCache.clear_all() + cleared["cache_dict_clear"] = CacheDict.clear_all() + cleared["cache_list_clear"] = CacheList.clear_all() + cleared["runtime_negative"] = _clear_runtime_negative_caches() + cleared["auth_local"] = _clear_auth_local_caches() + cleared["avatar_l1"] = _clear_avatar_memory_cache() + cleared["renderer_runtime"] = await _clear_renderer_runtime_caches() + cleared["message_manager"] = _clear_message_manager_cache() + cleared["aiocache_memory"] = await _clear_simple_memory_backend() + + collected = gc.collect(2) + malloc_trimmed = _malloc_trim() + after_rss = _get_total_rss() + _last_reclaim_at = time.monotonic() + + logger.info( + "idle memory reclaim completed: " + f"cost={time.monotonic() - start:.3f}s " + f"rss_before={_format_bytes(before_rss)} " + f"rss_after={_format_bytes(after_rss)} " + f"gc={collected} malloc_trim={malloc_trimmed} " + f"cleared={cleared} cache_stats_before={cache_stats_before}", + LOG_COMMAND, + ) + + +async def _flush_statistics_buffer() -> int: + module = sys.modules.get("zhenxun.builtin_plugins.statistics.statistics_hook") + if module is None: + return 0 + flush = getattr(module, "_flush_statistics_buffer", None) + if flush is None: + return 0 + result = await flush("内存回收") + return int(result or 0) + + +async def _flush_user_gold_log_buffer() -> int: + module = sys.modules.get("zhenxun.services.buffered_writers") + if module is None: + return 0 + flush = getattr(module, "flush_user_gold_log_buffer", None) + if flush is None: + return 0 + result = await flush("内存回收") + return int(result or 0) + + +async def _clear_simple_memory_backend() -> bool: + backend = getattr(CacheRoot, "_cache_backend", None) + if not isinstance(backend, SimpleMemoryCache): + return False + await backend.clear() + return True + + +def _clear_runtime_negative_caches() -> dict[str, int]: + module = sys.modules.get("zhenxun.services.cache.runtime_cache") + if module is None: + return {} + result: dict[str, int] = {} + for name in ( + "BotMemoryCache", + "GroupMemoryCache", + "LevelUserMemoryCache", + "TaskInfoMemoryCache", + "PluginLimitMemoryCache", + "BanMemoryCache", + ): + cache_cls = getattr(module, name, None) + negative = getattr(cache_cls, "_negative", None) + if isinstance(negative, dict) and negative: + result[name] = len(negative) + negative.clear() + return result + + +def _clear_auth_local_caches() -> dict[str, int]: + module = sys.modules.get("zhenxun.builtin_plugins.hooks.auth_checker") + if module is None: + return {} + result: dict[str, int] = {} + for name in ( + "_MATCHER_COMMAND_TYPE_CACHE", + "_MATCHER_COMMAND_LITERAL_CACHE", + "_MATCHER_ALCONNA_SHORTCUT_CACHE", + ): + cache = getattr(module, name, None) + if isinstance(cache, dict) and cache: + result[name] = len(cache) + cache.clear() + return result + + +def _clear_avatar_memory_cache() -> int: + module = sys.modules.get("zhenxun.services.avatar_service") + if module is None: + return 0 + service = getattr(module, "avatar_service", None) + clear = getattr(service, "clear_memory_cache", None) + if not callable(clear): + return 0 + result = clear() + return result if isinstance(result, int) and result > 0 else 0 + + +async def _clear_renderer_runtime_caches() -> dict[str, int]: + module = sys.modules.get("zhenxun.services.renderer.service") + if module is None: + return {} + service = getattr(module, "renderer_service", None) + clear = getattr(service, "clear_runtime_caches", None) + if not callable(clear): + return {} + result = clear() + if inspect.isawaitable(result): + result = await result + if not isinstance(result, dict): + return {} + return { + str(key): int(value) + for key, value in result.items() + if isinstance(value, int) and value > 0 + } + + +def _clear_message_manager_cache() -> int: + module = sys.modules.get("zhenxun.utils.manager.message_manager") + if module is None: + return 0 + manager_cls = getattr(module, "MessageManager", None) + clear = getattr(manager_cls, "clear_all", None) + if not callable(clear): + return 0 + result = clear() + return result if isinstance(result, int) and result > 0 else 0 + + +def _get_total_rss() -> int | None: + try: + import psutil + + process = psutil.Process() + total = process.memory_info().rss + for child in process.children(recursive=True): + with contextlib.suppress(Exception): + total += child.memory_info().rss + return int(total) + except Exception: + return None + + +def _malloc_trim() -> bool: + if sys.platform.startswith(("win", "darwin")): + return False + try: + import ctypes + + libc = ctypes.CDLL("libc.so.6") + return bool(libc.malloc_trim(0)) + except Exception: + return False + + +def _format_bytes(value: int | None) -> str: + if value is None: + return "unknown" + return f"{value / 1024 / 1024:.2f}MiB" diff --git a/zhenxun/services/message_load.py b/zhenxun/services/message_load.py index 6823ed5b..6b7294d0 100644 --- a/zhenxun/services/message_load.py +++ b/zhenxun/services/message_load.py @@ -3,6 +3,17 @@ from __future__ import annotations import time _OVERLOAD_UNTIL = 0.0 +_LAST_ACTIVITY = time.monotonic() + + +def mark_activity() -> None: + """Record lightweight runtime activity for idle-only maintenance jobs.""" + global _LAST_ACTIVITY + _LAST_ACTIVITY = time.monotonic() + + +def idle_seconds() -> float: + return max(0.0, time.monotonic() - _LAST_ACTIVITY) def signal_overload(duration: float = 5.0) -> None: diff --git a/zhenxun/services/renderer/engine.py b/zhenxun/services/renderer/engine.py index 933d7b9c..1a095490 100644 --- a/zhenxun/services/renderer/engine.py +++ b/zhenxun/services/renderer/engine.py @@ -20,6 +20,12 @@ from zhenxun.services.log import logger from .types import BaseScreenshotEngine _PLAYWRIGHT_DISCONNECT_ERROR = "Connection closed while reading from the driver" +_PLAYWRIGHT_TARGET_CLOSED_ERROR_MARKERS = ( + "TargetClosedError", + "Target page, context or browser has been closed", + "browser has been closed", + "BrowserContext.new_page", +) _UNRETRIEVED_FUTURE_MESSAGE = "Future exception was never retrieved" _LOOP_EXCEPTION_FILTER_STATE_ATTR = "_zhenxun_playwright_exception_filter_state" _DISCONNECT_SUPPRESSION_WINDOW_SECONDS = 10.0 @@ -135,6 +141,14 @@ def _is_ignorable_playwright_disconnect(ctx: dict[str, Any]) -> bool: ) +def _is_playwright_target_closed_error(exc: Exception) -> bool: + exc_name = type(exc).__name__ + if exc_name == "TargetClosedError": + return True + message = str(exc) + return any(marker in message for marker in _PLAYWRIGHT_TARGET_CLOSED_ERROR_MARKERS) + + def _get_loop_exception_filter_state( loop: asyncio.AbstractEventLoop, ) -> dict[str, Any] | None: @@ -1103,29 +1117,54 @@ class PlaywrightEngine(BaseScreenshotEngine): template_path: str, render_options: dict[str, Any], ) -> bytes: - generation, context = await self._acquire_context() - page = None - broken = False - try: - page = await context.new_page() - page_options = self._build_page_options(render_options, pooled=True) - viewport = page_options.get("viewport") - if isinstance(viewport, dict): - width = viewport.get("width") - height = viewport.get("height") - if isinstance(width, int) and isinstance(height, int): - await page.set_viewport_size({"width": width, "height": height}) - return await self._render_with_page( - page, html, template_path, render_options - ) - except Exception: - broken = True - raise - finally: - if page is not None: - with contextlib.suppress(Exception): - await page.close() - await self._release_context(generation, context, broken=broken) + last_error: Exception | None = None + for attempt in range(2): + generation, context = await self._acquire_context() + page = None + broken = False + try: + page = await context.new_page() + page_options = self._build_page_options(render_options, pooled=True) + viewport = page_options.get("viewport") + if isinstance(viewport, dict): + width = viewport.get("width") + height = viewport.get("height") + if isinstance(width, int) and isinstance(height, int): + await page.set_viewport_size({"width": width, "height": height}) + return await self._render_with_page( + page, html, template_path, render_options + ) + except Exception as e: + broken = True + last_error = e + if attempt == 0: + if _is_playwright_target_closed_error(e): + logger.warning( + "截图引擎浏览器上下文代已失效,切换新代后重试一次。", + "PlaywrightEngine", + e=e, + ) + try: + await self._swap_generation("target_closed") + except Exception: + raise e + else: + logger.warning( + "截图引擎上下文已失效,丢弃后重试一次。", + "PlaywrightEngine", + e=e, + ) + continue + raise + finally: + if page is not None: + with contextlib.suppress(Exception): + await page.close() + await self._release_context(generation, context, broken=broken) + + if last_error is not None: + raise last_error + raise RuntimeError("截图引擎上下文池渲染失败。") async def _render_html( self, diff --git a/zhenxun/services/renderer/result_cache.py b/zhenxun/services/renderer/result_cache.py index 75fa5a08..51a68a99 100644 --- a/zhenxun/services/renderer/result_cache.py +++ b/zhenxun/services/renderer/result_cache.py @@ -1,11 +1,9 @@ from __future__ import annotations -import asyncio -from collections import OrderedDict import hashlib -import time from typing import Any +from zhenxun.services.cache.bounded_ttl import BoundedTTLCache from zhenxun.utils.pydantic_compat import dump_json_safely @@ -23,9 +21,12 @@ class RenderResultMemoryCache: if isinstance(max_total_bytes, int) and max_total_bytes > 0 else None ) - self._cache: OrderedDict[str, tuple[float, bytes]] = OrderedDict() - self._total_bytes = 0 - self._lock = asyncio.Lock() + self._cache = BoundedTTLCache[str, bytes]( + "RENDER_RESULT", + ttl_seconds=self._ttl_seconds, + max_items=self._max_items, + max_total_bytes=self._max_total_bytes, + ) @staticmethod def build_key(payload: Any) -> str: @@ -37,55 +38,8 @@ class RenderResultMemoryCache: ) return hashlib.sha256(payload_text.encode("utf-8")).hexdigest() - def _pop_oldest(self) -> None: - if not self._cache: - return - _, (_, value) = self._cache.popitem(last=False) - self._total_bytes -= len(value) - if self._total_bytes < 0: - self._total_bytes = 0 - - def _cleanup(self, now: float) -> None: - while self._cache: - expire_at, _ = next(iter(self._cache.values())) - if expire_at > now: - break - self._pop_oldest() - while len(self._cache) > self._max_items: - self._pop_oldest() - if self._max_total_bytes is not None: - while self._total_bytes > self._max_total_bytes and self._cache: - self._pop_oldest() - async def get(self, key: str) -> bytes | None: - now = time.monotonic() - async with self._lock: - self._cleanup(now) - item = self._cache.get(key) - if item is None: - return None - expire_at, value = item - if expire_at <= now: - removed = self._cache.pop(key, None) - if removed: - self._total_bytes -= len(removed[1]) - if self._total_bytes < 0: - self._total_bytes = 0 - return None - self._cache.move_to_end(key) - return value + return await self._cache.get(key) async def set(self, key: str, value: bytes) -> None: - value_size = len(value) - if self._max_total_bytes is not None and value_size > self._max_total_bytes: - return - now = time.monotonic() - async with self._lock: - if old := self._cache.pop(key, None): - self._total_bytes -= len(old[1]) - if self._total_bytes < 0: - self._total_bytes = 0 - self._cache[key] = (now + self._ttl_seconds, value) - self._total_bytes += value_size - self._cache.move_to_end(key) - self._cleanup(now) + await self._cache.set(key, value) diff --git a/zhenxun/services/renderer/service.py b/zhenxun/services/renderer/service.py index 2bc320cd..40d9f7d1 100644 --- a/zhenxun/services/renderer/service.py +++ b/zhenxun/services/renderer/service.py @@ -475,6 +475,18 @@ class RendererService: raise RuntimeError("ThemeManager尚未初始化。") return self._theme_manager.list_available_themes() + def clear_runtime_caches(self) -> dict[str, int]: + cleared: dict[str, int] = {} + if self._theme_manager: + cleared.update(self._theme_manager.clear_runtime_caches()) + if self._template_engine and self._template_engine.env.cache: + jinja_cache = self._template_engine.env.cache + cache_size = len(jinja_cache) + jinja_cache.clear() + if cache_size: + cleared["jinja_env"] = cache_size + return cleared + async def switch_theme(self, theme_name: str) -> str: """ 切换UI主题,加载新主题并持久化配置。 diff --git a/zhenxun/services/renderer/theme.py b/zhenxun/services/renderer/theme.py index 2793ac7e..fc93229d 100644 --- a/zhenxun/services/renderer/theme.py +++ b/zhenxun/services/renderer/theme.py @@ -51,8 +51,10 @@ class ManifestRegistry: self._manifest_cache: dict[str, TemplateManifest] = {} self._lock = asyncio.Lock() - def clear_cache(self): + def clear_cache(self) -> int: + size = len(self._manifest_cache) self._manifest_cache.clear() + return size async def get_manifest( self, component_path: str, skin: str | None = None @@ -362,6 +364,21 @@ class ThemeManager: tuple[type, str, str | None], ComponentDependency ] = OrderedDict() + def clear_runtime_caches(self) -> dict[str, int]: + cleared = { + "asset_resolution": len(self._asset_resolution_cache), + "global_template": len(self._global_template_cache), + "component_dependency": len(self._component_dependency_cache), + } + self._asset_resolution_cache.clear() + self._global_template_cache.clear() + self._component_dependency_cache.clear() + if self.manifest_registry: + manifest_count = self.manifest_registry.clear_cache() + if manifest_count: + cleared["manifest"] = manifest_count + return {key: value for key, value in cleared.items() if value} + @staticmethod def _get_lru_entry(cache: OrderedDict, key: Any) -> Any: value = cache.get(key) diff --git a/zhenxun/services/runtime_bootstrap.py b/zhenxun/services/runtime_bootstrap.py index 3b12487b..a813fea9 100644 --- a/zhenxun/services/runtime_bootstrap.py +++ b/zhenxun/services/runtime_bootstrap.py @@ -2,10 +2,17 @@ import asyncio from concurrent.futures import ThreadPoolExecutor import contextlib import os +import signal import anyio.to_thread -from nonebot.drivers import Driver +from zhenxun.services.log import logger +from zhenxun.services.memory_governor import ( + start_memory_governor, + stop_memory_governor, +) +from zhenxun.services.send_queue import start_send_queue, stop_send_queue +from zhenxun.services.uninfo_patch import apply_uninfo_onebot11_patch from zhenxun.utils.manager.priority_manager import PriorityLifecycle DEFAULT_EXECUTOR_MIN_WORKERS = 16 @@ -14,6 +21,7 @@ DEFAULT_ANYIO_MIN_TOKENS = 32 DEFAULT_ANYIO_MAX_TOKENS = 128 _thread_executor: ThreadPoolExecutor | None = None +_launcher_watchdog_task: asyncio.Task[None] | None = None _runtime_hooks_registered = False _alconna_patch_applied = False @@ -57,14 +65,60 @@ def _apply_alconna_conflict_patch() -> None: _alconna_patch_applied = True -def register_runtime_bootstrap(driver: Driver) -> None: +async def _launcher_watchdog_loop(launcher_pid: int) -> None: + try: + import psutil + except Exception: + return + current_pid = os.getpid() + while True: + await asyncio.sleep(2) + if psutil.pid_exists(launcher_pid): + continue + logger.warning( + f"检测到 launcher 进程 {launcher_pid} 已退出,worker 将主动结束...", + "RuntimeBootstrap", + ) + with contextlib.suppress(Exception): + os.kill(current_pid, signal.SIGTERM) + return + + +def _start_launcher_watchdog() -> None: + global _launcher_watchdog_task + if _launcher_watchdog_task is not None and not _launcher_watchdog_task.done(): + return + launcher_pid_text = os.getenv("ZHENXUN_LAUNCHER_PID", "").strip() + if not launcher_pid_text: + return + with contextlib.suppress(ValueError): + launcher_pid = int(launcher_pid_text) + if launcher_pid > 0: + _launcher_watchdog_task = asyncio.create_task( + _launcher_watchdog_loop(launcher_pid) + ) + + +async def _stop_launcher_watchdog() -> None: + global _launcher_watchdog_task + task = _launcher_watchdog_task + _launcher_watchdog_task = None + if task is None or task.done(): + return + task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await task + + +def register_runtime_bootstrap(_driver) -> None: _apply_alconna_conflict_patch() + apply_uninfo_onebot11_patch() global _runtime_hooks_registered if _runtime_hooks_registered: return _runtime_hooks_registered = True - @driver.on_startup + @PriorityLifecycle.on_startup(priority=-100) async def _setup_runtime_concurrency() -> None: global _thread_executor workers = _get_executor_workers() @@ -77,10 +131,16 @@ def register_runtime_bootstrap(driver: Driver) -> None: with contextlib.suppress(Exception): limiter = anyio.to_thread.current_default_thread_limiter() limiter.total_tokens = _get_anyio_tokens(workers) + _start_launcher_watchdog() + await start_send_queue() + await start_memory_governor() @PriorityLifecycle.on_shutdown(priority=50) async def _shutdown_runtime_concurrency() -> None: global _thread_executor + await _stop_launcher_watchdog() + await stop_send_queue() + await stop_memory_governor() executor = _thread_executor _thread_executor = None if executor is not None: diff --git a/zhenxun/services/send_queue.py b/zhenxun/services/send_queue.py index 8ba5cc15..e381da2d 100644 --- a/zhenxun/services/send_queue.py +++ b/zhenxun/services/send_queue.py @@ -2,21 +2,29 @@ import asyncio import time from typing import Any -import nonebot from nonebot.adapters import Bot from zhenxun.services.log import logger -_SEND_APIS = {"send_msg", "send_like"} -_QUEUE: asyncio.Queue[tuple[Bot, str, dict[str, Any], asyncio.Future]] = asyncio.Queue() +_SEND_APIS = {"send_msg", "send_group_msg", "send_private_msg", "send_like"} _WORKERS = 3 _MIN_INTERVAL = 0.05 +_QUEUE_MAXSIZE = 2000 +_SHUTDOWN_DRAIN_TIMEOUT_SECONDS = 3.0 +_QUEUE_PRESSURE_LOG_INTERVAL = 10.0 +_QUEUE: asyncio.Queue[tuple[Bot, str, dict[str, Any], asyncio.Future[Any]]] = ( + asyncio.Queue(maxsize=_QUEUE_MAXSIZE) +) _SEND_LOCK = asyncio.Lock() _LAST_SEND_TS = 0.0 _API_SEMAPHORE = asyncio.Semaphore(3) _ORIG_CALL_API = Bot.call_api _PATCHED = False _WORKER_TASKS: list[asyncio.Task] = [] +_QUEUE_TIMEOUT_COUNT = 0 +_SEND_LIKE_DROP_COUNT = 0 +_LAST_QUEUE_PRESSURE_LOG = 0.0 +_STOPPING = False async def _rate_limit(): @@ -29,15 +37,36 @@ async def _rate_limit(): _LAST_SEND_TS = time.monotonic() +def _log_queue_pressure(reason: str) -> None: + global _LAST_QUEUE_PRESSURE_LOG + now = time.monotonic() + if now - _LAST_QUEUE_PRESSURE_LOG < _QUEUE_PRESSURE_LOG_INTERVAL: + return + _LAST_QUEUE_PRESSURE_LOG = now + logger.warning( + f"{reason}; qsize={_QUEUE.qsize()}/{_QUEUE_MAXSIZE} " + f"timeouts={_QUEUE_TIMEOUT_COUNT} dropped_like={_SEND_LIKE_DROP_COUNT}", + "SendQueue", + ) + + +async def _direct_call_api(bot: Bot, api: str, data: dict[str, Any]) -> Any: + await _rate_limit() + async with _API_SEMAPHORE: + return await _ORIG_CALL_API(bot, api, **data) + + async def _worker(worker_id: int): while True: bot, api, data, future = await _QUEUE.get() try: - await _rate_limit() - async with _API_SEMAPHORE: - result = await _ORIG_CALL_API(bot, api, **data) + result = await _direct_call_api(bot, api, data) if not future.done(): future.set_result(result) + except asyncio.CancelledError: + if not future.done(): + future.set_exception(RuntimeError("send queue worker cancelled")) + raise except Exception as exc: if not future.done(): future.set_exception(exc) @@ -54,12 +83,41 @@ async def _worker(worker_id: int): async def _queued_call_api(self: Bot, api: str, **data: Any): if api not in _SEND_APIS: return await _ORIG_CALL_API(self, api, **data) + if _STOPPING: + return await _direct_call_api(self, api, data) + loop = asyncio.get_running_loop() - future: asyncio.Future = loop.create_future() - await _QUEUE.put((self, api, data, future)) + future: asyncio.Future[Any] = loop.create_future() + queue_item = (self, api, data, future) + try: + _QUEUE.put_nowait(queue_item) + except asyncio.QueueFull: + if api == "send_like": + global _SEND_LIKE_DROP_COUNT + _SEND_LIKE_DROP_COUNT += 1 + _log_queue_pressure("send_like dropped because send queue is full") + return None + global _QUEUE_TIMEOUT_COUNT + _QUEUE_TIMEOUT_COUNT += 1 + _log_queue_pressure(f"{api} fallback to direct send because queue is full") + return await _direct_call_api(self, api, data) return await future +def _drain_pending_futures(reason: str) -> int: + drained = 0 + while True: + try: + _, _, _, future = _QUEUE.get_nowait() + except asyncio.QueueEmpty: + break + if not future.done(): + future.set_exception(RuntimeError(reason)) + _QUEUE.task_done() + drained += 1 + return drained + + def patch_send_queue() -> None: global _PATCHED if _PATCHED: @@ -68,21 +126,41 @@ def patch_send_queue() -> None: _PATCHED = True -driver = nonebot.get_driver() +def unpatch_send_queue() -> None: + global _PATCHED + if not _PATCHED: + return + Bot.call_api = _ORIG_CALL_API # type: ignore[assignment] + _PATCHED = False -@driver.on_startup -async def _start_send_queue(): +async def start_send_queue() -> None: + global _STOPPING patch_send_queue() + _STOPPING = False + if _WORKER_TASKS: + return for idx in range(_WORKERS): _WORKER_TASKS.append(asyncio.create_task(_worker(idx))) -@driver.on_shutdown -async def _stop_send_queue(): +async def stop_send_queue() -> None: + global _STOPPING + _STOPPING = True + try: + await asyncio.wait_for(_QUEUE.join(), timeout=_SHUTDOWN_DRAIN_TIMEOUT_SECONDS) + except asyncio.TimeoutError: + drained = _drain_pending_futures("send queue shutdown before drain completed") + logger.warning( + f"send queue shutdown timed out, dropped pending futures={drained}, " + f"qsize={_QUEUE.qsize()}", + "SendQueue", + ) tasks = _WORKER_TASKS.copy() _WORKER_TASKS.clear() for task in tasks: task.cancel() if tasks: await asyncio.gather(*tasks, return_exceptions=True) + unpatch_send_queue() + _STOPPING = False diff --git a/zhenxun/services/uninfo_patch.py b/zhenxun/services/uninfo_patch.py new file mode 100644 index 00000000..2793f53b --- /dev/null +++ b/zhenxun/services/uninfo_patch.py @@ -0,0 +1,149 @@ +import asyncio +from collections.abc import Awaitable, Callable +import contextlib +from typing import Any, cast + +from nonebot.adapters import Bot, Event +from nonebot.adapters.onebot.v11.event import GroupMessageEvent +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 + + +def _sender_value(sender: Any, key: str, default: Any = None) -> Any: + value = getattr(sender, key, default) + return default if value is None else value + + +def _event_value(event: Event, key: str, default: Any = None) -> Any: + value = getattr(event, key, default) + return default if value is None else value + + +def _event_group_name(event: Event) -> str | None: + group_name = _event_value(event, "group_name") + if isinstance(group_name, str) and group_name: + return group_name + group = _event_value(event, "group") + if group is not None: + name = _sender_value(group, "name") or _sender_value(group, "group_name") + if isinstance(name, str) and name: + return name + return None + + +def _has_compatible_onebot11_sender(event: Event) -> bool: + if getattr(event, "_zx_uninfo_full_fetch", False): + return False + sender = _event_value(event, "sender") + if sender is None: + return False + return ( + _event_value(event, "user_id") is not None + and _event_value(event, "group_id") is not None + and _sender_value(sender, "nickname") is not None + and _sender_value(sender, "role") is not None + ) + + +async def _fast_onebot11_group_message(bot: Bot, event: Event) -> dict[str, Any]: + """Build Uninfo session data from OneBot v11 group message event fields. + + nonebot-plugin-uninfo's default OneBot v11 fetcher always calls + get_group_info and get_group_member_info for group messages. For normal + matcher rule checks, event-provided sender fields are enough and avoid + multiplying protocol API calls by the number of candidate matchers. + """ + + original = _ORIGINAL_ONEBOT11_GROUP_MESSAGE + if not _has_compatible_onebot11_sender(event): + if original is not None: + return await original(bot, event) + logger.debug("Uninfo OneBot11 fast fetch fallback unavailable") + + sender = _event_value(event, "sender") + user_id = str(_event_value(event, "user_id", "")) + group_id = str(_event_value(event, "group_id", "")) + nickname = _sender_value(sender, "nickname", "") + card = _sender_value(sender, "card", "") or nickname + return { + "group_id": group_id, + "group_name": _event_group_name(event), + "user_id": user_id, + "name": nickname, + "nickname": card, + "card": card, + "role": _sender_value(sender, "role", "member"), + "join_time": _event_value(event, "join_time"), + "gender": _sender_value(sender, "sex", "unknown") or "unknown", + } + + +async def _singleflight_fetch(self: Any, bot: Bot, event: Event) -> Any: + original = _ORIGINAL_FETCH + if original is None: + return None + + try: + sess_id = self.get_session_id(event) + except ValueError: + return await original(self, bot, event) + + session_cache = getattr(self, "session_cache", None) + if isinstance(session_cache, dict) and sess_id in session_cache: + return session_cache[sess_id] + + inflight = getattr(self, "_zx_fetch_inflight", None) + if not isinstance(inflight, dict): + inflight = {} + setattr(self, "_zx_fetch_inflight", inflight) + + key = (str(getattr(bot, "self_id", "")), event.__class__, sess_id) + task = inflight.get(key) + if task is None or task.done(): + task = asyncio.ensure_future(original(self, bot, event)) + inflight[key] = task + try: + return await task + finally: + if inflight.get(key) is task and task.done(): + inflight.pop(key, None) + + +def apply_uninfo_onebot11_patch() -> None: + global _ORIGINAL_FETCH, _ORIGINAL_ONEBOT11_GROUP_MESSAGE, _PATCHED + if _PATCHED: + return + + with contextlib.suppress(Exception): + from nonebot_plugin_uninfo.adapters.onebot11.main import fetcher + + original_endpoint = fetcher.endpoint.get(GroupMessageEvent) + if not getattr(original_endpoint, "__zhenxun_fast_onebot11__", False): + _ORIGINAL_ONEBOT11_GROUP_MESSAGE = cast( + Callable[..., Awaitable[dict[str, Any]]] | None, + original_endpoint, + ) + setattr(_fast_onebot11_group_message, "__zhenxun_fast_onebot11__", True) + fetcher.endpoint[GroupMessageEvent] = _fast_onebot11_group_message + + try: + from nonebot_plugin_uninfo.fetch import InfoFetcher + except Exception as e: + logger.warning("Uninfo patch skipped", e=e) + return + + original_fetch = getattr(InfoFetcher, "fetch", None) + if getattr(original_fetch, "__zhenxun_singleflight__", False): + _PATCHED = True + return + if original_fetch is None: + return + + _ORIGINAL_FETCH = cast(Callable[..., Awaitable[Any]], original_fetch) + setattr(_singleflight_fetch, "__zhenxun_singleflight__", True) + setattr(InfoFetcher, "fetch", _singleflight_fetch) + _PATCHED = True + logger.debug("Uninfo OneBot11 fast fetch and singleflight patch applied") diff --git a/zhenxun/utils/enum.py b/zhenxun/utils/enum.py index 2fc94e11..0412ee4a 100644 --- a/zhenxun/utils/enum.py +++ b/zhenxun/utils/enum.py @@ -55,8 +55,6 @@ class CacheType(StrEnum): """全局全部群组""" GROUP_PLUGIN_SETTINGS = "GROUP_PLUGIN_SETTINGS" """插件分群配置""" - GROUP_PLUGIN_SETTINGS_VIEW = "GROUP_PLUGIN_SETTINGS_VIEW" - """插件分群配置视图缓存(聚合 dict)""" USERS = "GLOBAL_ALL_USERS" """全部用户""" BAN = "GLOBAL_ALL_BAN" diff --git a/zhenxun/utils/http_utils.py b/zhenxun/utils/http_utils.py index c2f4b0f1..f568db35 100644 --- a/zhenxun/utils/http_utils.py +++ b/zhenxun/utils/http_utils.py @@ -1,5 +1,4 @@ import asyncio -from collections import OrderedDict from collections.abc import AsyncGenerator, Awaitable, Callable, Sequence from contextlib import asynccontextmanager import os @@ -21,6 +20,7 @@ from rich.progress import ( import ujson as json from zhenxun.configs.config import BotConfig +from zhenxun.services.cache.bounded_ttl import BoundedTTLCache from zhenxun.services.log import logger from zhenxun.utils.decorator.retry import Retry from zhenxun.utils.exception import AllURIsFailedError @@ -168,7 +168,12 @@ class AsyncHttpx: _CONTENT_CACHE_TTL: ClassVar[float] = 3.0 _CONTENT_CACHE_MAX_ITEMS: ClassVar[int] = 256 _CONTENT_CACHE_MAX_BYTES: ClassVar[int] = 2 * 1024 * 1024 - _content_cache: ClassVar[OrderedDict[str, tuple[float, bytes]]] = OrderedDict() + _content_cache: ClassVar[BoundedTTLCache[str, bytes]] = BoundedTTLCache( + "HTTP_IMAGE_CONTENT", + ttl_seconds=_CONTENT_CACHE_TTL, + max_items=_CONTENT_CACHE_MAX_ITEMS, + max_total_bytes=_CONTENT_CACHE_MAX_BYTES, + ) _content_inflight: ClassVar[dict[str, asyncio.Task[Response]]] = {} _content_cache_lock: ClassVar[asyncio.Lock] = asyncio.Lock() @@ -191,29 +196,6 @@ class AsyncHttpx: return True return "qpic.cn" in lower_url or "qlogo.cn" in lower_url - @classmethod - def _get_cached_content_nolock(cls, key: str) -> bytes | None: - entry = cls._content_cache.get(key) - if not entry: - return None - expire_at, content = entry - if expire_at <= time.monotonic(): - cls._content_cache.pop(key, None) - return None - cls._content_cache.move_to_end(key) - return content - - @classmethod - def _cleanup_content_cache_nolock(cls) -> None: - now = time.monotonic() - while cls._content_cache: - expire_at, _ = next(iter(cls._content_cache.values())) - if expire_at > now: - break - cls._content_cache.popitem(last=False) - while len(cls._content_cache) > cls._CONTENT_CACHE_MAX_ITEMS: - cls._content_cache.popitem(last=False) - @classmethod async def _try_cache_content(cls, key: str, response: Response) -> None: content = response.content @@ -223,13 +205,7 @@ class AsyncHttpx: is_image = content_type.startswith("image/") or cls._is_probably_image_url(key) if not is_image: return - async with cls._content_cache_lock: - cls._content_cache[key] = ( - time.monotonic() + cls._CONTENT_CACHE_TTL, - content, - ) - cls._content_cache.move_to_end(key) - cls._cleanup_content_cache_nolock() + await cls._content_cache.set(key, content) @classmethod def _prepare_temporary_client_config(cls, client_kwargs: dict) -> dict: @@ -450,9 +426,11 @@ class AsyncHttpx: return res.content cache_key = url + if cached := await cls._content_cache.get(cache_key): + return cached + async with cls._content_cache_lock: - cached = cls._get_cached_content_nolock(cache_key) - if cached is not None: + if cached := await cls._content_cache.get(cache_key): return cached task = cls._content_inflight.get(cache_key) if task is None: diff --git a/zhenxun/utils/manager/message_manager.py b/zhenxun/utils/manager/message_manager.py index ee34369d..271923f4 100644 --- a/zhenxun/utils/manager/message_manager.py +++ b/zhenxun/utils/manager/message_manager.py @@ -1,25 +1,77 @@ +from collections import OrderedDict +import time from typing import ClassVar class MessageManager: - data: ClassVar[dict[str, list[str]]] = {} + _MAX_USERS: ClassVar[int] = 4096 + _MAX_MESSAGES_PER_USER: ClassVar[int] = 200 + _TRIM_MESSAGES_TO: ClassVar[int] = 100 + _USER_TTL_SECONDS: ClassVar[float] = 6 * 60 * 60 + data: ClassVar[OrderedDict[str, tuple[float, list[str]]]] = OrderedDict() + + @classmethod + def _prune(cls, now: float | None = None) -> None: + now = time.monotonic() if now is None else now + stale_before = now - cls._USER_TTL_SECONDS + stale_uids = [ + uid for uid, (last_seen, _) in cls.data.items() if last_seen <= stale_before + ] + for uid in stale_uids: + cls.data.pop(uid, None) + while len(cls.data) > cls._MAX_USERS: + cls.data.popitem(last=False) + + @classmethod + def _touch(cls, uid: str, messages: list[str], now: float | None = None) -> None: + now = time.monotonic() if now is None else now + cls.data[uid] = (now, messages) + cls.data.move_to_end(uid) @classmethod def add(cls, uid: str, msg_id: str): - if uid not in cls.data: - cls.data[uid] = [] - cls.data[uid].append(msg_id) + now = time.monotonic() + cls._prune(now) + _, messages = cls.data.get(uid, (now, [])) + messages.append(msg_id) + cls._touch(uid, messages, now) cls.remove_check(uid) + cls._prune(now) @classmethod def check(cls, uid: str, msg_id: str) -> bool: - return msg_id in cls.data.get(uid, []) + now = time.monotonic() + cls._prune(now) + entry = cls.data.get(uid) + if entry is None: + return False + _, messages = entry + cls._touch(uid, messages, now) + return msg_id in messages @classmethod def remove_check(cls, uid: str): - if len(cls.data[uid]) > 200: - cls.data[uid] = cls.data[uid][100:] + entry = cls.data.get(uid) + if entry is None: + return + _, messages = entry + if len(messages) > cls._MAX_MESSAGES_PER_USER: + messages = messages[-cls._TRIM_MESSAGES_TO :] + cls._touch(uid, messages) @classmethod def get(cls, uid: str) -> list[str]: - return cls.data[uid] if uid in cls.data else [] + now = time.monotonic() + cls._prune(now) + entry = cls.data.get(uid) + if entry is None: + return [] + _, messages = entry + cls._touch(uid, messages, now) + return list(messages) + + @classmethod + def clear_all(cls) -> int: + size = len(cls.data) + cls.data.clear() + return size diff --git a/zhenxun/utils/manager/virtual_env_package_manager.py b/zhenxun/utils/manager/virtual_env_package_manager.py index e2455a1f..f520ac02 100644 --- a/zhenxun/utils/manager/virtual_env_package_manager.py +++ b/zhenxun/utils/manager/virtual_env_package_manager.py @@ -8,6 +8,7 @@ from zhenxun.configs.config import Config from zhenxun.services.log import logger LOG_COMMAND = "VirtualEnvPackageManager" +PROJECT_ROOT = Path(__file__).resolve().parents[3] Config.add_plugin_config( "virtualenv", @@ -191,6 +192,46 @@ class VirtualEnvPackageManager: ) return stderr + @classmethod + async def add_requirement(cls, requirement_file: Path): + """将依赖文件写入项目依赖并同步环境 + + 插件商店安装依赖需要持久化到 pyproject.toml/uv.lock,避免重建环境后丢失。 + """ + if not requirement_file.exists(): + raise FileNotFoundError(f"依赖文件 {requirement_file} 不存在", LOG_COMMAND) + cls._clean_requirements_file(requirement_file) + try: + command = [ + "uv", + "add", + "--requirements", + str(requirement_file.absolute()), + ] + logger.info(f"执行项目依赖添加指令: {command}", LOG_COMMAND) + result = await asyncio.to_thread( + subprocess.run, + command, + cwd=PROJECT_ROOT, + check=True, + capture_output=True, + text=True, + encoding="utf-8", + errors="replace", + ) + logger.debug( + f"项目依赖添加指令执行完成: {result.stdout}", + LOG_COMMAND, + ) + return result.stdout + except (CalledProcessError, FileNotFoundError) as e: + stderr = e.stderr if isinstance(e, CalledProcessError) else str(e) + logger.error( + f"项目依赖添加指令执行失败: {stderr}.", + LOG_COMMAND, + ) + return stderr + @classmethod async def list(cls) -> str: """列出已安装的依赖包"""