性能优化 (#2126)

* 性能优化

* 代码改进

* 优化浏览器代际切换逻辑

* 统一缓存与生命周期

* 添加aiomysql依赖

* 优化插件路径处理逻辑,简化条件判断;在虚拟环境包管理器中添加编码和错误处理参数以增强稳定性

* 🚨 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>
This commit is contained in:
Copaan
2026-04-26 15:50:15 +08:00
committed by GitHub
co-authored by HibiKier pre-commit-ci[bot]
parent 24c316cd2c
commit 5d92ccd3b0
56 changed files with 3092 additions and 806 deletions
+1
View File
@@ -45,6 +45,7 @@ dependencies = [
"alibabacloud-devops20210625>=5.0.2,<6.0.0", "alibabacloud-devops20210625>=5.0.2,<6.0.0",
"uvloop>=0.21.0; sys_platform != 'win32'", "uvloop>=0.21.0; sys_platform != 'win32'",
"pytest-timeout>=2.4.0", "pytest-timeout>=2.4.0",
"aiomysql>=0.3.2",
] ]
[project.scripts] [project.scripts]
Generated
+23
View File
@@ -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" }, { 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]] [[package]]
name = "aiosignal" name = "aiosignal"
version = "1.4.0" 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" }, { 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]] [[package]]
name = "pypika-tortoise" name = "pypika-tortoise"
version = "0.1.6" version = "0.1.6"
@@ -4157,6 +4178,7 @@ source = { editable = "." }
dependencies = [ dependencies = [
{ name = "aiocache", extra = ["redis"] }, { name = "aiocache", extra = ["redis"] },
{ name = "aiofiles" }, { name = "aiofiles" },
{ name = "aiomysql" },
{ name = "alibabacloud-devops20210625" }, { name = "alibabacloud-devops20210625" },
{ name = "asyncpg" }, { name = "asyncpg" },
{ name = "beautifulsoup4" }, { name = "beautifulsoup4" },
@@ -4212,6 +4234,7 @@ dev = [
requires-dist = [ requires-dist = [
{ name = "aiocache", extras = ["redis"], specifier = ">=0.12.3" }, { name = "aiocache", extras = ["redis"], specifier = ">=0.12.3" },
{ name = "aiofiles", specifier = ">=23.2.1" }, { name = "aiofiles", specifier = ">=23.2.1" },
{ name = "aiomysql", specifier = ">=0.3.2" },
{ name = "alibabacloud-devops20210625", specifier = ">=5.0.2,<6.0.0" }, { name = "alibabacloud-devops20210625", specifier = ">=5.0.2,<6.0.0" },
{ name = "asyncpg", specifier = ">=0.20.0" }, { name = "asyncpg", specifier = ">=0.20.0" },
{ name = "beautifulsoup4", specifier = ">=4.12.3,<5.0.0" }, { name = "beautifulsoup4", specifier = ">=4.12.3,<5.0.0" },
@@ -1,10 +1,9 @@
from nonebot.adapters import Bot from nonebot.adapters import Bot
from zhenxun.models.group_console import GroupConsole from zhenxun.models.group_console import GroupConsole
from zhenxun.services.cache import CacheRoot
from zhenxun.services.cache.runtime_cache import GroupMemoryCache from zhenxun.services.cache.runtime_cache import GroupMemoryCache
from zhenxun.utils.common_utils import CommonUtils 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 zhenxun.utils.platform import PlatformUtils
from .strategy import get_strategy from .strategy import get_strategy
@@ -134,7 +133,6 @@ class PluginManager:
await GroupConsole.bulk_update( await GroupConsole.bulk_update(
update_list, [norm_field, su_field], batch_size=500 update_list, [norm_field, su_field], batch_size=500
) )
await CacheRoot.clear(CacheType.GROUPS)
for group in update_list: for group in update_list:
await GroupMemoryCache.upsert_from_model(group) await GroupMemoryCache.upsert_from_model(group)
@@ -318,7 +316,6 @@ class PluginManager:
status=False status=False
) )
await CacheRoot.clear(CacheType.GROUPS)
await GroupMemoryCache.refresh() await GroupMemoryCache.refresh()
action_str = "醒来" if status else "休眠" action_str = "醒来" if status else "休眠"
@@ -4,12 +4,11 @@ from typing import Any, cast
from zhenxun.models.group_console import GroupConsole from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.task_info import TaskInfo from zhenxun.models.task_info import TaskInfo
from zhenxun.services.cache import CacheRoot
from zhenxun.services.cache.runtime_cache import ( from zhenxun.services.cache.runtime_cache import (
PluginInfoMemoryCache, PluginInfoMemoryCache,
TaskInfoMemoryCache, TaskInfoMemoryCache,
) )
from zhenxun.utils.enum import BlockType, CacheType, PluginType from zhenxun.utils.enum import BlockType, PluginType
class SwitchStrategy(ABC): class SwitchStrategy(ABC):
@@ -135,7 +134,6 @@ class PluginStrategy(SwitchStrategy):
await self.refresh_cache() await self.refresh_cache()
async def refresh_cache(self) -> None: async def refresh_cache(self) -> None:
await CacheRoot.invalidate_cache(CacheType.PLUGINS)
await PluginInfoMemoryCache.refresh() await PluginInfoMemoryCache.refresh()
@@ -10,6 +10,7 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config from zhenxun.configs.config import Config
from zhenxun.configs.utils import PluginExtraData, RegisterConfig from zhenxun.configs.utils import PluginExtraData, RegisterConfig
from zhenxun.models.chat_history import ChatHistory 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.log import logger
from zhenxun.services.message_load import is_overloaded, should_pause_tasks from zhenxun.services.message_load import is_overloaded, should_pause_tasks
from zhenxun.utils.enum import PluginType from zhenxun.utils.enum import PluginType
@@ -47,6 +48,9 @@ _HISTORY_QUEUE: asyncio.Queue[ChatHistory] = asyncio.Queue(maxsize=5000)
_DROP_COUNT = 0 _DROP_COUNT = 0
_LAST_DROP_LOG = 0.0 _LAST_DROP_LOG = 0.0
_DROP_LOG_INTERVAL = 10.0 _DROP_LOG_INTERVAL = 10.0
_FLUSH_BATCH_SIZE = 200
_FLUSH_MAX_PER_TICK = 1000
_FLUSH_DB_TIMEOUT = 5.0
@chat_history.handle() @chat_history.handle()
@@ -80,19 +84,34 @@ async def _(message: UniMsg, session: Uninfo):
@scheduler.scheduled_job( @scheduler.scheduled_job(
"interval", "interval",
minutes=1, minutes=1,
max_instances=1,
coalesce=True,
) )
async def _(): async def _():
try: try:
if should_pause_tasks(): if should_pause_tasks():
return return
message_list: list[ChatHistory] = [] flushed = 0
while True: while flushed < _FLUSH_MAX_PER_TICK:
try: message_list: list[ChatHistory] = []
message_list.append(_HISTORY_QUEUE.get_nowait()) limit = min(_FLUSH_BATCH_SIZE, _FLUSH_MAX_PER_TICK - flushed)
except asyncio.QueueEmpty: for _ in range(limit):
try:
message_list.append(_HISTORY_QUEUE.get_nowait())
except asyncio.QueueEmpty:
break
if not message_list:
break break
if message_list: await with_db_timeout(
await ChatHistory.bulk_create(message_list) ChatHistory.bulk_create(message_list, _FLUSH_BATCH_SIZE),
logger.debug(f"批量添加聊天记录 {len(message_list)} 条", "定时任务") 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: except Exception as e:
logger.warning("存储聊天记录失败", "chat_history", e=e) logger.warning("存储聊天记录失败", "chat_history", e=e)
@@ -7,9 +7,10 @@ from zhenxun.models.level_user import LevelUser
from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.cache.runtime_cache import LevelUserMemoryCache, LevelUserSnapshot from zhenxun.services.cache.runtime_cache import LevelUserMemoryCache, LevelUserSnapshot
from zhenxun.services.log import logger 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 .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import SkipPluginException from .exception import SkipPluginException
@@ -20,6 +21,9 @@ async def auth_admin(
LevelUser | LevelUserSnapshot | None, LevelUser | LevelUserSnapshot | None LevelUser | LevelUserSnapshot | None, LevelUser | LevelUserSnapshot | None
] ]
| None = None, | None = None,
*,
context: PermissionContext | None = None,
entity: EntityIDs | None = None,
): ):
"""管理员命令 个人权限 """管理员命令 个人权限
@@ -33,7 +37,12 @@ async def auth_admin(
return return
try: 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 global_user: LevelUser | LevelUserSnapshot | None = None
group_users: LevelUser | LevelUserSnapshot | None = None group_users: LevelUser | LevelUserSnapshot | None = None
@@ -42,7 +51,7 @@ async def auth_admin(
global_user, group_users = cached_levels global_user, group_users = cached_levels
else: else:
global_user, group_users = await LevelUserMemoryCache.get_levels( 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 user_level = global_user.user_level if global_user else 0
@@ -53,7 +62,7 @@ async def auth_admin(
raise SkipPluginException( raise SkipPluginException(
f"{plugin.name}({plugin.module}) 管理员权限不足...", f"{plugin.name}({plugin.module}) 管理员权限不足...",
tip_message=[ tip_message=[
At(flag="user", target=session.user.id), At(flag="user", target=entity.user_id),
f"你的权限不足喔,该功能需要的权限等级: {plugin.admin_level}", f"你的权限不足喔,该功能需要的权限等级: {plugin.admin_level}",
], ],
tip_check_tag=entity.user_id, tip_check_tag=entity.user_id,
+7 -48
View File
@@ -1,4 +1,3 @@
import asyncio
import time import time
from nonebot.matcher import Matcher from nonebot.matcher import Matcher
@@ -8,7 +7,6 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config from zhenxun.configs.config import Config
from zhenxun.models.ban_console import BanConsole from zhenxun.models.ban_console import BanConsole
from zhenxun.models.plugin_info import PluginInfo 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.cache.runtime_cache import BanMemoryCache
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger 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 zhenxun.utils.utils import EntityIDs, get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import SkipPluginException from .exception import SkipPluginException
from .utils import freq from .utils import freq
@@ -25,37 +24,6 @@ Config.add_plugin_config(
"才不会给你发消息.", "才不会给你发消息.",
help="对被ban用户发送的消息", 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: async def calculate_ban_time(ban_record: BanConsole | None) -> int:
@@ -214,6 +182,7 @@ async def auth_ban(
session: Uninfo, session: Uninfo,
plugin: PluginInfo, plugin: PluginInfo,
*, *,
context: PermissionContext | None = None,
entity: EntityIDs | None = None, entity: EntityIDs | None = None,
is_superuser: bool = False, is_superuser: bool = False,
) -> None: ) -> None:
@@ -229,28 +198,18 @@ async def auth_ban(
return return
if not matcher.plugin_name: if not matcher.plugin_name:
return return
if context is not None:
entity = context.entity
is_superuser = context.is_superuser
if entity is None: if entity is None:
entity = get_entity_ids(session) entity = get_entity_ids(session)
if is_superuser: if is_superuser:
return return
if entity.group_id: if entity.group_id:
try: await group_handle(entity.group_id)
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)
# 超时时不阻塞,继续执行
if entity.user_id: if entity.user_id:
try: await user_handle(plugin, entity, session)
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)
# 超时时不阻塞,继续执行
finally: finally:
# 记录总执行时间 # 记录总执行时间
elapsed = time.time() - start_time elapsed = time.time() - start_time
@@ -7,6 +7,7 @@ from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils from zhenxun.utils.common_utils import CommonUtils
from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import SkipPluginException from .exception import SkipPluginException
@@ -16,6 +17,8 @@ async def auth_bot(
bot_data: BotConsole | BotSnapshot | None = None, bot_data: BotConsole | BotSnapshot | None = None,
skip_fetch: bool = False, skip_fetch: bool = False,
allow_sleep_bypass: bool = False, allow_sleep_bypass: bool = False,
*,
context: PermissionContext | None = None,
): ):
"""bot层面的权限检查 """bot层面的权限检查
@@ -30,6 +33,9 @@ async def auth_bot(
start_time = time.time() start_time = time.time()
try: try:
if context is not None:
bot_id = context.event.bot_id
bot_data = context.bot_data
bot: BotConsole | BotSnapshot | None = bot_data bot: BotConsole | BotSnapshot | None = bot_data
if bot is None and not skip_fetch: if bot is None and not skip_fetch:
bot = await BotMemoryCache.get(bot_id) bot = await BotMemoryCache.get(bot_id)
@@ -7,13 +7,18 @@ from zhenxun.models.user_console import UserConsole
from zhenxun.services.log import logger from zhenxun.services.log import logger
from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import SkipPluginException from .exception import SkipPluginException
DEFAULT_GOLD = 100 DEFAULT_GOLD = 100
async def auth_cost( async def auth_cost(
user: UserConsole | None, plugin: PluginInfo, session: Uninfo user: UserConsole | None,
plugin: PluginInfo,
session: Uninfo,
*,
context: PermissionContext | None = None,
) -> int: ) -> int:
"""检测是否满足金币条件 """检测是否满足金币条件
@@ -28,6 +33,8 @@ async def auth_cost(
start_time = time.time() start_time = time.time()
try: try:
if context is not None and user is None:
user = context.user
user_gold = user.gold if user else DEFAULT_GOLD user_gold = user.gold if user else DEFAULT_GOLD
if user_gold < plugin.cost_gold: if user_gold < plugin.cost_gold:
"""插件消耗金币不足""" """插件消耗金币不足"""
@@ -7,6 +7,7 @@ from zhenxun.services.cache.runtime_cache import GroupSnapshot
from zhenxun.services.log import logger from zhenxun.services.log import logger
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
from .context import PermissionContext
from .exception import SkipPluginException from .exception import SkipPluginException
_GROUP_WAKE_PATTERN = re.compile(r"^醒来$", re.IGNORECASE) _GROUP_WAKE_PATTERN = re.compile(r"^醒来$", re.IGNORECASE)
@@ -34,6 +35,8 @@ async def auth_group(
group: GroupConsole | GroupSnapshot | None, group: GroupConsole | GroupSnapshot | None,
text: str | None, text: str | None,
group_id: str | None, group_id: str | None,
*,
context: PermissionContext | None = None,
): ):
"""群黑名单检测 群总开关检测 """群黑名单检测 群总开关检测
@@ -42,6 +45,11 @@ async def auth_group(
group: GroupConsole group: GroupConsole
message: UniMsg 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: if not group_id:
return return
@@ -19,9 +19,10 @@ from zhenxun.utils.limiters import CountLimiter, FreqLimiter, UserBlockLimiter
from zhenxun.utils.manager.priority_manager import PriorityLifecycle from zhenxun.utils.manager.priority_manager import PriorityLifecycle
from zhenxun.utils.message import MessageUtils from zhenxun.utils.message import MessageUtils
from zhenxun.utils.time_utils import TimeUtils 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 .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import SkipPluginException from .exception import SkipPluginException
driver = nonebot.get_driver() driver = nonebot.get_driver()
@@ -31,7 +32,7 @@ _LIMIT_NOTICE_LIMITER = FreqLimiter(_LIMIT_NOTICE_CD)
_LIMIT_NOTICE_TASKS: set[asyncio.Task] = set() _LIMIT_NOTICE_TASKS: set[asyncio.Task] = set()
@PriorityLifecycle.on_startup(priority=5) @PriorityLifecycle.on_startup(priority=7)
async def _(): async def _():
"""初始化限制""" """初始化限制"""
await LimitManager.init_limit() await LimitManager.init_limit()
@@ -117,6 +118,7 @@ class LimitManager:
cls.cd_limit = {} cls.cd_limit = {}
cls.block_limit = {} cls.block_limit = {}
cls.count_limit = {} cls.count_limit = {}
cls.module_limit_cache.clear()
# 添加新数据 # 添加新数据
for limit in limit_list: for limit in limit_list:
cls.add_limit(limit) cls.add_limit(limit)
@@ -137,22 +139,22 @@ class LimitManager:
""" """
if limit.module not in cls.add_module: if limit.module not in cls.add_module:
cls.add_module.append(limit.module) cls.add_module.append(limit.module)
if limit.limit_type == PluginLimitType.BLOCK: if limit.limit_type == PluginLimitType.BLOCK:
cls.block_limit[limit.module] = Limit( cls.block_limit[limit.module] = Limit(
limit=limit, limiter=UserBlockLimiter() limit=limit, limiter=UserBlockLimiter()
) )
elif limit.limit_type == PluginLimitType.CD: elif limit.limit_type == PluginLimitType.CD:
cd_value = int(limit.cd or 0) cd_value = int(limit.cd or 0)
cls.cd_limit[limit.module] = Limit( cls.cd_limit[limit.module] = Limit(
limit=limit, limiter=FreqLimiter(cd_value) limit=limit, limiter=FreqLimiter(cd_value)
) )
elif limit.limit_type == PluginLimitType.COUNT: elif limit.limit_type == PluginLimitType.COUNT:
max_count = int(limit.max_count or 0) max_count = int(limit.max_count or 0)
if max_count <= 0: if max_count <= 0:
return return
cls.count_limit[limit.module] = Limit( cls.count_limit[limit.module] = Limit(
limit=limit, limiter=CountLimiter(max_count) limit=limit, limiter=CountLimiter(max_count)
) )
@classmethod @classmethod
def unblock( def unblock(
@@ -322,14 +324,23 @@ class LimitManager:
limiter.increase(key_type) 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 plugin: PluginInfo
session: Uninfo 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: try:
await asyncio.wait_for( await asyncio.wait_for(
LimitManager.check( LimitManager.check(
@@ -10,6 +10,7 @@ from zhenxun.services.log import logger
from zhenxun.utils.enum import BlockType from zhenxun.utils.enum import BlockType
from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import IsSuperuserException, SkipPluginException from .exception import IsSuperuserException, SkipPluginException
from .utils import freq, is_poke from .utils import freq, is_poke
@@ -107,11 +108,16 @@ class GroupCheck:
class PluginCheck: class PluginCheck:
def __init__( 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.session = session
self.is_poke = is_poke self.is_poke = is_poke
self.group_data = group self.group_data = group
self.user_id = user_id or session.user.id
self.group_id = None self.group_id = None
if group: if group:
self.group_id = group.group_id self.group_id = group.group_id
@@ -126,13 +132,11 @@ class PluginCheck:
IgnoredException: 忽略插件 IgnoredException: 忽略插件
""" """
if plugin.block_type == BlockType.PRIVATE: if plugin.block_type == BlockType.PRIVATE:
should_tip = freq.is_send_limit_message( should_tip = freq.is_send_limit_message(plugin, self.user_id, self.is_poke)
plugin, self.session.user.id, self.is_poke
)
raise SkipPluginException( raise SkipPluginException(
f"{plugin.name}({plugin.module}) 该插件在私聊中已被禁用...", f"{plugin.name}({plugin.module}) 该插件在私聊中已被禁用...",
tip_message="该功能在私聊中已被禁用..." if should_tip else None, 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, tip_background=should_tip,
) )
@@ -153,7 +157,7 @@ class PluginCheck:
if self.group_data and self.group_data.is_super: if self.group_data and self.group_data.is_super:
raise IsSuperuserException() 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) should_tip = freq.is_send_limit_message(plugin, sid, self.is_poke)
raise SkipPluginException( raise SkipPluginException(
f"{plugin.name}({plugin.module}) 全局未开启此功能...", f"{plugin.name}({plugin.module}) 全局未开启此功能...",
@@ -176,7 +180,9 @@ async def auth_plugin(
session: Uninfo, session: Uninfo,
event: Event, event: Event,
*, *,
context: PermissionContext | None = None,
skip_group_block: bool = False, skip_group_block: bool = False,
user_id: str | None = None,
): ):
"""插件状态 """插件状态
@@ -187,8 +193,11 @@ async def auth_plugin(
""" """
start_time = time.time() start_time = time.time()
try: try:
if context is not None:
group = context.group or group
user_id = context.user_id
is_poke_event = is_poke(event) 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: if group:
block_set, super_block_set = _get_group_block_sets(group) block_set, super_block_set = _get_group_block_sets(group)
@@ -3,6 +3,7 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config from zhenxun.configs.config import Config
from .context import PermissionContext
from .exception import SkipPluginException from .exception import SkipPluginException
Config.add_plugin_config( 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 """过滤bot调用bot
参数: 参数:
@@ -26,10 +32,13 @@ def bot_filter(session: Uninfo):
""" """
if not Config.get_config("hook", "FILTER_BOT"): if not Config.get_config("hook", "FILTER_BOT"):
return return
if context is not None:
user_id = context.user_id
bot_ids = list(nonebot.get_bots().keys()) 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 return
if session.user.id in bot_ids: if checked_user_id in bot_ids:
raise SkipPluginException( raise SkipPluginException(
f"bot:{session.self_id} 尝试调用 bot:{session.user.id}" f"bot:{session.self_id} 尝试调用 bot:{checked_user_id}"
) )
@@ -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
+427 -181
View File
@@ -1,7 +1,7 @@
import asyncio import asyncio
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
import contextlib import contextlib
import os import importlib
import re import re
import time import time
from typing import cast from typing import cast
@@ -11,7 +11,6 @@ from nonebot.adapters import Bot, Event
from nonebot.exception import IgnoredException from nonebot.exception import IgnoredException
from nonebot.matcher import Matcher from nonebot.matcher import Matcher
import nonebot.message as nb_message import nonebot.message as nb_message
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.utils import PluginExtraData 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.enum import BlockType, GoldHandle, PluginType
from zhenxun.utils.exception import InsufficientGold from zhenxun.utils.exception import InsufficientGold
from zhenxun.utils.platform import PlatformUtils from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.utils import get_entity_ids
from .auth.auth_admin import auth_admin from .auth.auth_admin import auth_admin
from .auth.auth_ban import auth_ban 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.auth_plugin import auth_plugin
from .auth.bot_filter import bot_filter from .auth.bot_filter import bot_filter
from .auth.config import LOGGER_COMMAND, WARNING_THRESHOLD 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 ( from .auth.exception import (
IsSuperuserException, IsSuperuserException,
PermissionExemption, PermissionExemption,
@@ -53,7 +61,6 @@ from .auth.utils import send_message
AUTH_HOOKS_CONCURRENCY_LIMIT = 5 AUTH_HOOKS_CONCURRENCY_LIMIT = 5
AUTH_DB_CONCURRENCY_LIMIT = 6 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 HOOKS_CONCURRENCY_LIMIT = AUTH_HOOKS_CONCURRENCY_LIMIT
DB_CONCURRENCY_LIMIT = AUTH_DB_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_LOCK = asyncio.Lock()
_ROUTE_INDEX_READY = False _ROUTE_INDEX_READY = False
@@ -91,21 +91,17 @@ MATCHER_ROUTE_PREFILTER_TTL = 2
PREFILTER_STATS_LOG_INTERVAL = 10.0 PREFILTER_STATS_LOG_INTERVAL = 10.0
CACHE_SWEEP_INTERVAL = 1.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_ACTIVE_COUNT = 0
HOOKS_SEMAPHORE = asyncio.Semaphore(HOOKS_CONCURRENCY_LIMIT) 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_SEMAPHORE = asyncio.Semaphore(DB_CONCURRENCY_LIMIT)
DB_ACTIVE_COUNT = 0 DB_ACTIVE_COUNT = 0
_CHECK_MATCHER_PATCHED = False _CHECK_MATCHER_PATCHED = False
_ORIGINAL_CHECK_AND_RUN_MATCHER: Callable[..., Awaitable[None]] | None = None _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_TYPE_CACHE: dict[type[Matcher], bool] = {}
_MATCHER_COMMAND_LITERAL_CACHE: dict[type[Matcher], tuple[str, ...] | None] = {} _MATCHER_COMMAND_LITERAL_CACHE: dict[type[Matcher], tuple[str, ...] | None] = {}
_MATCHER_ALCONNA_SHORTCUT_CACHE: dict[type[Matcher], bool] = {} _MATCHER_ALCONNA_SHORTCUT_CACHE: dict[type[Matcher], bool] = {}
@@ -115,6 +111,10 @@ _CHECK_MATCHER_ROUTE_CACHE = CacheDict(
_PREFILTER_STATS = { _PREFILTER_STATS = {
"checked": 0, "checked": 0,
"skipped": 0, "skipped": 0,
"before_task_checked": 0,
"before_task_skipped": 0,
"inside_task_checked": 0,
"inside_task_skipped": 0,
"type_miss": 0, "type_miss": 0,
"route_miss": 0, "route_miss": 0,
"command_miss": 0, "command_miss": 0,
@@ -163,33 +163,6 @@ def _debug_log(message: str, *args, **kwargs) -> None:
logger.debug(message, *args, **kwargs) 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: def _normalize_command(command: str) -> str:
text = command.strip() text = command.strip()
if not text: if not text:
@@ -488,6 +461,9 @@ def _event_plain_text(event: Event) -> str:
def _state_plain_text(state: dict | None) -> str: def _state_plain_text(state: dict | None) -> str:
if state is None: if state is None:
return "" return ""
context = get_event_context(state)
if context is not None:
return context.plain_text.strip()
text = state.get("_zx_plain_text") text = state.get("_zx_plain_text")
if isinstance(text, str): if isinstance(text, str):
return text.strip() 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]: def _get_route_modules_for_event(event: Event, state: dict | None = None) -> set[str]:
if state is not None: 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") route_modules = state.get("_zx_route_modules")
if isinstance(route_modules, set): if isinstance(route_modules, set):
return route_modules 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)) route_modules = _match_route_modules(_event_plain_text(event))
_CHECK_MATCHER_ROUTE_CACHE[key] = route_modules _CHECK_MATCHER_ROUTE_CACHE[key] = route_modules
if state is not None: 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 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 global _PREFILTER_LAST_LOG
_PREFILTER_STATS["checked"] += 1 _PREFILTER_STATS["checked"] += 1
if skipped: if skipped:
_PREFILTER_STATS["skipped"] += 1 _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": if reason == "type_miss":
_PREFILTER_STATS["type_miss"] += 1 _PREFILTER_STATS["type_miss"] += 1
elif reason == "route_miss": elif reason == "route_miss":
@@ -537,6 +568,10 @@ def _record_prefilter_stats(skipped: bool, reason: str | None) -> None:
"matcher prefilter stats: " "matcher prefilter stats: "
f"checked={_PREFILTER_STATS['checked']} " f"checked={_PREFILTER_STATS['checked']} "
f"skipped={_PREFILTER_STATS['skipped']} " 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"type_miss={_PREFILTER_STATS['type_miss']} "
f"route_miss={_PREFILTER_STATS['route_miss']} " f"route_miss={_PREFILTER_STATS['route_miss']} "
f"command_miss={_PREFILTER_STATS['command_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 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( async def _check_matcher_prefilter(
matcher_cls: type[Matcher], event: Event, state: dict | None = None matcher_cls: type[Matcher], event: Event, state: dict | None = None
) -> tuple[bool, str | None]: ) -> tuple[bool, str | None]:
@@ -686,6 +712,17 @@ async def _check_matcher_prefilter(
if not module: if not module:
return False, None 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_modules = _collect_ai_route_modules(event, state)
ai_route_heads = _collect_ai_route_heads(event, state) ai_route_heads = _collect_ai_route_heads(event, state)
if ai_route_modules and module not in ai_route_modules: if ai_route_modules and module not in ai_route_modules:
@@ -696,25 +733,85 @@ async def _check_matcher_prefilter(
await _ensure_route_index() await _ensure_route_index()
if module not in _ROUTE_MODULES_WITH_COMMANDS: 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 return False, None
route_modules = _get_route_modules_for_event(event, state) route_modules = _get_route_modules_for_event(event, state)
if module not in route_modules: if module not in route_modules:
if command_matched:
return False, None
if _matcher_has_alconna_shortcuts(matcher_cls): if _matcher_has_alconna_shortcuts(matcher_cls):
return False, None return False, None
return True, "route_miss" return True, "route_miss"
return False, None 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 _MAX_MATCHER_CACHE = 512
@@ -729,7 +826,7 @@ async def _patched_check_and_run_matcher(
skip, reason = await _check_matcher_prefilter( skip, reason = await _check_matcher_prefilter(
Matcher, event, state if isinstance(state, dict) else None Matcher, event, state if isinstance(state, dict) else None
) )
_record_prefilter_stats(skip, reason) _record_prefilter_stats(skip, reason, "inside_task")
if skip: if skip:
return return
@@ -744,28 +841,6 @@ async def _patched_check_and_run_matcher(
"stack": stack, "stack": stack,
"dependency_cache": dependency_cache, "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) await original(**kwargs)
@@ -788,27 +863,137 @@ def _uninstall_matcher_prefilter() -> None:
_ORIGINAL_CHECK_AND_RUN_MATCHER = None _ORIGINAL_CHECK_AND_RUN_MATCHER = None
def _get_message_text( async def _patched_handle_event(bot: Bot, event: Event) -> None:
message: UniMsg | None, show_log = True
event_cache: dict | None, escape_tag = getattr(nb_message, "escape_tag")
event: Event | None = None, logger_ = getattr(nb_message, "logger")
) -> str: no_log_exception = getattr(nb_message, "NoLogException")
if event_cache is not None:
cached = event_cache.get("plain_text")
if isinstance(cached, str):
return cached
text = "" log_msg = f"<m>{escape_tag(bot.type)} {escape_tag(bot.self_id)}</m> | "
if message is not None: try:
with contextlib.suppress(Exception): log_msg += event.get_log_string()
text = message.extract_plain_text() except no_log_exception:
if not text and event is not None: show_log = False
with contextlib.suppress(Exception): if show_log:
text = (event.get_plaintext() or "").strip() logger_.opt(colors=True).success(log_msg)
if event_cache is not None: state = {}
event_cache["plain_text"] = text dependency_cache = {}
return text 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(
"<r><bg #f8bbd0>Error when checking Matcher.</bg #f8bbd0></r>"
),
}
):
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]: 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 global _CACHE_SWEEP_TASK
await _ensure_route_index() await _ensure_route_index()
_install_matcher_prefilter() _install_matcher_prefilter()
_install_handle_event_selector()
if _CACHE_SWEEP_TASK is None or _CACHE_SWEEP_TASK.done(): if _CACHE_SWEEP_TASK is None or _CACHE_SWEEP_TASK.done():
_CACHE_SWEEP_TASK = asyncio.create_task(_cache_sweep_loop()) _CACHE_SWEEP_TASK = asyncio.create_task(_cache_sweep_loop())
async def stop_auth_runtime_tasks() -> None: async def stop_auth_runtime_tasks() -> None:
global _CACHE_SWEEP_TASK global _CACHE_SWEEP_TASK
_uninstall_handle_event_selector()
_uninstall_matcher_prefilter() _uninstall_matcher_prefilter()
task = _CACHE_SWEEP_TASK task = _CACHE_SWEEP_TASK
_CACHE_SWEEP_TASK = None _CACHE_SWEEP_TASK = None
@@ -873,6 +1060,12 @@ async def _has_limits_cached(module: str, event_cache: dict | None) -> bool:
@contextlib.asynccontextmanager @contextlib.asynccontextmanager
async def _db_section(): async def _db_section():
global DB_ACTIVE_COUNT 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() await DB_SEMAPHORE.acquire()
DB_ACTIVE_COUNT += 1 DB_ACTIVE_COUNT += 1
try: 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 plugin.block_type == BlockType.ALL and not plugin.status:
if group and getattr(group, "is_super", False): if group and getattr(group, "is_super", False):
return False return False
@@ -953,11 +1148,11 @@ async def _get_bot_data_cached(
async def _get_admin_levels_cached( async def _get_admin_levels_cached(
session: Uninfo, entity, event_cache entity, event_cache
) -> tuple[tuple[LevelUserSnapshot | None, LevelUserSnapshot | None] | None, bool]: ) -> tuple[tuple[LevelUserSnapshot | None, LevelUserSnapshot | None] | None, bool]:
if event_cache is not None and "admin_levels" in event_cache: if event_cache is not None and "admin_levels" in event_cache:
return event_cache.get("admin_levels"), event_cache.get("admin_timeout", False) 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: if event_cache is not None:
event_cache["admin_levels"] = levels event_cache["admin_levels"] = levels
event_cache["admin_timeout"] = False event_cache["admin_timeout"] = False
@@ -1074,12 +1269,18 @@ async def get_plugin_and_user(
if user_id in user_cache: if user_id in user_cache:
user = user_cache[user_id] user = user_cache[user_id]
else: else:
async with _db_section(): try:
user = await _fetch_user_readonly(user_dao, user_id) async with _db_section():
user = await _fetch_user_readonly(user_dao, user_id)
except PermissionExemption:
user = None
user_cache[user_id] = user user_cache[user_id] = user
else: else:
async with _db_section(): try:
user = await _fetch_user_readonly(user_dao, user_id) async with _db_section():
user = await _fetch_user_readonly(user_dao, user_id)
except PermissionExemption:
user = None
return plugin, user return plugin, user
@@ -1089,7 +1290,7 @@ async def get_plugin_cost(
plugin: PluginInfo, plugin: PluginInfo,
session: Uninfo, session: Uninfo,
*, *,
is_superuser: bool = False, context: PermissionContext | None = None,
) -> int: ) -> int:
"""获取插件费用 """获取插件费用
@@ -1106,7 +1307,10 @@ async def get_plugin_cost(
返回: 返回:
int: 调用插件金币费用 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 is_superuser:
if plugin.plugin_type == PluginType.SUPERUSER: if plugin.plugin_type == PluginType.SUPERUSER:
raise IsSuperuserException() raise IsSuperuserException()
@@ -1124,7 +1328,7 @@ async def reduce_gold(user_id: str, module: str, cost_gold: int, session: Uninfo
cost_gold: 消耗金币 cost_gold: 消耗金币
session: Uninfo session: Uninfo
""" """
user_dao = DataAccess(UserConsole) should_clear_cache = False
try: try:
await with_timeout( await with_timeout(
UserConsole.reduce_gold( UserConsole.reduce_gold(
@@ -1141,14 +1345,16 @@ async def reduce_gold(user_id: str, module: str, cost_gold: int, session: Uninfo
u.gold = 0 u.gold = 0
await u.save(update_fields=["gold"]) await u.save(update_fields=["gold"])
except asyncio.TimeoutError: except asyncio.TimeoutError:
should_clear_cache = True
logger.error( logger.error(
f"扣除金币超时,用户: {user_id}, 金币: {cost_gold}", f"扣除金币超时,用户: {user_id}, 金币: {cost_gold}",
LOGGER_COMMAND, LOGGER_COMMAND,
session=session, session=session,
) )
# 清除缓存,使下次查询时从数据库获取最新数据 # 正常写入路径由 UserConsole.save() 统一失效缓存;超时状态不确定时兜底清理。
await user_dao.clear_cache(user_id=user_id) if should_clear_cache:
await DataAccess(UserConsole).clear_cache(user_id=user_id)
logger.debug(f"调用功能花费金币: {cost_gold}", LOGGER_COMMAND, session=session) 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(): async def _enter_hooks_section():
"""尝试获取全局信号量并更新计数器,超时则抛出 PermissionExemption。""" """尝试获取全局信号量并更新计数器,饱和时快速放行。"""
global HOOKS_ACTIVE_COUNT global HOOKS_ACTIVE_COUNT
try: if HOOKS_SEMAPHORE.locked():
await asyncio.wait_for(HOOKS_SEMAPHORE.acquire(), timeout=TIMEOUT_SECONDS)
except asyncio.TimeoutError:
logger.warning( logger.warning(
"hooks semaphore acquire timeout, allowing pass", "hooks semaphore saturated, allowing pass",
LOGGER_COMMAND, LOGGER_COMMAND,
) )
raise PermissionExemption("hooks semaphore timeout, allow pass") raise PermissionExemption("hooks semaphore saturated, allow pass")
await HOOKS_SEMAPHORE.acquire()
HOOKS_ACTIVE_COUNT += 1 HOOKS_ACTIVE_COUNT += 1
@@ -1197,14 +1402,7 @@ async def _leave_hooks_section():
async def route_precheck( async def route_precheck(
matcher: Matcher, matcher: Matcher,
event: Event, context: EventContext,
session: Uninfo,
message: UniMsg | None,
*,
entity=None,
event_cache: dict | None = None,
text: str | None = None,
route_modules: set[str] | None = None,
) -> bool: ) -> bool:
module = matcher.plugin_name or "" module = matcher.plugin_name or ""
if not module: if not module:
@@ -1213,19 +1411,20 @@ async def route_precheck(
return False return False
if not _is_command_matcher_class(type(matcher)): if not _is_command_matcher_class(type(matcher)):
return False return False
if entity is None:
entity = get_entity_ids(session) route_modules = context.route_modules if context.route_modules_loaded else None
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)
if route_modules is 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 module in _ROUTE_MODULES_WITH_COMMANDS and module not in route_modules:
if _matcher_has_alconna_shortcuts(type(matcher)): if _matcher_has_alconna_shortcuts(type(matcher)):
return False return False
if event_cache is not None: if context.event_cache is not None:
event_cache["route_skip"] = True context.event_cache["route_skip"] = True
return True return True
return False return False
@@ -1235,14 +1434,10 @@ async def auth(
event: Event, event: Event,
bot: Bot, bot: Bot,
session: Uninfo, session: Uninfo,
message: UniMsg | None,
*, *,
context: EventContext,
skip_ban: bool = False, skip_ban: bool = False,
entity=None, state: dict | None = None,
event_cache: dict | None = None,
text: str | None = None,
route_modules: set[str] | None = None,
is_superuser: bool = False,
): ):
"""权限检查 """权限检查
@@ -1251,20 +1446,28 @@ async def auth(
event: Event event: Event
bot: bot bot: bot
session: Uninfo session: Uninfo
message: UniMsg context: EventContext
""" """
start_time = time.time() start_time = time.time()
cost_gold = 0 cost_gold = 0
ignore_flag = False ignore_flag = False
if entity is None: entity = context.entity
entity = get_entity_ids(session) 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 "" module = matcher.plugin_name or ""
is_command_matcher = _is_command_matcher_class(type(matcher)) 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_allowed = None
auth_result_cache = None auth_result_cache = None
admin_checked_pre = False 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 明细,避免热路径高频构造字符串
hook_recorder = HookTraceRecorder(start_time) hook_recorder = HookTraceRecorder(start_time)
@@ -1278,14 +1481,17 @@ async def auth(
auth_allowed = True auth_allowed = True
return return
if event_cache is not None: side_effect_lock = side_effect_cache.lock_for(module)
auth_result_cache = event_cache.setdefault("auth_result", {}) await side_effect_lock.acquire()
cached_result = auth_result_cache.get(module) entered_side_effect_lock = True
if cached_result is not None:
allowed, reason = cached_result auth_result_cache = side_effect_cache.auth_results
if not allowed: cached_result = auth_result_cache.get(module)
raise SkipPluginException(reason or "auth cached skip") if cached_result is not None:
return allowed, reason = cached_result
if not allowed:
raise SkipPluginException(reason or "auth cached skip")
return
if _is_hidden_plugin(matcher): if _is_hidden_plugin(matcher):
auth_allowed = True auth_allowed = True
@@ -1293,10 +1499,9 @@ async def auth(
if event_cache is not None and event_cache.get("ban_state") is True: if event_cache is not None and event_cache.get("ban_state") is True:
raise SkipPluginException("user or group banned (cached)") raise SkipPluginException("user or group banned (cached)")
if text is None:
text = _get_message_text(message, event_cache, event)
if route_modules is None: if route_modules is None:
route_modules = await _get_route_context(text, event_cache) route_modules = await _get_route_context(text, event_cache)
set_route_modules(state, context, route_modules)
route_skip_checks = ( route_skip_checks = (
is_command_matcher is_command_matcher
and module in _ROUTE_MODULES_WITH_COMMANDS and module in _ROUTE_MODULES_WITH_COMMANDS
@@ -1310,7 +1515,7 @@ async def auth(
auth_allowed = True auth_allowed = True
return return
platform = PlatformUtils.get_platform(session) platform = context.platform
# 获取插件和用户数据 # 获取插件和用户数据
plugin_user_start = time.time() plugin_user_start = time.time()
try: try:
@@ -1336,6 +1541,14 @@ async def auth(
auth_allowed = True auth_allowed = True
return 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 not route_skip_checks and _needs_admin_check(plugin):
if plugin.plugin_type in { if plugin.plugin_type in {
PluginType.SUPERUSER, PluginType.SUPERUSER,
@@ -1356,13 +1569,18 @@ async def auth(
admin_timeout = False admin_timeout = False
if event_cache is not None: if event_cache is not None:
admin_levels, admin_timeout = await _get_admin_levels_cached( 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: if admin_timeout:
hook_recorder.set("auth_admin", "timeout") hook_recorder.set("auth_admin", "timeout")
else: else:
admin_start = time.time() admin_start = time.time()
await auth_admin(plugin, session, cached_levels=admin_levels) await auth_admin(
plugin,
session,
context=permission_context,
)
hook_recorder.set( hook_recorder.set(
"auth_admin", f"{time.time() - admin_start:.3f}s(pre)" "auth_admin", f"{time.time() - admin_start:.3f}s(pre)"
) )
@@ -1386,8 +1604,7 @@ async def auth(
matcher, matcher,
session, session,
plugin, plugin,
entity=entity, context=permission_context,
is_superuser=is_superuser,
) )
hook_recorder.set("auth_ban", f"{time.time() - ban_start:.3f}s") hook_recorder.set("auth_ban", f"{time.time() - ban_start:.3f}s")
if event_cache is not None: if event_cache is not None:
@@ -1407,7 +1624,7 @@ async def auth(
user, user,
plugin, plugin,
session, session,
is_superuser=is_superuser, context=permission_context,
), ),
name="get_plugin_cost", name="get_plugin_cost",
) )
@@ -1421,7 +1638,7 @@ async def auth(
hook_recorder.set("cost_gold", "skipped") hook_recorder.set("cost_gold", "skipped")
# 执行 bot_filter # 执行 bot_filter
bot_filter(session) bot_filter(session, context=permission_context)
group = await _get_group_cached(entity, event_cache) group = await _get_group_cached(entity, event_cache)
@@ -1439,13 +1656,23 @@ async def auth(
and not route_skip_checks and not route_skip_checks
): ):
admin_levels, admin_timeout = await _get_admin_levels_cached( 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 检查,并记录执行时间 # 并行执行所有 hook 检查,并记录执行时间
hooks_start = time.time() hooks_start = time.time()
allow_sleep_bypass = _is_bot_wake_command(module, text) allow_sleep_bypass = _is_bot_wake_command(module, text)
# 先进入 hooks 并行检查区域;饱和时快速放行,避免创建并积压协程。
await _enter_hooks_section()
entered_hooks = True
# 创建所有 hook 任务 # 创建所有 hook 任务
hook_tasks = [] hook_tasks = []
if event_cache is None: if event_cache is None:
@@ -1455,6 +1682,7 @@ async def auth(
plugin, plugin,
bot.self_id, bot.self_id,
allow_sleep_bypass=allow_sleep_bypass, allow_sleep_bypass=allow_sleep_bypass,
context=permission_context,
), ),
"auth_bot", "auth_bot",
hook_recorder, hook_recorder,
@@ -1472,6 +1700,7 @@ async def auth(
bot_data=bot_data, bot_data=bot_data,
skip_fetch=True, skip_fetch=True,
allow_sleep_bypass=allow_sleep_bypass, allow_sleep_bypass=allow_sleep_bypass,
context=permission_context,
), ),
"auth_bot", "auth_bot",
hook_recorder, hook_recorder,
@@ -1483,7 +1712,13 @@ async def auth(
else: else:
hook_tasks.append( hook_tasks.append(
time_hook( time_hook(
auth_group(plugin, group, text, entity.group_id), auth_group(
plugin,
group,
text,
entity.group_id,
context=permission_context,
),
"auth_group", "auth_group",
hook_recorder, hook_recorder,
) )
@@ -1492,7 +1727,11 @@ async def auth(
if not route_skip_checks and plugin.admin_level and not admin_checked_pre: if not route_skip_checks and plugin.admin_level and not admin_checked_pre:
if event_cache is None: if event_cache is None:
hook_tasks.append( 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: else:
if admin_timeout: if admin_timeout:
@@ -1500,7 +1739,11 @@ async def auth(
else: else:
hook_tasks.append( hook_tasks.append(
time_hook( time_hook(
auth_admin(plugin, session, cached_levels=admin_levels), auth_admin(
plugin,
session,
context=permission_context,
),
"auth_admin", "auth_admin",
hook_recorder, hook_recorder,
) )
@@ -1510,7 +1753,7 @@ async def auth(
if is_superuser: if is_superuser:
hook_recorder.set("auth_plugin", "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( hook_tasks.append(
time_hook( time_hook(
auth_plugin( auth_plugin(
@@ -1518,6 +1761,7 @@ async def auth(
group, group,
session, session,
event, event,
context=permission_context,
skip_group_block=is_superuser, skip_group_block=is_superuser,
), ),
"auth_plugin", "auth_plugin",
@@ -1531,18 +1775,17 @@ async def auth(
has_limits = await _has_limits_cached(module, event_cache) has_limits = await _has_limits_cached(module, event_cache)
if has_limits: if has_limits:
hook_tasks.append( 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: else:
hook_recorder.set("auth_limit", "skipped") hook_recorder.set("auth_limit", "skipped")
else: else:
hook_recorder.set("auth_limit", "skipped") hook_recorder.set("auth_limit", "skipped")
if hook_tasks:
# 进入 hooks 并行检查区域(会在高并发时排队)
await _enter_hooks_section()
entered_hooks = True
# 使用 gather 并行执行所有 hook,但添加总体超时控制 # 使用 gather 并行执行所有 hook,但添加总体超时控制
try: try:
await with_timeout( await with_timeout(
@@ -1599,6 +1842,9 @@ async def auth(
) )
if auth_result_cache is not None and auth_allowed is not None: if auth_result_cache is not None and auth_allowed is not None:
auth_result_cache[module] = (auth_allowed, 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: if not ignore_flag and cost_gold > 0:
gold_start = time.time() gold_start = time.time()
+47 -98
View File
@@ -1,4 +1,3 @@
import contextlib
import time import time
from nonebot import get_driver 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.cache.runtime_cache import is_cache_ready
from zhenxun.services.log import logger 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.services.runtime_bootstrap import register_runtime_bootstrap
from zhenxun.utils.utils import get_entity_ids
from .auth.config import LOGGER_COMMAND 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 ( from .auth_checker import (
LimitManager, LimitManager,
_get_event_cache,
_get_route_context, _get_route_context,
auth, auth,
route_precheck, route_precheck,
@@ -41,17 +46,6 @@ async def _mark_bot_connected(bot: Bot):
_BOT_CONNECT_TS = time.time() _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 @driver.on_startup
async def _start_auth_runtime_tasks(): async def _start_auth_runtime_tasks():
await 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 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 @event_preprocessor
async def _drop_message_before_cache_ready(event: Event): async def _drop_message_before_cache_ready(event: Event):
mark_activity()
if event.get_type() != "message": if event.get_type() != "message":
return return
if not is_cache_ready(): if not is_cache_ready():
@@ -130,46 +96,22 @@ async def _auth_preprocessor(
return return
start_time = time.time() start_time = time.time()
entity = state.get("_zx_entity") event_context = get_or_create_event_context(
if entity is None: bot,
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, event,
session, session,
message, state,
entity=entity, message=message,
event_cache=event_cache, )
text=text,
route_modules=route_modules, 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 return
try: try:
@@ -178,13 +120,9 @@ async def _auth_preprocessor(
event, event,
bot, bot,
session, session,
message, context=event_context,
skip_ban=False, skip_ban=False,
entity=entity, state=state,
event_cache=event_cache,
text=text,
route_modules=route_modules,
is_superuser=is_superuser,
) )
except IgnoredException: except IgnoredException:
raise raise
@@ -203,16 +141,27 @@ async def _auth_preprocessor(
@run_postprocessor @run_postprocessor
async def _unblock_after_matcher(matcher: Matcher, session: Uninfo, event: Event): async def _unblock_after_matcher(
user_id = _resolve_actor_user_id(event, session.user.id) matcher: Matcher,
group_id = _resolve_event_group_id(event, None) session: Uninfo,
channel_id = _resolve_event_channel_id(event, None) event: Event,
if session.group: state: T_State,
if session.group.parent: ):
group_id = session.group.parent.id context = get_event_context(state)
channel_id = session.group.id if context is not None:
else: user_id = context.user_id
group_id = session.group.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: if user_id and matcher.plugin:
module = matcher.plugin.name module = matcher.plugin.name
LimitManager.unblock(module, user_id, group_id, channel_id) LimitManager.unblock(module, user_id, group_id, channel_id)
+24 -14
View File
@@ -16,13 +16,7 @@ from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType from zhenxun.utils.enum import PluginType
from zhenxun.utils.message import MessageUtils from zhenxun.utils.message import MessageUtils
malicious_check_time = Config.get_config("hook", "MALICIOUS_CHECK_TIME") from .auth.context import resolve_actor_user_id, resolve_event_group_id
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")
class BanCheckLimiter: class BanCheckLimiter:
@@ -36,6 +30,10 @@ class BanCheckLimiter:
self.default_check_time = default_check_time self.default_check_time = default_check_time
self.default_count = default_count 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): def add(self, key: str | float):
if self.mint[key] == 1: if self.mint[key] == 1:
self.mtime[key] = time.time() self.mtime[key] = time.time()
@@ -59,11 +57,22 @@ class BanCheckLimiter:
_blmt = BanCheckLimiter( _blmt = BanCheckLimiter(
malicious_check_time, 5,
malicious_ban_count, 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 @run_preprocessor
async def _( async def _(
@@ -88,11 +97,12 @@ async def _(
else: else:
return return
user_id = session.id1 user_id = resolve_actor_user_id(event, session.id1)
group_id = session.id3 or session.id2 group_id = resolve_event_group_id(event, session.id3 or session.id2)
malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME") malicious_check_time = float(_get_positive_config("MALICIOUS_CHECK_TIME", float))
if not malicious_ban_time: malicious_ban_count = int(_get_positive_config("MALICIOUS_BAN_COUNT", int))
raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_TIME] 为空或小于0") 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 user_id and module:
if _blmt.check(f"{user_id}__{module}"): if _blmt.check(f"{user_id}__{module}"):
await BanConsole.ban( await BanConsole.ban(
@@ -29,7 +29,6 @@ def register_cache_types():
GroupPluginSetting, GroupPluginSetting,
key_format="{group_id}_{plugin_name}_{key}", key_format="{group_id}_{plugin_name}_{key}",
) )
CacheRegistry.register(CacheType.GROUP_PLUGIN_SETTINGS_VIEW, dict)
CacheRegistry.register( CacheRegistry.register(
CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}" CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}"
) )
+1 -1
View File
@@ -88,7 +88,7 @@ async def _handle_setting(
) )
@PriorityLifecycle.on_startup(priority=5) @PriorityLifecycle.on_startup(priority=4)
async def _(): async def _():
""" """
初始化插件数据配置 初始化插件数据配置
@@ -3,12 +3,12 @@ from pathlib import Path
import random import random
import shutil import shutil
from aiocache import cached
import ujson as json import ujson as json
from zhenxun.builtin_plugins.plugin_store.models import StorePluginInfo from zhenxun.builtin_plugins.plugin_store.models import StorePluginInfo
from zhenxun.configs.path_config import TEMP_PATH from zhenxun.configs.path_config import TEMP_PATH
from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.cache.bounded_ttl import BoundedTTLCache
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.services.plugin_init import PluginInitManager from zhenxun.services.plugin_init import PluginInitManager
from zhenxun.utils.enum import PluginType from zhenxun.utils.enum import PluginType
@@ -26,6 +26,14 @@ from .config import (
) )
from .exceptions import PluginStoreException 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: 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_parts = [part for part in plugin_info.module_path.split(".") if part]
relative_path = Path(*relative_parts) if relative_parts else Path(plugin_name) relative_path = Path(*relative_parts) if relative_parts else Path(plugin_name)
path = BASE_PATH.parent / relative_path path = BASE_PATH.parent / relative_path
if plugin_info.is_dir: return path if plugin_info.is_dir else path.parent / f"{plugin_name}.py"
return path
return path.parent / f"{plugin_name}.py"
@classmethod @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]]: async def get_data(cls) -> tuple[list[StorePluginInfo], list[StorePluginInfo]]:
"""获取插件信息数据 """获取插件信息数据
@@ -77,15 +74,22 @@ class StoreManager:
tuple[list[StorePluginInfo], list[StorePluginInfo]]: 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( plugins = await RepoFileManager.get_file_content(
DEFAULT_GITHUB_URL, "plugins.json" DEFAULT_GITHUB_URL, "plugins.json"
) )
extra_plugins = await RepoFileManager.get_file_content( extra_plugins = await RepoFileManager.get_file_content(
EXTRA_GITHUB_URL, "plugins.json", "index" EXTRA_GITHUB_URL, "plugins.json", "index"
) )
return [StorePluginInfo(**plugin) for plugin in json.loads(plugins)], [ result = (
StorePluginInfo(**plugin) for plugin in json.loads(extra_plugins) [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 @classmethod
def version_check(cls, plugin_info: StorePluginInfo, suc_plugin: dict[str, str]): def version_check(cls, plugin_info: StorePluginInfo, suc_plugin: dict[str, str]):
@@ -330,13 +334,12 @@ class StoreManager:
source: 源 source: 源
""" """
repo_type = RepoType.GITHUB if is_external else None 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 repo_type = RepoType.ALIYUN
elif source == "git": elif source == "git":
repo_type = RepoType.GITHUB repo_type = RepoType.GITHUB
else:
if plugin_info.ali_url:
repo_type = RepoType.ALIYUN
module_path = plugin_info.module_path module_path = plugin_info.module_path
is_dir = plugin_info.is_dir is_dir = plugin_info.is_dir
github_url = plugin_info.github_url github_url = plugin_info.github_url
@@ -380,7 +383,7 @@ class StoreManager:
requirement_file = target_dir / requirement_path.path requirement_file = target_dir / requirement_path.path
if requirement_file.exists(): if requirement_file.exists():
is_install_req = True is_install_req = True
await VirtualEnvPackageManager.install_requirement(requirement_file) await VirtualEnvPackageManager.add_requirement(requirement_file)
if not is_install_req: if not is_install_req:
# 从仓库根目录查找文件 # 从仓库根目录查找文件
@@ -401,13 +404,13 @@ class StoreManager:
f"开始安装插件 {module_path} 依赖文件: {requirement_path}", f"开始安装插件 {module_path} 依赖文件: {requirement_path}",
LOG_COMMAND, LOG_COMMAND,
) )
await VirtualEnvPackageManager.install_requirement(requirement_path) await VirtualEnvPackageManager.add_requirement(requirement_path)
if requirements_path.exists(): if requirements_path.exists():
logger.info( logger.info(
f"开始安装插件 {module_path} 依赖文件: {requirements_path}", f"开始安装插件 {module_path} 依赖文件: {requirements_path}",
LOG_COMMAND, LOG_COMMAND,
) )
await VirtualEnvPackageManager.install_requirement(requirements_path) await VirtualEnvPackageManager.add_requirement(requirements_path)
@classmethod @classmethod
async def remove_plugin(cls, index_or_module: str) -> str: async def remove_plugin(cls, index_or_module: str) -> str:
+7 -5
View File
@@ -19,9 +19,9 @@ from zhenxun.models.friend_user import FriendUser
from zhenxun.models.goods_info import GoodsInfo from zhenxun.models.goods_info import GoodsInfo
from zhenxun.models.group_member_info import GroupInfoUser from zhenxun.models.group_member_info import GroupInfoUser
from zhenxun.models.user_console import UserConsole 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.models.user_props_log import UserPropsLog
from zhenxun.services import avatar_service from zhenxun.services import avatar_service
from zhenxun.services.buffered_writers import append_user_gold_log
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.ui.models import ImageCell, TextCell from zhenxun.ui.models import ImageCell, TextCell
from zhenxun.utils.enum import GoldHandle, PropHandle from zhenxun.utils.enum import GoldHandle, PropHandle
@@ -480,10 +480,6 @@ class ShopManage:
).count() ).count()
if goods.daily_limit and count >= goods.daily_limit: if goods.daily_limit and count >= goods.daily_limit:
return "今天的购买已达限制了喔!" 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( logger.info(
f"花费 {price} 金币购买 {goods.goods_name} ×{num} 成功!", f"花费 {price} 金币购买 {goods.goods_name} ×{num} 成功!",
"购买道具", "购买道具",
@@ -494,6 +490,12 @@ class ShopManage:
user.props[goods.uuid] = 0 user.props[goods.uuid] = 0
user.props[goods.uuid] += num user.props[goods.uuid] += num
await user.save(update_fields=["gold", "props"]) 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} 成功!" return f"花费 {price} 金币购买 {goods.goods_name} ×{num} 成功!"
@classmethod @classmethod
@@ -9,13 +9,17 @@ import pytz
from zhenxun import ui from zhenxun import ui
from zhenxun.configs.path_config import IMAGE_PATH from zhenxun.configs.path_config import IMAGE_PATH
from zhenxun.models.friend_user import FriendUser 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.group_member_info import GroupInfoUser
from zhenxun.models.sign_log import SignLog from zhenxun.models.sign_log import SignLog
from zhenxun.models.sign_user import SignUser from zhenxun.models.sign_user import SignUser
from zhenxun.models.user_console import UserConsole from zhenxun.models.user_console import UserConsole
from zhenxun.services.avatar_service import avatar_service 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.services.log import logger
from zhenxun.ui.models import ImageCell, TextCell 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 zhenxun.utils.platform import PlatformUtils
from ._random_event import random_event from ._random_event import random_event
@@ -182,11 +186,32 @@ class SignManage:
gift = random_event(float(user.impression)) gift = random_event(float(user.impression))
if isinstance(gift, int): if isinstance(gift, int):
gold += gift 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}" gift = f"额外金币 +{gift}"
else: else:
await UserConsole.add_gold(user.user_id, gold, "sign_in", platform) goods = await GoodsInfo.get_or_none(goods_name=gift)
await UserConsole.add_props_by_name(user.user_id, gift, 1, platform) 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" gift += " + 1"
logger.info( logger.info(
f"签到成功. score: {user.impression:.2f} " f"签到成功. score: {user.impression:.2f} "
@@ -1,5 +1,7 @@
import asyncio
from datetime import datetime from datetime import datetime
from nonebot import get_driver
from nonebot.adapters import Bot, Event from nonebot.adapters import Bot, Event
from nonebot.adapters.onebot.v11 import PokeNotifyEvent from nonebot.adapters.onebot.v11 import PokeNotifyEvent
from nonebot.matcher import Matcher from nonebot.matcher import Matcher
@@ -25,7 +27,35 @@ __plugin_meta__ = PluginMetadata(
).to_dict(), ).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 @run_postprocessor
@@ -50,7 +80,7 @@ async def _(
if plugin_type == PluginType.NORMAL: if plugin_type == PluginType.NORMAL:
entity = get_entity_ids(session) entity = get_entity_ids(session)
logger.debug(f"提交调用记录: {matcher.plugin_name}...", session=session) logger.debug(f"提交调用记录: {matcher.plugin_name}...", session=session)
TEMP_LIST.append( await _append_statistics(
Statistics( Statistics(
user_id=entity.user_id, user_id=entity.user_id,
group_id=entity.group_id, group_id=entity.group_id,
@@ -66,10 +96,11 @@ async def _():
try: try:
if should_pause_tasks(): if should_pause_tasks():
return return
call_list = TEMP_LIST.copy() await _flush_statistics_buffer("定时")
TEMP_LIST.clear()
if call_list:
await Statistics.bulk_create(call_list)
logger.debug(f"批量添加调用记录 {len(call_list)} 条", "定时任务")
except Exception as e: except Exception as e:
logger.error("定时批量添加调用记录", "定时任务", e=e) logger.error("定时批量添加调用记录", "定时任务", e=e)
@driver.on_shutdown
async def _flush_statistics_on_shutdown():
await _flush_statistics_buffer("关闭")
@@ -1,3 +1,5 @@
import contextlib
from nonebot.permission import SUPERUSER from nonebot.permission import SUPERUSER
from nonebot.plugin import PluginMetadata from nonebot.plugin import PluginMetadata
from nonebot.rule import to_me 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.llm.manager import clear_model_cache
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType from zhenxun.utils.enum import PluginType
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
from zhenxun.utils.message import MessageUtils from zhenxun.utils.message import MessageUtils
AUTO_RELOAD_JOB_ID = "zhenxun.reload_setting.auto_reload"
__plugin_meta__ = PluginMetadata( __plugin_meta__ = PluginMetadata(
name="重载配置", name="重载配置",
description="重新加载config.yaml", description="重新加载config.yaml",
@@ -53,22 +58,75 @@ _matcher = on_alconna(
) )
@_matcher.handle() def _get_auto_reload_interval() -> int:
async def _(session: EventSession, arparma: Arparma): 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() Config.reload()
get_llm_config.cache_clear() get_llm_config.cache_clear()
clear_model_cache() 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) logger.debug("自动重载配置文件", arparma.header_result, session=session)
await MessageUtils.build_message("重载完成!").send(reply_to=True) await MessageUtils.build_message("重载完成!").send(reply_to=True)
@scheduler.scheduled_job( async def _auto_reload_config() -> None:
"interval",
seconds=Config.get_config("reload_setting", "AUTO_RELOAD_TIME", 180),
)
async def _():
if Config.get_config("reload_setting", "AUTO_RELOAD"): if Config.get_config("reload_setting", "AUTO_RELOAD"):
Config.reload() await _reload_runtime_config()
get_llm_config.cache_clear()
clear_model_cache()
logger.debug("已自动重载配置文件...") logger.debug("已自动重载配置文件...")
+1 -23
View File
@@ -1,20 +1,17 @@
import asyncio
import secrets import secrets
from fastapi import APIRouter, FastAPI from fastapi import APIRouter, FastAPI
import nonebot import nonebot
from nonebot.log import default_filter, default_format
from nonebot.plugin import PluginMetadata from nonebot.plugin import PluginMetadata
from zhenxun.configs.config import Config as gConfig from zhenxun.configs.config import Config as gConfig
from zhenxun.configs.utils import PluginExtraData, RegisterConfig 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.enum import PluginType
from zhenxun.utils.manager.priority_manager import PriorityLifecycle from zhenxun.utils.manager.priority_manager import PriorityLifecycle
from .api.configure import router as configure_router from .api.configure import router as configure_router
from .api.logs import router as ws_log_routes 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.menu import router as menu_router
from .api.tabs.dashboard import router as dashboard_router from .api.tabs.dashboard import router as dashboard_router
from .api.tabs.database import router as database_router from .api.tabs.database import router as database_router
@@ -95,25 +92,6 @@ WsApiRouter.include_router(chat_routes)
@PriorityLifecycle.on_startup(priority=0) @PriorityLifecycle.on_startup(priority=0)
async def _(): async def _():
try: 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: FastAPI = nonebot.get_app()
app.include_router(BaseApiRouter) app.include_router(BaseApiRouter)
app.include_router(WsApiRouter) app.include_router(WsApiRouter)
@@ -1,33 +1,97 @@
import asyncio import asyncio
from collections import deque
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
from typing import Generic, TypeVar import contextlib
_T = TypeVar("_T") from nonebot.log import default_filter, default_format
LogListener = Callable[[_T], Awaitable[None]]
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.count, self.rotation = 0, rotation
self.max_logs = max_logs
self.max_listeners = max_listeners
self.logs: dict[int, str] = {} 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): async def add(self, log: str):
seq = self.count = self.count + 1 seq = self.count = self.count + 1
self.logs[seq] = log self.logs[seq] = log
self._order.append(seq)
self._trim()
asyncio.get_running_loop().call_later(self.rotation, self.remove, seq) asyncio.get_running_loop().call_later(self.rotation, self.remove, seq)
await asyncio.gather( listeners = tuple(self.listeners)
*(listener(log) for listener in self.listeners), if listeners:
return_exceptions=True, 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 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): 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
@@ -3,7 +3,7 @@ from loguru import logger
from nonebot.utils import escape_tag from nonebot.utils import escape_tag
from starlette.websockets import WebSocket, WebSocketDisconnect, WebSocketState 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() router = APIRouter()
@@ -11,11 +11,16 @@ router = APIRouter()
@router.websocket("/logs") @router.websocket("/logs")
async def system_logs_realtime(websocket: WebSocket): async def system_logs_realtime(websocket: WebSocket):
await websocket.accept() await websocket.accept()
await ensure_log_sink_started()
async def log_listener(log: str): async def log_listener(log: str):
await websocket.send_text(log) 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: try:
while websocket.client_state == WebSocketState.CONNECTED: while websocket.client_state == WebSocketState.CONNECTED:
recv = await websocket.receive() recv = await websocket.receive()
@@ -26,4 +31,5 @@ async def system_logs_realtime(websocket: WebSocket):
except WebSocketDisconnect: except WebSocketDisconnect:
pass pass
finally: finally:
LOG_STORAGE.listeners.remove(log_listener) LOG_STORAGE.remove_listener(log_listener)
stop_log_sink_if_idle()
@@ -34,6 +34,31 @@ run_time = time.time()
ws_router = APIRouter() ws_router = APIRouter()
router = APIRouter(prefix="/main") 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( @router.get(
@@ -243,11 +268,39 @@ async def _(param: BotManageUpdateParam):
@ws_router.websocket("/system_status") @ws_router.websocket("/system_status")
async def system_logs_realtime(websocket: WebSocket, sleep: int = 5): async def system_logs_realtime(websocket: WebSocket, sleep: int = 5):
await websocket.accept() await websocket.accept()
_SYSTEM_STATUS_CONNECTIONS.add(websocket)
logger.debug("ws system_status is connect") logger.debug("ws system_status is connect")
with contextlib.suppress(
WebSocketDisconnect, ConnectionClosedError, ConnectionClosedOK disconnect_event = asyncio.Event()
):
while websocket.client_state == WebSocketState.CONNECTED: 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() system_status = await get_system_status()
await websocket.send_text(system_status.json()) await asyncio.wait_for(websocket.send_text(system_status.json()), timeout=5)
await asyncio.sleep(sleep) 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)
@@ -52,34 +52,20 @@ async def _(
async def _() -> Result[PluginCount]: async def _() -> Result[PluginCount]:
try: try:
plugin_count = PluginCount() plugin_count = PluginCount()
plugin_count.normal = len( plugins = await DbPluginInfo.get_plugins(
await DbPluginInfo.get_plugins( load_status=True,
plugin_type=PluginType.NORMAL, filter_parent=False,
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,
)
) )
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, "拿到信息啦!") return Result.ok(plugin_count, "拿到信息啦!")
except Exception as e: except Exception as e:
logger.error(f"{router.prefix}/get_plugin_count 调用错误", "WebUi", e=e) logger.error(f"{router.prefix}/get_plugin_count 调用错误", "WebUi", e=e)
@@ -111,10 +111,19 @@ class ApiDataSource:
other_update_fields = set() other_update_fields = set()
updated_count = 0 updated_count = 0
errors = [] 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: for item in params.updates:
try: 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_other = False
plugin_changed_block = False plugin_changed_block = False
+129 -13
View File
@@ -8,12 +8,27 @@
from __future__ import annotations from __future__ import annotations
import atexit
import importlib.metadata import importlib.metadata
import os
from pathlib import Path from pathlib import Path
import signal
import subprocess import subprocess
import sys import sys
import time 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: def _print_version() -> None:
try: try:
@@ -94,32 +109,69 @@ def _run_worker() -> None:
nonebot.logger.info(f"加载第三方插件目录: {ext}") nonebot.logger.info(f"加载第三方插件目录: {ext}")
nonebot.load_plugins(ext) nonebot.load_plugins(ext)
nonebot.run() nonebot.run(timeout_graceful_shutdown=GRACEFUL_SHUTDOWN_TIMEOUT)
def _build_worker_command() -> list[str]: def _build_worker_command() -> list[str]:
return [sys.executable, "-m", "zhenxun.cli", "run-worker"] 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: def _wait_worker_exit(proc: subprocess.Popen, timeout_seconds: float) -> bool:
deadline = time.monotonic() + timeout_seconds deadline = time.monotonic() + timeout_seconds
while time.monotonic() < deadline: while time.monotonic() < deadline:
if proc.poll() is not None: if proc.poll() is not None:
return True return True
time.sleep(0.1) time.sleep(WORKER_POLL_INTERVAL)
return proc.poll() is not None return proc.poll() is not None
def _terminate_worker(proc: subprocess.Popen) -> None: def _terminate_worker(proc: subprocess.Popen) -> None:
if proc.poll() is not None: if proc.poll() is not None:
return return
if _wait_worker_exit(proc, 8.0): _launcher_log(f"stopping worker pid={proc.pid}")
return if os.name == "nt":
proc.terminate() ctrl_break_event = getattr(signal, "CTRL_BREAK_EVENT", None)
if _wait_worker_exit(proc, 5.0): 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 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.kill()
proc.wait(timeout=5) proc.wait(timeout=WORKER_KILL_TIMEOUT)
def _run_launcher() -> None: def _run_launcher() -> None:
@@ -130,19 +182,83 @@ def _run_launcher() -> None:
) )
clear_launcher_restart_signal() clear_launcher_restart_signal()
while True: current_worker: subprocess.Popen | None = None
worker = subprocess.Popen(_build_worker_command(), cwd=str(cwd)) 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: 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: except KeyboardInterrupt:
clear_launcher_restart_signal() clear_launcher_restart_signal()
_terminate_worker(worker) _terminate_worker(worker)
return return
finally:
if current_worker is worker:
current_worker = None
should_restart = consume_launcher_restart_signal() if restart_requested or consume_launcher_restart_signal():
if should_restart:
continue continue
raise SystemExit(return_code) raise SystemExit(return_code if return_code is not None else 1)
def main() -> None: def main() -> None:
+108 -19
View File
@@ -4,12 +4,12 @@ from pathlib import Path
from typing import Any, TypeVar from typing import Any, TypeVar
import cattrs import cattrs
from nonebot.log import logger as _nonebot_logger
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from ruamel.yaml import YAML from ruamel.yaml import YAML
from ruamel.yaml.scanner import ScannerError from ruamel.yaml.scanner import ScannerError
from zhenxun.configs.path_config import DATA_PATH from zhenxun.configs.path_config import DATA_PATH
from zhenxun.services.log import logger
from zhenxun.utils.pydantic_compat import ( from zhenxun.utils.pydantic_compat import (
_dump_pydantic_obj, _dump_pydantic_obj,
_is_pydantic_type, _is_pydantic_type,
@@ -38,6 +38,33 @@ _yaml.indent = 2
_yaml.allow_unicode = True _yaml.allow_unicode = True
T = TypeVar("T") 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): class NoSuchConfig(Exception):
@@ -114,21 +141,83 @@ class ConfigsManager:
self._simple_data: dict = {} self._simple_data: dict = {}
self._simple_file = DATA_PATH / "config.yaml" self._simple_file = DATA_PATH / "config.yaml"
self.add_module = [] self.add_module = []
_yaml = YAML()
if file: if file:
file.parent.mkdir(exist_ok=True, parents=True) file.parent.mkdir(exist_ok=True, parents=True)
self.file = file self.file = file
self.load_data() self.load_data()
if self._simple_file.exists(): if self._simple_file.exists():
try: self._load_simple_data(raise_on_error=True)
with self._simple_file.open(encoding="utf8") as f: self._apply_simple_data(warn_unknown=False)
self._simple_data = _yaml.load(f)
except ScannerError as e: def _load_simple_data(self, *, raise_on_error: bool = False) -> None:
raise ScannerError( if not self._simple_file.exists():
f"{e}\n**********************************************\n" self._simple_data = {}
f"****** 可能为config.yaml配置文件填写不规范 ******\n" return
f"**********************************************" try:
) from e 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): 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)): if module in self._data and (config := self._data[module].configs.get(key)):
existing_value = config.value 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) processed_default_value = self._normalize_config_data(default_value)
self.add_module.append(f"{module}:{key}".lower()) self.add_module.append(f"{module}:{key}".lower())
@@ -231,7 +324,7 @@ class ConfigsManager:
config.help = help config.help = help
config.arg_parser = arg_parser config.arg_parser = arg_parser
config.type = type config.type = type
if _override: if simple_value is not _MISSING or _override:
config.value = processed_value config.value = processed_value
config.default_value = processed_default_value config.default_value = processed_default_value
else: else:
@@ -371,12 +464,8 @@ class ConfigsManager:
def reload(self): def reload(self):
"""重新加载配置文件""" """重新加载配置文件"""
if self._simple_file.exists(): self._load_simple_data()
with open(self._simple_file, encoding="utf8") as f: self._apply_simple_data(warn_unknown=True)
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.save() self.save()
def load_data(self): def load_data(self):
+52 -1
View File
@@ -2,6 +2,7 @@ from typing import ClassVar
from tortoise import fields from tortoise import fields
from zhenxun.configs.config import BotConfig
from zhenxun.services.db_context import Model from zhenxun.services.db_context import Model
@@ -51,4 +52,54 @@ class GroupInfoUser(Model):
@classmethod @classmethod
async def _run_script(cls): 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
+18 -8
View File
@@ -5,12 +5,11 @@ from tortoise import fields
from tortoise.exceptions import IntegrityError from tortoise.exceptions import IntegrityError
from zhenxun.models.goods_info import GoodsInfo 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.services.db_context import Model
from zhenxun.utils.enum import CacheType, GoldHandle from zhenxun.utils.enum import CacheType, GoldHandle
from zhenxun.utils.exception import GoodsNotFound, InsufficientGold from zhenxun.utils.exception import GoodsNotFound, InsufficientGold
from .user_gold_log import UserGoldLog
class UserConsole(Model): class UserConsole(Model):
id = fields.IntField(pk=True, generated=True, auto_increment=True) 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) user, _ = await cls.get_or_create_user(user_id=user_id, platform=platform)
return user 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 @classmethod
async def get_new_uid(cls) -> int: async def get_new_uid(cls) -> int:
"""获取最新uid """获取最新uid
@@ -103,10 +113,10 @@ class UserConsole(Model):
source: 来源 source: 来源
platform: 平台. 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 user.gold += gold
await user.save(update_fields=["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 user_id=user_id, gold=gold, handle=GoldHandle.GET, source=source
) )
@@ -131,12 +141,12 @@ class UserConsole(Model):
异常: 异常:
InsufficientGold: 金币不足 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: if user.gold < gold:
raise InsufficientGold() raise InsufficientGold()
user.gold -= gold user.gold -= gold
await user.save(update_fields=["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 user_id=user_id, gold=gold, handle=handle, source=plugin_module
) )
@@ -152,7 +162,7 @@ class UserConsole(Model):
num: 道具数量. num: 道具数量.
platform: 平台. 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: if goods_uuid not in user.props:
user.props[goods_uuid] = 0 user.props[goods_uuid] = 0
user.props[goods_uuid] += num user.props[goods_uuid] += num
@@ -186,7 +196,7 @@ class UserConsole(Model):
num: 道具数量. num: 道具数量.
platform: 平台. 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: if goods_uuid not in user.props or user.props[goods_uuid] < num:
raise GoodsNotFound("未找到商品或道具数量不足...") raise GoodsNotFound("未找到商品或道具数量不足...")
+5
View File
@@ -63,6 +63,11 @@ class AvatarService:
identifier = str(identifier) identifier = str(identifier)
return self.cache_path / platform / f"{identifier}.png" 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( async def get_avatar_path(
self, platform: str, identifier: str, force_refresh: bool = False self, platform: str, identifier: str, force_refresh: bool = False
) -> Path | None: ) -> Path | None:
+128
View File
@@ -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()
+10 -6
View File
@@ -74,6 +74,7 @@ from .config import (
) )
__all__ = [ __all__ = [
"BoundedTTLCache",
"Cache", "Cache",
"CacheDict", "CacheDict",
"CacheManager", "CacheManager",
@@ -82,6 +83,7 @@ __all__ = [
] ]
from . import runtime_cache as _runtime_cache # noqa: F401 from . import runtime_cache as _runtime_cache # noqa: F401
from .bounded_ttl import BoundedTTLCache
T = TypeVar("T") T = TypeVar("T")
U = TypeVar("U") U = TypeVar("U")
@@ -400,9 +402,10 @@ class CacheManager:
"""清除缓存 """清除缓存
参数: 参数:
cache_type: 缓存类型,为None时清除所有缓存。 cache_type: 缓存类型。为 None 时清除整个 backend。
注意:受 aiocache 限制,无法按类型精确删除, 指定 cache_type 时不再退化为清除整个 backend,避免误删其他类型缓存。
指定 cache_type 时仅清除整个 backend(行为与不指定相同)。 需要刷新模型运行态缓存时,应调用对应
RuntimeCache.refresh/upsert/remove。
返回: 返回:
bool: 是否成功 bool: 是否成功
@@ -413,11 +416,12 @@ class CacheManager:
try: try:
if cache_type: if cache_type:
logger.debug( logger.warning(
f"清除缓存类型 {cache_type}" f"拒绝清除缓存类型 {cache_type}: "
"(aiocache 不支持按前缀删除,清除整个 backend)", "当前后端不支持可靠的按类型清理,已避免清除整个 backend",
LOG_COMMAND, LOG_COMMAND,
) )
return False
await self.cache_backend.clear() # type: ignore await self.cache_backend.clear() # type: ignore
return True return True
except Exception as e: except Exception as e:
+218
View File
@@ -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
+84 -2
View File
@@ -1,9 +1,12 @@
from dataclasses import dataclass from dataclasses import dataclass
import time import time
from typing import Any, Generic, TypeVar from typing import Any, Generic, TypeVar
import weakref
T = TypeVar("T") T = TypeVar("T")
DEFAULT_CACHE_MAX_ITEMS = 10000
@dataclass @dataclass
class CacheData(Generic[T]): class CacheData(Generic[T]):
@@ -16,16 +19,21 @@ class CacheData(Generic[T]):
class CacheDict(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: 字典名称 name: 字典名称
expire: 过期时间(秒),默认为0表示永不过期 expire: 过期时间(秒),默认为0表示永不过期
max_items: 最大缓存项数,None 使用统一默认值,0 表示不限制
""" """
self.name = name.upper() self.name = name.upper()
self.expire = expire 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._data: dict[str, CacheData[T]] = {}
self.__class__._instances.add(self)
def expire_time(self, key: str) -> float: 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 expire_time = time.time() + self.expire if self.expire > 0 else 0
self._data[key] = CacheData(value=value, expire_time=expire_time) self._data[key] = CacheData(value=value, expire_time=expire_time)
self._enforce_limit()
def __delitem__(self, key: str) -> None: def __delitem__(self, key: str) -> None:
"""删除字典项 """删除字典项
@@ -122,6 +131,7 @@ class CacheDict(Generic[T]):
expire_time = time.time() + self.expire expire_time = time.time() + self.expire
self._data[key] = CacheData(value=value, expire_time=expire_time) self._data[key] = CacheData(value=value, expire_time=expire_time)
self._enforce_limit()
def pop(self, key: str, default: Any = None) -> T: def pop(self, key: str, default: Any = None) -> T:
"""删除并返回字典项 """删除并返回字典项
@@ -146,6 +156,32 @@ class CacheDict(Generic[T]):
"""清空字典""" """清空字典"""
self._data.clear() 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]: def keys(self) -> list[str]:
"""获取所有键 """获取所有键
@@ -187,6 +223,12 @@ class CacheDict(Generic[T]):
for key in expired_keys: for key in expired_keys:
del self._data[key] 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: def __len__(self) -> int:
"""获取字典长度 """获取字典长度
@@ -211,17 +253,22 @@ class CacheDict(Generic[T]):
class CacheList(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: 列表名称 name: 列表名称
expire: 过期时间(秒),默认为0表示永不过期 expire: 过期时间(秒),默认为0表示永不过期
max_items: 最大缓存项数,None 使用统一默认值,0 表示不限制
""" """
self.name = name.upper() self.name = name.upper()
self.expire = expire self.expire = expire
self.max_items = DEFAULT_CACHE_MAX_ITEMS if max_items is None else max_items
self._data: list[CacheData[T]] = [] self._data: list[CacheData[T]] = []
self._expire_time = 0 self._expire_time = 0
self.__class__._instances.add(self)
# 如果设置了过期时间,计算整个列表的过期时间 # 如果设置了过期时间,计算整个列表的过期时间
if self.expire > 0: if self.expire > 0:
@@ -303,6 +350,7 @@ class CacheList(Generic[T]):
self.clear() self.clear()
self._data.append(CacheData(value=value)) self._data.append(CacheData(value=value))
self._enforce_limit()
# 更新过期时间 # 更新过期时间
self._update_expire_time() self._update_expire_time()
@@ -318,6 +366,7 @@ class CacheList(Generic[T]):
self.clear() self.clear()
self._data.extend([CacheData(value=v) for v in values]) self._data.extend([CacheData(value=v) for v in values])
self._enforce_limit()
# 更新过期时间 # 更新过期时间
self._update_expire_time() self._update_expire_time()
@@ -334,6 +383,7 @@ class CacheList(Generic[T]):
self.clear() self.clear()
self._data.insert(index, CacheData(value=value)) self._data.insert(index, CacheData(value=value))
self._enforce_limit()
# 更新过期时间 # 更新过期时间
self._update_expire_time() self._update_expire_time()
@@ -389,6 +439,32 @@ class CacheList(Generic[T]):
# 重置过期时间 # 重置过期时间
self._update_expire_time() 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: 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 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: def __str__(self) -> str:
"""字符串表示 """字符串表示
+169 -42
View File
@@ -10,7 +10,13 @@ import uuid
from zhenxun.services.cache.config import CacheMode from zhenxun.services.cache.config import CacheMode
from zhenxun.services.log import logger 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 from zhenxun.utils.manager.priority_manager import PriorityLifecycle
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -18,24 +24,6 @@ if TYPE_CHECKING:
LOG_COMMAND = "RuntimeCache" 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: def _coerce_int(value, default: int) -> int:
try: try:
@@ -45,6 +33,27 @@ def _coerce_int(value, default: int) -> int:
return value_int if value_int >= 0 else default 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 INSTANCE_ID = uuid.uuid4().hex
_CACHE_READY_EVENT = asyncio.Event() _CACHE_READY_EVENT = asyncio.Event()
@@ -89,6 +98,89 @@ def _parse_block_modules(value: str) -> frozenset[str]:
return frozenset(items) 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) @dataclass(frozen=True)
class BanEntry: class BanEntry:
user_id: str | None user_id: str | None
@@ -465,6 +557,20 @@ class RuntimeCacheSync:
@classmethod @classmethod
async def stop(cls) -> None: 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(): if cls._task and not cls._task.done():
cls._task.cancel() cls._task.cancel()
cls._task = None cls._task = None
@@ -480,7 +586,6 @@ class RuntimeCacheSync:
except Exception: except Exception:
pass pass
cls._redis = None cls._redis = None
cls._ready = False
@classmethod @classmethod
def publish_event(cls, cache_type: str, action: str, data: dict[str, Any]) -> None: def publish_event(cls, cache_type: str, action: str, data: dict[str, Any]) -> None:
@@ -559,25 +664,43 @@ class RuntimeCacheSync:
class PluginInfoMemoryCache: class PluginInfoMemoryCache:
_lock: ClassVar[asyncio.Lock] = asyncio.Lock() _lock: ClassVar[asyncio.Lock] = asyncio.Lock()
_by_module: ClassVar[dict[str, "PluginInfo"]] = {} _by_module: ClassVar[dict[str, PluginInfoSnapshot]] = {}
_by_module_path: ClassVar[dict[str, "PluginInfo"]] = {} _by_module_path: ClassVar[dict[str, PluginInfoSnapshot]] = {}
_loaded: ClassVar[bool] = False _loaded: ClassVar[bool] = False
_refresh_task: ClassVar[asyncio.Task | None] = None _refresh_task: ClassVar[asyncio.Task | None] = None
_last_refresh: ClassVar[float] = 0.0 _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 @classmethod
async def refresh(cls) -> None: async def refresh(cls) -> None:
from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_info import PluginInfo
async with cls._lock: async with cls._lock:
plugins = await PluginInfo.all() plugins = await PluginInfo.all()
by_module: dict[str, "PluginInfo"] = {} by_module: dict[str, PluginInfoSnapshot] = {}
by_module_path: dict[str, "PluginInfo"] = {} by_module_path: dict[str, PluginInfoSnapshot] = {}
for plugin in plugins: for plugin in plugins:
if plugin.module: snapshot = PluginInfoSnapshot.from_model(plugin)
by_module[plugin.module] = plugin if snapshot.module:
if plugin.module_path: by_module[snapshot.module] = snapshot
by_module_path[plugin.module_path] = plugin if snapshot.module_path:
by_module_path[snapshot.module_path] = snapshot
cls._by_module = by_module cls._by_module = by_module
cls._by_module_path = by_module_path cls._by_module_path = by_module_path
cls._loaded = True cls._loaded = True
@@ -596,42 +719,42 @@ class PluginInfoMemoryCache:
async def get_by_module(cls, module: str) -> "PluginInfo | None": async def get_by_module(cls, module: str) -> "PluginInfo | None":
if not cls._loaded: if not cls._loaded:
await cls.ensure_loaded() await cls.ensure_loaded()
return cls._by_module.get(module) return cls._to_model(cls._by_module.get(module))
@classmethod @classmethod
async def get_all(cls) -> dict[str, "PluginInfo"]: async def get_all(cls) -> dict[str, "PluginInfo"]:
if not cls._loaded: if not cls._loaded:
await cls.ensure_loaded() await cls.ensure_loaded()
return dict(cls._by_module) return {
module: snapshot.to_model() for module, snapshot in cls._by_module.items()
}
@classmethod @classmethod
def get_by_module_path(cls, module_path: str) -> "PluginInfo | None": 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 @classmethod
def set_plugin(cls, plugin) -> None: def set_plugin(cls, plugin) -> None:
if not plugin: if not plugin:
return return
if plugin.module: snapshot = PluginInfoSnapshot.from_model(plugin)
cls._by_module[plugin.module] = plugin cls._store_snapshot(snapshot)
if getattr(plugin, "module_path", None):
cls._by_module_path[plugin.module_path] = plugin
cls._loaded = True cls._loaded = True
cls._last_refresh = time.time() cls._last_refresh = time.time()
@classmethod @classmethod
def remove_by_module(cls, module: str) -> None: 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 @classmethod
async def upsert_from_model(cls, plugin) -> None: async def upsert_from_model(cls, plugin) -> None:
if not plugin: if not plugin:
return return
async with cls._lock: async with cls._lock:
if getattr(plugin, "module", None): snapshot = PluginInfoSnapshot.from_model(plugin)
cls._by_module[plugin.module] = plugin cls._store_snapshot(snapshot)
if getattr(plugin, "module_path", None):
cls._by_module_path[plugin.module_path] = plugin
cls._loaded = True cls._loaded = True
cls._last_refresh = time.time() cls._last_refresh = time.time()
@@ -643,9 +766,13 @@ class PluginInfoMemoryCache:
return return
async with cls._lock: async with cls._lock:
if module: 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: 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 @classmethod
async def _refresh_loop(cls, interval: int) -> None: async def _refresh_loop(cls, interval: int) -> None:
+19 -58
View File
@@ -10,7 +10,11 @@ T = TypeVar("T", bound=Model)
class DataAccess(Generic[T]): class DataAccess(Generic[T]):
"""数据访问层,根据配置决定是否使用缓存 """数据访问兼容层,根据配置保留单点缓存读取和清理能力
新的高频运行态路径应优先使用 RuntimeCache 或 BoundedTTLCache。
这里不再把 filter/all/create/update_or_create 结果写入通用缓存,
create/update_or_create 只负责清理旧缓存,避免旧值残留。
使用示例: 使用示例:
```python ```python
@@ -395,34 +399,25 @@ class DataAccess(Generic[T]):
return COMPOSITE_KEY_SEPARATOR.join(key_parts) return COMPOSITE_KEY_SEPARATOR.join(key_parts)
async def _cache_items(self, data_list: list[T]) -> None: async def _invalidate_item_cache(self, item: T, action: str) -> None:
"""将数据列表存入缓存 if not self.cache_type or cache_config.cache_mode == CacheMode.NONE:
参数:
data_list: 数据列表
"""
if (
not data_list
or not self.cache_type
or cache_config.cache_mode == CacheMode.NONE
):
return return
try: try:
# 遍历数据列表,将每条数据存入缓存 cache_key = self._build_cache_key_for_item(item)
cached_count = 0 if cache_key is None:
for item in data_list: return
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
await self.cache.delete(cache_key)
self._cache_stats[self.cache_type]["deletes"] += 1
logger.debug( logger.debug(
f"{self.model_cls.__name__} 批量缓存: {cached_count}/{len(data_list)}项" f"{self.model_cls.__name__} {action}: 已失效兼容缓存: {cache_key}"
) )
except Exception as e: 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]: 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)}" f"{self.model_cls.__name__} filter: 查询结果数量: {len(data_list)}"
) )
# 将数据存入缓存
await self._cache_items(data_list)
return data_list return data_list
async def all(self) -> list[T]: async def all(self) -> list[T]:
@@ -457,9 +449,6 @@ class DataAccess(Generic[T]):
data_list = await self.model_cls.all() data_list = await self.model_cls.all()
logger.debug(f"{self.model_cls.__name__} all: 查询结果数量: {len(data_list)}") logger.debug(f"{self.model_cls.__name__} all: 查询结果数量: {len(data_list)}")
# 将数据存入缓存
await self._cache_items(data_list)
return data_list return data_list
async def count(self, *args, **kwargs) -> int: async def count(self, *args, **kwargs) -> int:
@@ -501,24 +490,7 @@ class DataAccess(Generic[T]):
logger.debug(f"{self.model_cls.__name__} create: 创建数据, 参数: {kwargs}") logger.debug(f"{self.model_cls.__name__} create: 创建数据, 参数: {kwargs}")
data = await self.model_cls.create(**kwargs) data = await self.model_cls.create(**kwargs)
# 如果有缓存类型,将数据存入缓存 await self._invalidate_item_cache(data, "create")
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,
)
return data return data
@@ -539,18 +511,7 @@ class DataAccess(Generic[T]):
defaults=defaults, **kwargs defaults=defaults, **kwargs
) )
# 如果有缓存类型,将数据存入缓存 await self._invalidate_item_cache(data, "update_or_create")
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)
return data, created return data, created
+1 -1
View File
@@ -150,7 +150,7 @@ class Model(TortoiseModel):
obj = await cls.filter(**kwargs).using_db(connection).get() obj = await cls.filter(**kwargs).using_db(connection).get()
result = (obj, False) 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])) await CacheRoot.invalidate_cache(cache_type, cls.get_cache_key(result[0]))
return result return result
+6 -3
View File
@@ -5,10 +5,9 @@ import ujson as json
from zhenxun.configs.config import Config from zhenxun.configs.config import Config
from zhenxun.models.group_plugin_setting import GroupPluginSetting 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.data_access import DataAccess
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.enum import CacheType
from zhenxun.utils.pydantic_compat import model_dump, model_validate, parse_as from zhenxun.utils.pydantic_compat import model_dump, model_validate, parse_as
T = TypeVar("T", bound=BaseModel) T = TypeVar("T", bound=BaseModel)
@@ -22,7 +21,11 @@ class GroupSettingsService:
def __init__(self): def __init__(self):
self.dao = DataAccess(GroupPluginSetting) 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 @staticmethod
def _build_cache_key(group_id: str, plugin_name: str) -> str: def _build_cache_key(group_id: str, plugin_name: str) -> str:
+307
View File
@@ -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"
+11
View File
@@ -3,6 +3,17 @@ from __future__ import annotations
import time import time
_OVERLOAD_UNTIL = 0.0 _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: def signal_overload(duration: float = 5.0) -> None:
+62 -23
View File
@@ -20,6 +20,12 @@ from zhenxun.services.log import logger
from .types import BaseScreenshotEngine from .types import BaseScreenshotEngine
_PLAYWRIGHT_DISCONNECT_ERROR = "Connection closed while reading from the driver" _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" _UNRETRIEVED_FUTURE_MESSAGE = "Future exception was never retrieved"
_LOOP_EXCEPTION_FILTER_STATE_ATTR = "_zhenxun_playwright_exception_filter_state" _LOOP_EXCEPTION_FILTER_STATE_ATTR = "_zhenxun_playwright_exception_filter_state"
_DISCONNECT_SUPPRESSION_WINDOW_SECONDS = 10.0 _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( def _get_loop_exception_filter_state(
loop: asyncio.AbstractEventLoop, loop: asyncio.AbstractEventLoop,
) -> dict[str, Any] | None: ) -> dict[str, Any] | None:
@@ -1103,29 +1117,54 @@ class PlaywrightEngine(BaseScreenshotEngine):
template_path: str, template_path: str,
render_options: dict[str, Any], render_options: dict[str, Any],
) -> bytes: ) -> bytes:
generation, context = await self._acquire_context() last_error: Exception | None = None
page = None for attempt in range(2):
broken = False generation, context = await self._acquire_context()
try: page = None
page = await context.new_page() broken = False
page_options = self._build_page_options(render_options, pooled=True) try:
viewport = page_options.get("viewport") page = await context.new_page()
if isinstance(viewport, dict): page_options = self._build_page_options(render_options, pooled=True)
width = viewport.get("width") viewport = page_options.get("viewport")
height = viewport.get("height") if isinstance(viewport, dict):
if isinstance(width, int) and isinstance(height, int): width = viewport.get("width")
await page.set_viewport_size({"width": width, "height": height}) height = viewport.get("height")
return await self._render_with_page( if isinstance(width, int) and isinstance(height, int):
page, html, template_path, render_options await page.set_viewport_size({"width": width, "height": height})
) return await self._render_with_page(
except Exception: page, html, template_path, render_options
broken = True )
raise except Exception as e:
finally: broken = True
if page is not None: last_error = e
with contextlib.suppress(Exception): if attempt == 0:
await page.close() if _is_playwright_target_closed_error(e):
await self._release_context(generation, context, broken=broken) 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( async def _render_html(
self, self,
+9 -55
View File
@@ -1,11 +1,9 @@
from __future__ import annotations from __future__ import annotations
import asyncio
from collections import OrderedDict
import hashlib import hashlib
import time
from typing import Any from typing import Any
from zhenxun.services.cache.bounded_ttl import BoundedTTLCache
from zhenxun.utils.pydantic_compat import dump_json_safely 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 if isinstance(max_total_bytes, int) and max_total_bytes > 0
else None else None
) )
self._cache: OrderedDict[str, tuple[float, bytes]] = OrderedDict() self._cache = BoundedTTLCache[str, bytes](
self._total_bytes = 0 "RENDER_RESULT",
self._lock = asyncio.Lock() ttl_seconds=self._ttl_seconds,
max_items=self._max_items,
max_total_bytes=self._max_total_bytes,
)
@staticmethod @staticmethod
def build_key(payload: Any) -> str: def build_key(payload: Any) -> str:
@@ -37,55 +38,8 @@ class RenderResultMemoryCache:
) )
return hashlib.sha256(payload_text.encode("utf-8")).hexdigest() 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: async def get(self, key: str) -> bytes | None:
now = time.monotonic() return await self._cache.get(key)
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
async def set(self, key: str, value: bytes) -> None: async def set(self, key: str, value: bytes) -> None:
value_size = len(value) await self._cache.set(key, 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)
+12
View File
@@ -475,6 +475,18 @@ class RendererService:
raise RuntimeError("ThemeManager尚未初始化。") raise RuntimeError("ThemeManager尚未初始化。")
return self._theme_manager.list_available_themes() 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: async def switch_theme(self, theme_name: str) -> str:
""" """
切换UI主题,加载新主题并持久化配置。 切换UI主题,加载新主题并持久化配置。
+18 -1
View File
@@ -51,8 +51,10 @@ class ManifestRegistry:
self._manifest_cache: dict[str, TemplateManifest] = {} self._manifest_cache: dict[str, TemplateManifest] = {}
self._lock = asyncio.Lock() self._lock = asyncio.Lock()
def clear_cache(self): def clear_cache(self) -> int:
size = len(self._manifest_cache)
self._manifest_cache.clear() self._manifest_cache.clear()
return size
async def get_manifest( async def get_manifest(
self, component_path: str, skin: str | None = None self, component_path: str, skin: str | None = None
@@ -362,6 +364,21 @@ class ThemeManager:
tuple[type, str, str | None], ComponentDependency tuple[type, str, str | None], ComponentDependency
] = OrderedDict() ] = 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 @staticmethod
def _get_lru_entry(cache: OrderedDict, key: Any) -> Any: def _get_lru_entry(cache: OrderedDict, key: Any) -> Any:
value = cache.get(key) value = cache.get(key)
+63 -3
View File
@@ -2,10 +2,17 @@ import asyncio
from concurrent.futures import ThreadPoolExecutor from concurrent.futures import ThreadPoolExecutor
import contextlib import contextlib
import os import os
import signal
import anyio.to_thread 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 from zhenxun.utils.manager.priority_manager import PriorityLifecycle
DEFAULT_EXECUTOR_MIN_WORKERS = 16 DEFAULT_EXECUTOR_MIN_WORKERS = 16
@@ -14,6 +21,7 @@ DEFAULT_ANYIO_MIN_TOKENS = 32
DEFAULT_ANYIO_MAX_TOKENS = 128 DEFAULT_ANYIO_MAX_TOKENS = 128
_thread_executor: ThreadPoolExecutor | None = None _thread_executor: ThreadPoolExecutor | None = None
_launcher_watchdog_task: asyncio.Task[None] | None = None
_runtime_hooks_registered = False _runtime_hooks_registered = False
_alconna_patch_applied = False _alconna_patch_applied = False
@@ -57,14 +65,60 @@ def _apply_alconna_conflict_patch() -> None:
_alconna_patch_applied = True _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_alconna_conflict_patch()
apply_uninfo_onebot11_patch()
global _runtime_hooks_registered global _runtime_hooks_registered
if _runtime_hooks_registered: if _runtime_hooks_registered:
return return
_runtime_hooks_registered = True _runtime_hooks_registered = True
@driver.on_startup @PriorityLifecycle.on_startup(priority=-100)
async def _setup_runtime_concurrency() -> None: async def _setup_runtime_concurrency() -> None:
global _thread_executor global _thread_executor
workers = _get_executor_workers() workers = _get_executor_workers()
@@ -77,10 +131,16 @@ def register_runtime_bootstrap(driver: Driver) -> None:
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
limiter = anyio.to_thread.current_default_thread_limiter() limiter = anyio.to_thread.current_default_thread_limiter()
limiter.total_tokens = _get_anyio_tokens(workers) limiter.total_tokens = _get_anyio_tokens(workers)
_start_launcher_watchdog()
await start_send_queue()
await start_memory_governor()
@PriorityLifecycle.on_shutdown(priority=50) @PriorityLifecycle.on_shutdown(priority=50)
async def _shutdown_runtime_concurrency() -> None: async def _shutdown_runtime_concurrency() -> None:
global _thread_executor global _thread_executor
await _stop_launcher_watchdog()
await stop_send_queue()
await stop_memory_governor()
executor = _thread_executor executor = _thread_executor
_thread_executor = None _thread_executor = None
if executor is not None: if executor is not None:
+91 -13
View File
@@ -2,21 +2,29 @@ import asyncio
import time import time
from typing import Any from typing import Any
import nonebot
from nonebot.adapters import Bot from nonebot.adapters import Bot
from zhenxun.services.log import logger from zhenxun.services.log import logger
_SEND_APIS = {"send_msg", "send_like"} _SEND_APIS = {"send_msg", "send_group_msg", "send_private_msg", "send_like"}
_QUEUE: asyncio.Queue[tuple[Bot, str, dict[str, Any], asyncio.Future]] = asyncio.Queue()
_WORKERS = 3 _WORKERS = 3
_MIN_INTERVAL = 0.05 _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() _SEND_LOCK = asyncio.Lock()
_LAST_SEND_TS = 0.0 _LAST_SEND_TS = 0.0
_API_SEMAPHORE = asyncio.Semaphore(3) _API_SEMAPHORE = asyncio.Semaphore(3)
_ORIG_CALL_API = Bot.call_api _ORIG_CALL_API = Bot.call_api
_PATCHED = False _PATCHED = False
_WORKER_TASKS: list[asyncio.Task] = [] _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(): async def _rate_limit():
@@ -29,15 +37,36 @@ async def _rate_limit():
_LAST_SEND_TS = time.monotonic() _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): async def _worker(worker_id: int):
while True: while True:
bot, api, data, future = await _QUEUE.get() bot, api, data, future = await _QUEUE.get()
try: try:
await _rate_limit() result = await _direct_call_api(bot, api, data)
async with _API_SEMAPHORE:
result = await _ORIG_CALL_API(bot, api, **data)
if not future.done(): if not future.done():
future.set_result(result) future.set_result(result)
except asyncio.CancelledError:
if not future.done():
future.set_exception(RuntimeError("send queue worker cancelled"))
raise
except Exception as exc: except Exception as exc:
if not future.done(): if not future.done():
future.set_exception(exc) 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): async def _queued_call_api(self: Bot, api: str, **data: Any):
if api not in _SEND_APIS: if api not in _SEND_APIS:
return await _ORIG_CALL_API(self, api, **data) return await _ORIG_CALL_API(self, api, **data)
if _STOPPING:
return await _direct_call_api(self, api, data)
loop = asyncio.get_running_loop() loop = asyncio.get_running_loop()
future: asyncio.Future = loop.create_future() future: asyncio.Future[Any] = loop.create_future()
await _QUEUE.put((self, api, data, 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 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: def patch_send_queue() -> None:
global _PATCHED global _PATCHED
if _PATCHED: if _PATCHED:
@@ -68,21 +126,41 @@ def patch_send_queue() -> None:
_PATCHED = True _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() -> None:
async def _start_send_queue(): global _STOPPING
patch_send_queue() patch_send_queue()
_STOPPING = False
if _WORKER_TASKS:
return
for idx in range(_WORKERS): for idx in range(_WORKERS):
_WORKER_TASKS.append(asyncio.create_task(_worker(idx))) _WORKER_TASKS.append(asyncio.create_task(_worker(idx)))
@driver.on_shutdown async def stop_send_queue() -> None:
async def _stop_send_queue(): 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() tasks = _WORKER_TASKS.copy()
_WORKER_TASKS.clear() _WORKER_TASKS.clear()
for task in tasks: for task in tasks:
task.cancel() task.cancel()
if tasks: if tasks:
await asyncio.gather(*tasks, return_exceptions=True) await asyncio.gather(*tasks, return_exceptions=True)
unpatch_send_queue()
_STOPPING = False
+149
View File
@@ -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")
-2
View File
@@ -55,8 +55,6 @@ class CacheType(StrEnum):
"""全局全部群组""" """全局全部群组"""
GROUP_PLUGIN_SETTINGS = "GROUP_PLUGIN_SETTINGS" GROUP_PLUGIN_SETTINGS = "GROUP_PLUGIN_SETTINGS"
"""插件分群配置""" """插件分群配置"""
GROUP_PLUGIN_SETTINGS_VIEW = "GROUP_PLUGIN_SETTINGS_VIEW"
"""插件分群配置视图缓存(聚合 dict)"""
USERS = "GLOBAL_ALL_USERS" USERS = "GLOBAL_ALL_USERS"
"""全部用户""" """全部用户"""
BAN = "GLOBAL_ALL_BAN" BAN = "GLOBAL_ALL_BAN"
+12 -34
View File
@@ -1,5 +1,4 @@
import asyncio import asyncio
from collections import OrderedDict
from collections.abc import AsyncGenerator, Awaitable, Callable, Sequence from collections.abc import AsyncGenerator, Awaitable, Callable, Sequence
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
import os import os
@@ -21,6 +20,7 @@ from rich.progress import (
import ujson as json import ujson as json
from zhenxun.configs.config import BotConfig from zhenxun.configs.config import BotConfig
from zhenxun.services.cache.bounded_ttl import BoundedTTLCache
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.decorator.retry import Retry from zhenxun.utils.decorator.retry import Retry
from zhenxun.utils.exception import AllURIsFailedError from zhenxun.utils.exception import AllURIsFailedError
@@ -168,7 +168,12 @@ class AsyncHttpx:
_CONTENT_CACHE_TTL: ClassVar[float] = 3.0 _CONTENT_CACHE_TTL: ClassVar[float] = 3.0
_CONTENT_CACHE_MAX_ITEMS: ClassVar[int] = 256 _CONTENT_CACHE_MAX_ITEMS: ClassVar[int] = 256
_CONTENT_CACHE_MAX_BYTES: ClassVar[int] = 2 * 1024 * 1024 _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_inflight: ClassVar[dict[str, asyncio.Task[Response]]] = {}
_content_cache_lock: ClassVar[asyncio.Lock] = asyncio.Lock() _content_cache_lock: ClassVar[asyncio.Lock] = asyncio.Lock()
@@ -191,29 +196,6 @@ class AsyncHttpx:
return True return True
return "qpic.cn" in lower_url or "qlogo.cn" in lower_url 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 @classmethod
async def _try_cache_content(cls, key: str, response: Response) -> None: async def _try_cache_content(cls, key: str, response: Response) -> None:
content = response.content content = response.content
@@ -223,13 +205,7 @@ class AsyncHttpx:
is_image = content_type.startswith("image/") or cls._is_probably_image_url(key) is_image = content_type.startswith("image/") or cls._is_probably_image_url(key)
if not is_image: if not is_image:
return return
async with cls._content_cache_lock: await cls._content_cache.set(key, content)
cls._content_cache[key] = (
time.monotonic() + cls._CONTENT_CACHE_TTL,
content,
)
cls._content_cache.move_to_end(key)
cls._cleanup_content_cache_nolock()
@classmethod @classmethod
def _prepare_temporary_client_config(cls, client_kwargs: dict) -> dict: def _prepare_temporary_client_config(cls, client_kwargs: dict) -> dict:
@@ -450,9 +426,11 @@ class AsyncHttpx:
return res.content return res.content
cache_key = url cache_key = url
if cached := await cls._content_cache.get(cache_key):
return cached
async with cls._content_cache_lock: async with cls._content_cache_lock:
cached = cls._get_cached_content_nolock(cache_key) if cached := await cls._content_cache.get(cache_key):
if cached is not None:
return cached return cached
task = cls._content_inflight.get(cache_key) task = cls._content_inflight.get(cache_key)
if task is None: if task is None:
+60 -8
View File
@@ -1,25 +1,77 @@
from collections import OrderedDict
import time
from typing import ClassVar from typing import ClassVar
class MessageManager: 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 @classmethod
def add(cls, uid: str, msg_id: str): def add(cls, uid: str, msg_id: str):
if uid not in cls.data: now = time.monotonic()
cls.data[uid] = [] cls._prune(now)
cls.data[uid].append(msg_id) _, messages = cls.data.get(uid, (now, []))
messages.append(msg_id)
cls._touch(uid, messages, now)
cls.remove_check(uid) cls.remove_check(uid)
cls._prune(now)
@classmethod @classmethod
def check(cls, uid: str, msg_id: str) -> bool: 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 @classmethod
def remove_check(cls, uid: str): def remove_check(cls, uid: str):
if len(cls.data[uid]) > 200: entry = cls.data.get(uid)
cls.data[uid] = cls.data[uid][100:] 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 @classmethod
def get(cls, uid: str) -> list[str]: 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
@@ -8,6 +8,7 @@ from zhenxun.configs.config import Config
from zhenxun.services.log import logger from zhenxun.services.log import logger
LOG_COMMAND = "VirtualEnvPackageManager" LOG_COMMAND = "VirtualEnvPackageManager"
PROJECT_ROOT = Path(__file__).resolve().parents[3]
Config.add_plugin_config( Config.add_plugin_config(
"virtualenv", "virtualenv",
@@ -191,6 +192,46 @@ class VirtualEnvPackageManager:
) )
return stderr 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 @classmethod
async def list(cls) -> str: async def list(cls) -> str:
"""列出已安装的依赖包""" """列出已安装的依赖包"""