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:
"""列出已安装的依赖包"""