性能优化 (#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
@@ -7,9 +7,10 @@ from zhenxun.models.level_user import LevelUser
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.cache.runtime_cache import LevelUserMemoryCache, LevelUserSnapshot
from zhenxun.services.log import logger
from zhenxun.utils.utils import get_entity_ids
from zhenxun.utils.utils import EntityIDs, get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import SkipPluginException
@@ -20,6 +21,9 @@ async def auth_admin(
LevelUser | LevelUserSnapshot | None, LevelUser | LevelUserSnapshot | None
]
| None = None,
*,
context: PermissionContext | None = None,
entity: EntityIDs | None = None,
):
"""管理员命令 个人权限
@@ -33,7 +37,12 @@ async def auth_admin(
return
try:
entity = get_entity_ids(session)
if context is not None:
entity = context.entity
if cached_levels is None:
cached_levels = context.admin_levels
if entity is None:
entity = get_entity_ids(session)
global_user: LevelUser | LevelUserSnapshot | None = None
group_users: LevelUser | LevelUserSnapshot | None = None
@@ -42,7 +51,7 @@ async def auth_admin(
global_user, group_users = cached_levels
else:
global_user, group_users = await LevelUserMemoryCache.get_levels(
session.user.id, entity.group_id
entity.user_id, entity.group_id
)
user_level = global_user.user_level if global_user else 0
@@ -53,7 +62,7 @@ async def auth_admin(
raise SkipPluginException(
f"{plugin.name}({plugin.module}) 管理员权限不足...",
tip_message=[
At(flag="user", target=session.user.id),
At(flag="user", target=entity.user_id),
f"你的权限不足喔,该功能需要的权限等级: {plugin.admin_level}",
],
tip_check_tag=entity.user_id,
+7 -48
View File
@@ -1,4 +1,3 @@
import asyncio
import time
from nonebot.matcher import Matcher
@@ -8,7 +7,6 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config
from zhenxun.models.ban_console import BanConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.cache.cache_containers import CacheDict
from zhenxun.services.cache.runtime_cache import BanMemoryCache
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger
@@ -16,6 +14,7 @@ from zhenxun.utils.enum import PluginType
from zhenxun.utils.utils import EntityIDs, get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import SkipPluginException
from .utils import freq
@@ -25,37 +24,6 @@ Config.add_plugin_config(
"才不会给你发消息.",
help="对被ban用户发送的消息",
)
BAN_CACHE_TTL = 2
BAN_CACHE_TTL_POSITIVE = 30
BAN_CACHE_TTL_NEGATIVE = 5
BAN_CACHE = (
CacheDict("AUTH_BAN_CACHE", expire=0)
if max(BAN_CACHE_TTL_POSITIVE, BAN_CACHE_TTL_NEGATIVE) > 0
else None
)
def _ban_cache_key(user_id: str | None, group_id: str | None) -> str:
return f"{user_id or ''}:{group_id or ''}"
def _ban_cache_get(key: str) -> int | None:
if not BAN_CACHE:
return None
try:
return BAN_CACHE[key]
except KeyError:
return None
def _ban_cache_set(key: str, value: int) -> None:
if not BAN_CACHE:
return
ttl = BAN_CACHE_TTL_POSITIVE if value else BAN_CACHE_TTL_NEGATIVE
if ttl <= 0:
return
BAN_CACHE.set(key, value, expire=ttl)
async def calculate_ban_time(ban_record: BanConsole | None) -> int:
@@ -214,6 +182,7 @@ async def auth_ban(
session: Uninfo,
plugin: PluginInfo,
*,
context: PermissionContext | None = None,
entity: EntityIDs | None = None,
is_superuser: bool = False,
) -> None:
@@ -229,28 +198,18 @@ async def auth_ban(
return
if not matcher.plugin_name:
return
if context is not None:
entity = context.entity
is_superuser = context.is_superuser
if entity is None:
entity = get_entity_ids(session)
if is_superuser:
return
if entity.group_id:
try:
await asyncio.wait_for(
group_handle(entity.group_id), timeout=DB_TIMEOUT_SECONDS
)
except asyncio.TimeoutError:
logger.error(f"群组ban检查超时: {entity.group_id}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
await group_handle(entity.group_id)
if entity.user_id:
try:
await asyncio.wait_for(
user_handle(plugin, entity, session),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error(f"用户ban检查超时: {entity.user_id}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
await user_handle(plugin, entity, session)
finally:
# 记录总执行时间
elapsed = time.time() - start_time
@@ -7,6 +7,7 @@ from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import SkipPluginException
@@ -16,6 +17,8 @@ async def auth_bot(
bot_data: BotConsole | BotSnapshot | None = None,
skip_fetch: bool = False,
allow_sleep_bypass: bool = False,
*,
context: PermissionContext | None = None,
):
"""bot层面的权限检查
@@ -30,6 +33,9 @@ async def auth_bot(
start_time = time.time()
try:
if context is not None:
bot_id = context.event.bot_id
bot_data = context.bot_data
bot: BotConsole | BotSnapshot | None = bot_data
if bot is None and not skip_fetch:
bot = await BotMemoryCache.get(bot_id)
@@ -7,13 +7,18 @@ from zhenxun.models.user_console import UserConsole
from zhenxun.services.log import logger
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import SkipPluginException
DEFAULT_GOLD = 100
async def auth_cost(
user: UserConsole | None, plugin: PluginInfo, session: Uninfo
user: UserConsole | None,
plugin: PluginInfo,
session: Uninfo,
*,
context: PermissionContext | None = None,
) -> int:
"""检测是否满足金币条件
@@ -28,6 +33,8 @@ async def auth_cost(
start_time = time.time()
try:
if context is not None and user is None:
user = context.user
user_gold = user.gold if user else DEFAULT_GOLD
if user_gold < plugin.cost_gold:
"""插件消耗金币不足"""
@@ -7,6 +7,7 @@ from zhenxun.services.cache.runtime_cache import GroupSnapshot
from zhenxun.services.log import logger
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
from .context import PermissionContext
from .exception import SkipPluginException
_GROUP_WAKE_PATTERN = re.compile(r"^醒来$", re.IGNORECASE)
@@ -34,6 +35,8 @@ async def auth_group(
group: GroupConsole | GroupSnapshot | None,
text: str | None,
group_id: str | None,
*,
context: PermissionContext | None = None,
):
"""群黑名单检测 群总开关检测
@@ -42,6 +45,11 @@ async def auth_group(
group: GroupConsole
message: UniMsg
"""
if context is not None:
group = context.group or group
text = context.plain_text
group_id = context.group_id
if not group_id:
return
@@ -19,9 +19,10 @@ from zhenxun.utils.limiters import CountLimiter, FreqLimiter, UserBlockLimiter
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.time_utils import TimeUtils
from zhenxun.utils.utils import get_entity_ids
from zhenxun.utils.utils import EntityIDs, get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import SkipPluginException
driver = nonebot.get_driver()
@@ -31,7 +32,7 @@ _LIMIT_NOTICE_LIMITER = FreqLimiter(_LIMIT_NOTICE_CD)
_LIMIT_NOTICE_TASKS: set[asyncio.Task] = set()
@PriorityLifecycle.on_startup(priority=5)
@PriorityLifecycle.on_startup(priority=7)
async def _():
"""初始化限制"""
await LimitManager.init_limit()
@@ -117,6 +118,7 @@ class LimitManager:
cls.cd_limit = {}
cls.block_limit = {}
cls.count_limit = {}
cls.module_limit_cache.clear()
# 添加新数据
for limit in limit_list:
cls.add_limit(limit)
@@ -137,22 +139,22 @@ class LimitManager:
"""
if limit.module not in cls.add_module:
cls.add_module.append(limit.module)
if limit.limit_type == PluginLimitType.BLOCK:
cls.block_limit[limit.module] = Limit(
limit=limit, limiter=UserBlockLimiter()
)
elif limit.limit_type == PluginLimitType.CD:
cd_value = int(limit.cd or 0)
cls.cd_limit[limit.module] = Limit(
limit=limit, limiter=FreqLimiter(cd_value)
)
elif limit.limit_type == PluginLimitType.COUNT:
max_count = int(limit.max_count or 0)
if max_count <= 0:
return
cls.count_limit[limit.module] = Limit(
limit=limit, limiter=CountLimiter(max_count)
)
if limit.limit_type == PluginLimitType.BLOCK:
cls.block_limit[limit.module] = Limit(
limit=limit, limiter=UserBlockLimiter()
)
elif limit.limit_type == PluginLimitType.CD:
cd_value = int(limit.cd or 0)
cls.cd_limit[limit.module] = Limit(
limit=limit, limiter=FreqLimiter(cd_value)
)
elif limit.limit_type == PluginLimitType.COUNT:
max_count = int(limit.max_count or 0)
if max_count <= 0:
return
cls.count_limit[limit.module] = Limit(
limit=limit, limiter=CountLimiter(max_count)
)
@classmethod
def unblock(
@@ -322,14 +324,23 @@ class LimitManager:
limiter.increase(key_type)
async def auth_limit(plugin: PluginInfo, session: Uninfo):
async def auth_limit(
plugin: PluginInfo,
session: Uninfo,
*,
context: PermissionContext | None = None,
entity: EntityIDs | None = None,
):
"""插件限制
参数:
plugin: PluginInfo
session: Uninfo
"""
entity = get_entity_ids(session)
if context is not None:
entity = context.entity
if entity is None:
entity = get_entity_ids(session)
try:
await asyncio.wait_for(
LimitManager.check(
@@ -10,6 +10,7 @@ from zhenxun.services.log import logger
from zhenxun.utils.enum import BlockType
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .context import PermissionContext
from .exception import IsSuperuserException, SkipPluginException
from .utils import freq, is_poke
@@ -107,11 +108,16 @@ class GroupCheck:
class PluginCheck:
def __init__(
self, group: GroupConsole | GroupSnapshot | None, session: Uninfo, is_poke: bool
self,
group: GroupConsole | GroupSnapshot | None,
session: Uninfo,
is_poke: bool,
user_id: str | None,
):
self.session = session
self.is_poke = is_poke
self.group_data = group
self.user_id = user_id or session.user.id
self.group_id = None
if group:
self.group_id = group.group_id
@@ -126,13 +132,11 @@ class PluginCheck:
IgnoredException: 忽略插件
"""
if plugin.block_type == BlockType.PRIVATE:
should_tip = freq.is_send_limit_message(
plugin, self.session.user.id, self.is_poke
)
should_tip = freq.is_send_limit_message(plugin, self.user_id, self.is_poke)
raise SkipPluginException(
f"{plugin.name}({plugin.module}) 该插件在私聊中已被禁用...",
tip_message="该功能在私聊中已被禁用..." if should_tip else None,
tip_check_tag=self.session.user.id if should_tip else None,
tip_check_tag=self.user_id if should_tip else None,
tip_background=should_tip,
)
@@ -153,7 +157,7 @@ class PluginCheck:
if self.group_data and self.group_data.is_super:
raise IsSuperuserException()
sid = self.group_id or self.session.user.id
sid = self.group_id or self.user_id
should_tip = freq.is_send_limit_message(plugin, sid, self.is_poke)
raise SkipPluginException(
f"{plugin.name}({plugin.module}) 全局未开启此功能...",
@@ -176,7 +180,9 @@ async def auth_plugin(
session: Uninfo,
event: Event,
*,
context: PermissionContext | None = None,
skip_group_block: bool = False,
user_id: str | None = None,
):
"""插件状态
@@ -187,8 +193,11 @@ async def auth_plugin(
"""
start_time = time.time()
try:
if context is not None:
group = context.group or group
user_id = context.user_id
is_poke_event = is_poke(event)
user_check = PluginCheck(group, session, is_poke_event)
user_check = PluginCheck(group, session, is_poke_event, user_id)
if group:
block_set, super_block_set = _get_group_block_sets(group)
@@ -3,6 +3,7 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config
from .context import PermissionContext
from .exception import SkipPluginException
Config.add_plugin_config(
@@ -15,7 +16,12 @@ Config.add_plugin_config(
)
def bot_filter(session: Uninfo):
def bot_filter(
session: Uninfo,
*,
context: PermissionContext | None = None,
user_id: str | None = None,
):
"""过滤bot调用bot
参数:
@@ -26,10 +32,13 @@ def bot_filter(session: Uninfo):
"""
if not Config.get_config("hook", "FILTER_BOT"):
return
if context is not None:
user_id = context.user_id
bot_ids = list(nonebot.get_bots().keys())
if session.user.id == session.self_id:
checked_user_id = user_id or session.user.id
if checked_user_id == session.self_id:
return
if session.user.id in bot_ids:
if checked_user_id in bot_ids:
raise SkipPluginException(
f"bot:{session.self_id} 尝试调用 bot:{session.user.id}"
f"bot:{session.self_id} 尝试调用 bot:{checked_user_id}"
)
@@ -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
from collections.abc import Awaitable, Callable
import contextlib
import os
import importlib
import re
import time
from typing import cast
@@ -11,7 +11,6 @@ from nonebot.adapters import Bot, Event
from nonebot.exception import IgnoredException
from nonebot.matcher import Matcher
import nonebot.message as nb_message
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.utils import PluginExtraData
@@ -33,7 +32,6 @@ from zhenxun.services.message_load import is_overloaded
from zhenxun.utils.enum import BlockType, GoldHandle, PluginType
from zhenxun.utils.exception import InsufficientGold
from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.utils import get_entity_ids
from .auth.auth_admin import auth_admin
from .auth.auth_ban import auth_ban
@@ -44,6 +42,16 @@ from .auth.auth_limit import LimitManager, auth_limit
from .auth.auth_plugin import auth_plugin
from .auth.bot_filter import bot_filter
from .auth.config import LOGGER_COMMAND, WARNING_THRESHOLD
from .auth.context import (
EVENT_CACHE,
STATE_PLAIN_TEXT,
EventContext,
PermissionContext,
get_event_context,
get_permission_side_effect_cache,
set_route_modules,
store_permission_context,
)
from .auth.exception import (
IsSuperuserException,
PermissionExemption,
@@ -53,7 +61,6 @@ from .auth.utils import send_message
AUTH_HOOKS_CONCURRENCY_LIMIT = 5
AUTH_DB_CONCURRENCY_LIMIT = 6
AUTH_EVENT_CACHE_TTL = 5 # 增加到5秒,减少缓存抖动
# 超时设置(秒)
@@ -74,13 +81,6 @@ CIRCUIT_RESET_TIME = 300 # 5分钟
HOOKS_CONCURRENCY_LIMIT = AUTH_HOOKS_CONCURRENCY_LIMIT
DB_CONCURRENCY_LIMIT = AUTH_DB_CONCURRENCY_LIMIT
EVENT_CACHE_TTL = AUTH_EVENT_CACHE_TTL
EVENT_CACHE = (
CacheDict("AUTH_EVENT_CACHE", expire=EVENT_CACHE_TTL)
if EVENT_CACHE_TTL > 0
else None
)
# 路由索引缓存
_ROUTE_INDEX_LOCK = asyncio.Lock()
_ROUTE_INDEX_READY = False
@@ -91,21 +91,17 @@ MATCHER_ROUTE_PREFILTER_TTL = 2
PREFILTER_STATS_LOG_INTERVAL = 10.0
CACHE_SWEEP_INTERVAL = 1.0
CPU_COUNT = os.cpu_count() or 4
COMMAND_MATCHER_CONCURRENCY = max(8, min(48, CPU_COUNT * 4))
HEAVY_COMMAND_CONCURRENCY = max(1, min(3, CPU_COUNT // 2))
HEAVY_COMMAND_MODULES = frozenset({"shop", "sign_in"})
# 全局信号量与计数器
HOOKS_ACTIVE_COUNT = 0
HOOKS_SEMAPHORE = asyncio.Semaphore(HOOKS_CONCURRENCY_LIMIT)
COMMAND_MATCHER_SEMAPHORE = asyncio.Semaphore(COMMAND_MATCHER_CONCURRENCY)
HEAVY_COMMAND_SEMAPHORE = asyncio.Semaphore(HEAVY_COMMAND_CONCURRENCY)
DB_SEMAPHORE = asyncio.Semaphore(DB_CONCURRENCY_LIMIT)
DB_ACTIVE_COUNT = 0
_CHECK_MATCHER_PATCHED = False
_ORIGINAL_CHECK_AND_RUN_MATCHER: Callable[..., Awaitable[None]] | None = None
_HANDLE_EVENT_PATCHED = False
_ORIGINAL_HANDLE_EVENT: Callable[..., Awaitable[None]] | None = None
_ORIGINAL_ADAPTER_HANDLE_EVENTS: dict[object, object] = {}
_MATCHER_COMMAND_TYPE_CACHE: dict[type[Matcher], bool] = {}
_MATCHER_COMMAND_LITERAL_CACHE: dict[type[Matcher], tuple[str, ...] | None] = {}
_MATCHER_ALCONNA_SHORTCUT_CACHE: dict[type[Matcher], bool] = {}
@@ -115,6 +111,10 @@ _CHECK_MATCHER_ROUTE_CACHE = CacheDict(
_PREFILTER_STATS = {
"checked": 0,
"skipped": 0,
"before_task_checked": 0,
"before_task_skipped": 0,
"inside_task_checked": 0,
"inside_task_skipped": 0,
"type_miss": 0,
"route_miss": 0,
"command_miss": 0,
@@ -163,33 +163,6 @@ def _debug_log(message: str, *args, **kwargs) -> None:
logger.debug(message, *args, **kwargs)
def _event_cache_key(event: Event, session: Uninfo, entity) -> str:
msg_id = getattr(event, "message_id", None)
if msg_id is None:
msg_id = getattr(event, "id", None)
if msg_id is None:
msg_id = id(event)
platform = PlatformUtils.get_platform(session)
group_id = entity.group_id or ""
channel_id = entity.channel_id or ""
return (
f"{platform}:{session.self_id}:{entity.user_id}:"
f"{group_id}:{channel_id}:{msg_id}"
)
def _get_event_cache(event: Event, session: Uninfo, entity):
if not EVENT_CACHE:
return None
key = _event_cache_key(event, session, entity)
try:
return EVENT_CACHE[key]
except KeyError:
cache = {}
EVENT_CACHE[key] = cache
return cache
def _normalize_command(command: str) -> str:
text = command.strip()
if not text:
@@ -488,6 +461,9 @@ def _event_plain_text(event: Event) -> str:
def _state_plain_text(state: dict | None) -> str:
if state is None:
return ""
context = get_event_context(state)
if context is not None:
return context.plain_text.strip()
text = state.get("_zx_plain_text")
if isinstance(text, str):
return text.strip()
@@ -496,6 +472,9 @@ def _state_plain_text(state: dict | None) -> str:
def _get_route_modules_for_event(event: Event, state: dict | None = None) -> set[str]:
if state is not None:
context = get_event_context(state)
if context is not None and context.route_modules_loaded:
return context.route_modules
route_modules = state.get("_zx_route_modules")
if isinstance(route_modules, set):
return route_modules
@@ -506,15 +485,67 @@ def _get_route_modules_for_event(event: Event, state: dict | None = None) -> set
route_modules = _match_route_modules(_event_plain_text(event))
_CHECK_MATCHER_ROUTE_CACHE[key] = route_modules
if state is not None:
state["_zx_route_modules"] = route_modules
context = get_event_context(state)
if context is not None:
set_route_modules(state, context, route_modules)
else:
state["_zx_route_modules"] = route_modules
return route_modules
def _record_prefilter_stats(skipped: bool, reason: str | None) -> None:
def _prepare_handle_event_state(event: Event, state: dict) -> None:
get_permission_side_effect_cache(state=state)
if event.get_type() != "message":
return
if _state_plain_text(state):
return
text = _event_plain_text(event)
if text:
state[STATE_PLAIN_TEXT] = text
def _build_matcher_state(base_state: dict) -> dict:
get_permission_side_effect_cache(state=base_state)
matcher_state = base_state.copy()
get_permission_side_effect_cache(state=matcher_state)
return matcher_state
async def _run_selected_matcher(
matcher: type[Matcher],
bot: Bot,
event: Event,
state: dict,
stack,
dependency_cache,
) -> None:
await nb_message.check_and_run_matcher(
matcher,
bot,
event,
state,
stack,
dependency_cache,
)
def _record_prefilter_stats(
skipped: bool,
reason: str | None,
stage: str = "inside_task",
) -> None:
global _PREFILTER_LAST_LOG
_PREFILTER_STATS["checked"] += 1
if skipped:
_PREFILTER_STATS["skipped"] += 1
if stage == "before_task":
_PREFILTER_STATS["before_task_checked"] += 1
if skipped:
_PREFILTER_STATS["before_task_skipped"] += 1
else:
_PREFILTER_STATS["inside_task_checked"] += 1
if skipped:
_PREFILTER_STATS["inside_task_skipped"] += 1
if reason == "type_miss":
_PREFILTER_STATS["type_miss"] += 1
elif reason == "route_miss":
@@ -537,6 +568,10 @@ def _record_prefilter_stats(skipped: bool, reason: str | None) -> None:
"matcher prefilter stats: "
f"checked={_PREFILTER_STATS['checked']} "
f"skipped={_PREFILTER_STATS['skipped']} "
f"before_task={_PREFILTER_STATS['before_task_skipped']}/"
f"{_PREFILTER_STATS['before_task_checked']} "
f"inside_task={_PREFILTER_STATS['inside_task_skipped']}/"
f"{_PREFILTER_STATS['inside_task_checked']} "
f"type_miss={_PREFILTER_STATS['type_miss']} "
f"route_miss={_PREFILTER_STATS['route_miss']} "
f"command_miss={_PREFILTER_STATS['command_miss']} "
@@ -643,15 +678,6 @@ def _matcher_has_alconna_shortcuts(matcher_cls: type[Matcher]) -> bool:
return has_shortcuts
def _is_heavy_command_module(module: str) -> bool:
normalized = module.strip().lower()
if not normalized:
return False
if normalized in HEAVY_COMMAND_MODULES:
return True
return any(normalized.endswith(f".{name}") for name in HEAVY_COMMAND_MODULES)
async def _check_matcher_prefilter(
matcher_cls: type[Matcher], event: Event, state: dict | None = None
) -> tuple[bool, str | None]:
@@ -686,6 +712,17 @@ async def _check_matcher_prefilter(
if not module:
return False, None
command_matched = False
matcher_commands = _extract_matcher_command_literals(matcher_cls)
if matcher_commands:
for command in matcher_commands:
if _command_matches(text, command):
command_matched = True
break
else:
if not _matcher_has_alconna_shortcuts(matcher_cls):
return True, "command_miss"
ai_route_modules = _collect_ai_route_modules(event, state)
ai_route_heads = _collect_ai_route_heads(event, state)
if ai_route_modules and module not in ai_route_modules:
@@ -696,25 +733,85 @@ async def _check_matcher_prefilter(
await _ensure_route_index()
if module not in _ROUTE_MODULES_WITH_COMMANDS:
matcher_commands = _extract_matcher_command_literals(matcher_cls)
if matcher_commands:
for command in matcher_commands:
if _command_matches(text, command):
return False, None
if _matcher_has_alconna_shortcuts(matcher_cls):
return False, None
return True, "command_miss"
return False, None
route_modules = _get_route_modules_for_event(event, state)
if module not in route_modules:
if command_matched:
return False, None
if _matcher_has_alconna_shortcuts(matcher_cls):
return False, None
return True, "route_miss"
return False, None
_MATCHER_SEMAPHORE_TIMEOUT = 8.0
def _check_matcher_prefilter_before_task(
matcher_cls: type[Matcher], event: Event, state: dict | None = None
) -> tuple[bool, str | None]:
"""Conservative selector before creating matcher task.
This mirrors the async matcher prefilter but never performs IO or route-index
rebuild. If anything is uncertain, let the existing check_and_run_matcher
patch handle it inside the task.
"""
event_type = event.get_type()
matcher_type = getattr(matcher_cls, "type", "") or ""
if isinstance(matcher_type, str) and matcher_type and matcher_type != event_type:
return True, "type_miss"
if event_type != "message":
return False, None
if getattr(matcher_cls, "temp", False):
return False, None
if not _is_command_matcher_class(matcher_cls):
return False, None
text = _state_plain_text(state)
if not text:
text = _event_plain_text(event)
if state is not None and text:
state["_zx_plain_text"] = text
if not text:
return True, "empty_text"
module = _matcher_module_name(matcher_cls)
if not module:
return False, None
command_matched = False
matcher_commands = _extract_matcher_command_literals(matcher_cls)
has_alconna_shortcuts = _matcher_has_alconna_shortcuts(matcher_cls)
if matcher_commands:
for command in matcher_commands:
if _command_matches(text, command):
command_matched = True
break
else:
if not has_alconna_shortcuts:
return True, "command_miss"
ai_route_modules = _collect_ai_route_modules(event, state)
ai_route_heads = _collect_ai_route_heads(event, state)
if ai_route_modules and module not in ai_route_modules:
if not _matcher_matches_ai_route_heads(matcher_cls, ai_route_heads):
return True, "route_miss"
if not _ROUTE_INDEX_READY:
return False, None
if module not in _ROUTE_MODULES_WITH_COMMANDS:
return False, None
route_modules = _get_route_modules_for_event(event, state)
if module not in route_modules:
if command_matched or has_alconna_shortcuts:
return False, None
return True, "route_miss"
return False, None
_MAX_MATCHER_CACHE = 512
@@ -729,7 +826,7 @@ async def _patched_check_and_run_matcher(
skip, reason = await _check_matcher_prefilter(
Matcher, event, state if isinstance(state, dict) else None
)
_record_prefilter_stats(skip, reason)
_record_prefilter_stats(skip, reason, "inside_task")
if skip:
return
@@ -744,28 +841,6 @@ async def _patched_check_and_run_matcher(
"stack": stack,
"dependency_cache": dependency_cache,
}
if _is_command_matcher_class(Matcher):
module = _matcher_module_name(Matcher)
sem = (
HEAVY_COMMAND_SEMAPHORE
if _is_heavy_command_module(module)
else COMMAND_MATCHER_SEMAPHORE
)
try:
await asyncio.wait_for(sem.acquire(), timeout=_MATCHER_SEMAPHORE_TIMEOUT)
except asyncio.TimeoutError:
logger.warning(
f"matcher semaphore acquire timeout for {module}, "
"executing without concurrency limit",
LOGGER_COMMAND,
)
await original(**kwargs)
return
try:
await original(**kwargs)
finally:
sem.release()
return
await original(**kwargs)
@@ -788,27 +863,137 @@ def _uninstall_matcher_prefilter() -> None:
_ORIGINAL_CHECK_AND_RUN_MATCHER = None
def _get_message_text(
message: UniMsg | None,
event_cache: dict | None,
event: Event | None = None,
) -> str:
if event_cache is not None:
cached = event_cache.get("plain_text")
if isinstance(cached, str):
return cached
async def _patched_handle_event(bot: Bot, event: Event) -> None:
show_log = True
escape_tag = getattr(nb_message, "escape_tag")
logger_ = getattr(nb_message, "logger")
no_log_exception = getattr(nb_message, "NoLogException")
text = ""
if message is not None:
with contextlib.suppress(Exception):
text = message.extract_plain_text()
if not text and event is not None:
with contextlib.suppress(Exception):
text = (event.get_plaintext() or "").strip()
log_msg = f"<m>{escape_tag(bot.type)} {escape_tag(bot.self_id)}</m> | "
try:
log_msg += event.get_log_string()
except no_log_exception:
show_log = False
if show_log:
logger_.opt(colors=True).success(log_msg)
if event_cache is not None:
event_cache["plain_text"] = text
return text
state = {}
dependency_cache = {}
async_exit_stack = getattr(nb_message, "AsyncExitStack")
apply_event_preprocessors = getattr(nb_message, "_apply_event_preprocessors")
apply_event_postprocessors = getattr(nb_message, "_apply_event_postprocessors")
trie_rule = getattr(nb_message, "TrieRule")
matchers = getattr(nb_message, "matchers")
catch = getattr(nb_message, "catch")
stop_propagation = getattr(nb_message, "StopPropagation")
handle_exception = getattr(nb_message, "_handle_exception")
anyio_mod = getattr(nb_message, "anyio")
run_coro_with_shield = getattr(nb_message, "run_coro_with_shield")
async with async_exit_stack() as stack:
if not await apply_event_preprocessors(
bot=bot,
event=event,
state=state,
stack=stack,
dependency_cache=dependency_cache,
):
return
try:
trie_rule.get_value(bot, event, state)
except Exception as e:
logger_.opt(colors=True, exception=e).warning(
"Error while parsing command for event"
)
_prepare_handle_event_state(event, state)
break_flag = False
def _handle_stop_propagation(_exc_group) -> None:
nonlocal break_flag
break_flag = True
logger_.debug("Stop event propagation")
for priority in sorted(matchers.keys()):
if break_flag:
break
if show_log:
logger_.debug(f"Checking for matchers in priority {priority}...")
if not (priority_matchers := matchers[priority]):
continue
with catch(
{
stop_propagation: _handle_stop_propagation,
Exception: handle_exception(
"<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]:
@@ -843,12 +1028,14 @@ async def start_auth_runtime_tasks() -> None:
global _CACHE_SWEEP_TASK
await _ensure_route_index()
_install_matcher_prefilter()
_install_handle_event_selector()
if _CACHE_SWEEP_TASK is None or _CACHE_SWEEP_TASK.done():
_CACHE_SWEEP_TASK = asyncio.create_task(_cache_sweep_loop())
async def stop_auth_runtime_tasks() -> None:
global _CACHE_SWEEP_TASK
_uninstall_handle_event_selector()
_uninstall_matcher_prefilter()
task = _CACHE_SWEEP_TASK
_CACHE_SWEEP_TASK = None
@@ -873,6 +1060,12 @@ async def _has_limits_cached(module: str, event_cache: dict | None) -> bool:
@contextlib.asynccontextmanager
async def _db_section():
global DB_ACTIVE_COUNT
if DB_SEMAPHORE.locked():
logger.warning(
"db semaphore saturated, allowing permission check to continue",
LOGGER_COMMAND,
)
raise PermissionExemption("db semaphore saturated, allow pass")
await DB_SEMAPHORE.acquire()
DB_ACTIVE_COUNT += 1
try:
@@ -918,7 +1111,9 @@ def _group_has_plugin_block(group, module: str) -> bool:
)
def _needs_auth_plugin(plugin: PluginInfo, group, entity) -> bool:
def _needs_auth_plugin(plugin: PluginInfo, context: PermissionContext) -> bool:
group = context.group
entity = context.entity
if plugin.block_type == BlockType.ALL and not plugin.status:
if group and getattr(group, "is_super", False):
return False
@@ -953,11 +1148,11 @@ async def _get_bot_data_cached(
async def _get_admin_levels_cached(
session: Uninfo, entity, event_cache
entity, event_cache
) -> tuple[tuple[LevelUserSnapshot | None, LevelUserSnapshot | None] | None, bool]:
if event_cache is not None and "admin_levels" in event_cache:
return event_cache.get("admin_levels"), event_cache.get("admin_timeout", False)
levels = await LevelUserMemoryCache.get_levels(session.user.id, entity.group_id)
levels = await LevelUserMemoryCache.get_levels(entity.user_id, entity.group_id)
if event_cache is not None:
event_cache["admin_levels"] = levels
event_cache["admin_timeout"] = False
@@ -1074,12 +1269,18 @@ async def get_plugin_and_user(
if user_id in user_cache:
user = user_cache[user_id]
else:
async with _db_section():
user = await _fetch_user_readonly(user_dao, user_id)
try:
async with _db_section():
user = await _fetch_user_readonly(user_dao, user_id)
except PermissionExemption:
user = None
user_cache[user_id] = user
else:
async with _db_section():
user = await _fetch_user_readonly(user_dao, user_id)
try:
async with _db_section():
user = await _fetch_user_readonly(user_dao, user_id)
except PermissionExemption:
user = None
return plugin, user
@@ -1089,7 +1290,7 @@ async def get_plugin_cost(
plugin: PluginInfo,
session: Uninfo,
*,
is_superuser: bool = False,
context: PermissionContext | None = None,
) -> int:
"""获取插件费用
@@ -1106,7 +1307,10 @@ async def get_plugin_cost(
返回:
int: 调用插件金币费用
"""
cost_gold = await with_timeout(auth_cost(user, plugin, session), name="auth_cost")
cost_gold = await with_timeout(
auth_cost(user, plugin, session, context=context), name="auth_cost"
)
is_superuser = context.is_superuser if context is not None else False
if is_superuser:
if plugin.plugin_type == PluginType.SUPERUSER:
raise IsSuperuserException()
@@ -1124,7 +1328,7 @@ async def reduce_gold(user_id: str, module: str, cost_gold: int, session: Uninfo
cost_gold: 消耗金币
session: Uninfo
"""
user_dao = DataAccess(UserConsole)
should_clear_cache = False
try:
await with_timeout(
UserConsole.reduce_gold(
@@ -1141,14 +1345,16 @@ async def reduce_gold(user_id: str, module: str, cost_gold: int, session: Uninfo
u.gold = 0
await u.save(update_fields=["gold"])
except asyncio.TimeoutError:
should_clear_cache = True
logger.error(
f"扣除金币超时,用户: {user_id}, 金币: {cost_gold}",
LOGGER_COMMAND,
session=session,
)
# 清除缓存,使下次查询时从数据库获取最新数据
await user_dao.clear_cache(user_id=user_id)
# 正常写入路径由 UserConsole.save() 统一失效缓存;超时状态不确定时兜底清理。
if should_clear_cache:
await DataAccess(UserConsole).clear_cache(user_id=user_id)
logger.debug(f"调用功能花费金币: {cost_gold}", LOGGER_COMMAND, session=session)
@@ -1174,16 +1380,15 @@ async def time_hook(coro, name, recorder: HookTraceRecorder | None = None):
async def _enter_hooks_section():
"""尝试获取全局信号量并更新计数器,超时则抛出 PermissionExemption。"""
"""尝试获取全局信号量并更新计数器,饱和时快速放行。"""
global HOOKS_ACTIVE_COUNT
try:
await asyncio.wait_for(HOOKS_SEMAPHORE.acquire(), timeout=TIMEOUT_SECONDS)
except asyncio.TimeoutError:
if HOOKS_SEMAPHORE.locked():
logger.warning(
"hooks semaphore acquire timeout, allowing pass",
"hooks semaphore saturated, allowing pass",
LOGGER_COMMAND,
)
raise PermissionExemption("hooks semaphore timeout, allow pass")
raise PermissionExemption("hooks semaphore saturated, allow pass")
await HOOKS_SEMAPHORE.acquire()
HOOKS_ACTIVE_COUNT += 1
@@ -1197,14 +1402,7 @@ async def _leave_hooks_section():
async def route_precheck(
matcher: Matcher,
event: Event,
session: Uninfo,
message: UniMsg | None,
*,
entity=None,
event_cache: dict | None = None,
text: str | None = None,
route_modules: set[str] | None = None,
context: EventContext,
) -> bool:
module = matcher.plugin_name or ""
if not module:
@@ -1213,19 +1411,20 @@ async def route_precheck(
return False
if not _is_command_matcher_class(type(matcher)):
return False
if entity is None:
entity = get_entity_ids(session)
if event_cache is None:
event_cache = _get_event_cache(event, session, entity)
if text is None:
text = _get_message_text(message, event_cache, event)
route_modules = context.route_modules if context.route_modules_loaded else None
if route_modules is None:
route_modules = await _get_route_context(text, event_cache)
route_modules = await _get_route_context(
context.plain_text,
context.event_cache,
)
set_route_modules(None, context, route_modules)
if module in _ROUTE_MODULES_WITH_COMMANDS and module not in route_modules:
if _matcher_has_alconna_shortcuts(type(matcher)):
return False
if event_cache is not None:
event_cache["route_skip"] = True
if context.event_cache is not None:
context.event_cache["route_skip"] = True
return True
return False
@@ -1235,14 +1434,10 @@ async def auth(
event: Event,
bot: Bot,
session: Uninfo,
message: UniMsg | None,
*,
context: EventContext,
skip_ban: bool = False,
entity=None,
event_cache: dict | None = None,
text: str | None = None,
route_modules: set[str] | None = None,
is_superuser: bool = False,
state: dict | None = None,
):
"""权限检查
@@ -1251,20 +1446,28 @@ async def auth(
event: Event
bot: bot
session: Uninfo
message: UniMsg
context: EventContext
"""
start_time = time.time()
cost_gold = 0
ignore_flag = False
if entity is None:
entity = get_entity_ids(session)
entity = context.entity
event_cache = context.event_cache
text = context.plain_text
is_superuser = context.is_superuser
route_modules = context.route_modules if context.route_modules_loaded else None
module = matcher.plugin_name or ""
is_command_matcher = _is_command_matcher_class(type(matcher))
if event_cache is None:
event_cache = _get_event_cache(event, session, entity)
auth_allowed = None
auth_result_cache = None
admin_checked_pre = False
permission_context: PermissionContext | None = None
side_effect_cache = get_permission_side_effect_cache(
state=state,
event_cache=event_cache,
)
side_effect_lock = None
entered_side_effect_lock = False
# 仅在慢请求时记录 hook 明细,避免热路径高频构造字符串
hook_recorder = HookTraceRecorder(start_time)
@@ -1278,14 +1481,17 @@ async def auth(
auth_allowed = True
return
if event_cache is not None:
auth_result_cache = event_cache.setdefault("auth_result", {})
cached_result = auth_result_cache.get(module)
if cached_result is not None:
allowed, reason = cached_result
if not allowed:
raise SkipPluginException(reason or "auth cached skip")
return
side_effect_lock = side_effect_cache.lock_for(module)
await side_effect_lock.acquire()
entered_side_effect_lock = True
auth_result_cache = side_effect_cache.auth_results
cached_result = auth_result_cache.get(module)
if cached_result is not None:
allowed, reason = cached_result
if not allowed:
raise SkipPluginException(reason or "auth cached skip")
return
if _is_hidden_plugin(matcher):
auth_allowed = True
@@ -1293,10 +1499,9 @@ async def auth(
if event_cache is not None and event_cache.get("ban_state") is True:
raise SkipPluginException("user or group banned (cached)")
if text is None:
text = _get_message_text(message, event_cache, event)
if route_modules is None:
route_modules = await _get_route_context(text, event_cache)
set_route_modules(state, context, route_modules)
route_skip_checks = (
is_command_matcher
and module in _ROUTE_MODULES_WITH_COMMANDS
@@ -1310,7 +1515,7 @@ async def auth(
auth_allowed = True
return
platform = PlatformUtils.get_platform(session)
platform = context.platform
# 获取插件和用户数据
plugin_user_start = time.time()
try:
@@ -1336,6 +1541,14 @@ async def auth(
auth_allowed = True
return
permission_context = PermissionContext(
event=context,
module=module,
plugin=plugin,
user=user,
)
store_permission_context(state, permission_context)
if not route_skip_checks and _needs_admin_check(plugin):
if plugin.plugin_type in {
PluginType.SUPERUSER,
@@ -1356,13 +1569,18 @@ async def auth(
admin_timeout = False
if event_cache is not None:
admin_levels, admin_timeout = await _get_admin_levels_cached(
session, entity, event_cache
entity, event_cache
)
permission_context.admin_levels = admin_levels
if admin_timeout:
hook_recorder.set("auth_admin", "timeout")
else:
admin_start = time.time()
await auth_admin(plugin, session, cached_levels=admin_levels)
await auth_admin(
plugin,
session,
context=permission_context,
)
hook_recorder.set(
"auth_admin", f"{time.time() - admin_start:.3f}s(pre)"
)
@@ -1386,8 +1604,7 @@ async def auth(
matcher,
session,
plugin,
entity=entity,
is_superuser=is_superuser,
context=permission_context,
)
hook_recorder.set("auth_ban", f"{time.time() - ban_start:.3f}s")
if event_cache is not None:
@@ -1407,7 +1624,7 @@ async def auth(
user,
plugin,
session,
is_superuser=is_superuser,
context=permission_context,
),
name="get_plugin_cost",
)
@@ -1421,7 +1638,7 @@ async def auth(
hook_recorder.set("cost_gold", "skipped")
# 执行 bot_filter
bot_filter(session)
bot_filter(session, context=permission_context)
group = await _get_group_cached(entity, event_cache)
@@ -1439,13 +1656,23 @@ async def auth(
and not route_skip_checks
):
admin_levels, admin_timeout = await _get_admin_levels_cached(
session, entity, event_cache
entity, event_cache
)
permission_context.group = group
permission_context.bot_data = bot_data
if admin_levels is not None:
permission_context.admin_levels = admin_levels
store_permission_context(state, permission_context)
# 并行执行所有 hook 检查,并记录执行时间
hooks_start = time.time()
allow_sleep_bypass = _is_bot_wake_command(module, text)
# 先进入 hooks 并行检查区域;饱和时快速放行,避免创建并积压协程。
await _enter_hooks_section()
entered_hooks = True
# 创建所有 hook 任务
hook_tasks = []
if event_cache is None:
@@ -1455,6 +1682,7 @@ async def auth(
plugin,
bot.self_id,
allow_sleep_bypass=allow_sleep_bypass,
context=permission_context,
),
"auth_bot",
hook_recorder,
@@ -1472,6 +1700,7 @@ async def auth(
bot_data=bot_data,
skip_fetch=True,
allow_sleep_bypass=allow_sleep_bypass,
context=permission_context,
),
"auth_bot",
hook_recorder,
@@ -1483,7 +1712,13 @@ async def auth(
else:
hook_tasks.append(
time_hook(
auth_group(plugin, group, text, entity.group_id),
auth_group(
plugin,
group,
text,
entity.group_id,
context=permission_context,
),
"auth_group",
hook_recorder,
)
@@ -1492,7 +1727,11 @@ async def auth(
if not route_skip_checks and plugin.admin_level and not admin_checked_pre:
if event_cache is None:
hook_tasks.append(
time_hook(auth_admin(plugin, session), "auth_admin", hook_recorder)
time_hook(
auth_admin(plugin, session, context=permission_context),
"auth_admin",
hook_recorder,
)
)
else:
if admin_timeout:
@@ -1500,7 +1739,11 @@ async def auth(
else:
hook_tasks.append(
time_hook(
auth_admin(plugin, session, cached_levels=admin_levels),
auth_admin(
plugin,
session,
context=permission_context,
),
"auth_admin",
hook_recorder,
)
@@ -1510,7 +1753,7 @@ async def auth(
if is_superuser:
hook_recorder.set("auth_plugin", "superuser")
elif not route_skip_checks and _needs_auth_plugin(plugin, group, entity):
elif not route_skip_checks and _needs_auth_plugin(plugin, permission_context):
hook_tasks.append(
time_hook(
auth_plugin(
@@ -1518,6 +1761,7 @@ async def auth(
group,
session,
event,
context=permission_context,
skip_group_block=is_superuser,
),
"auth_plugin",
@@ -1531,18 +1775,17 @@ async def auth(
has_limits = await _has_limits_cached(module, event_cache)
if has_limits:
hook_tasks.append(
time_hook(auth_limit(plugin, session), "auth_limit", hook_recorder)
time_hook(
auth_limit(plugin, session, context=permission_context),
"auth_limit",
hook_recorder,
)
)
else:
hook_recorder.set("auth_limit", "skipped")
else:
hook_recorder.set("auth_limit", "skipped")
if hook_tasks:
# 进入 hooks 并行检查区域(会在高并发时排队)
await _enter_hooks_section()
entered_hooks = True
# 使用 gather 并行执行所有 hook,但添加总体超时控制
try:
await with_timeout(
@@ -1599,6 +1842,9 @@ async def auth(
)
if auth_result_cache is not None and auth_allowed is not None:
auth_result_cache[module] = (auth_allowed, None)
if entered_side_effect_lock and side_effect_lock is not None:
with contextlib.suppress(Exception):
side_effect_lock.release()
# 扣除金币
if not ignore_flag and cost_gold > 0:
gold_start = time.time()
+47 -98
View File
@@ -1,4 +1,3 @@
import contextlib
import time
from nonebot import get_driver
@@ -12,14 +11,20 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.services.cache.runtime_cache import is_cache_ready
from zhenxun.services.log import logger
from zhenxun.services.message_load import is_overloaded
from zhenxun.services.message_load import is_overloaded, mark_activity
from zhenxun.services.runtime_bootstrap import register_runtime_bootstrap
from zhenxun.utils.utils import get_entity_ids
from .auth.config import LOGGER_COMMAND
from .auth.context import (
get_event_context,
get_or_create_event_context,
resolve_actor_user_id,
resolve_event_channel_id,
resolve_event_group_id,
set_route_modules,
)
from .auth_checker import (
LimitManager,
_get_event_cache,
_get_route_context,
auth,
route_precheck,
@@ -41,17 +46,6 @@ async def _mark_bot_connected(bot: Bot):
_BOT_CONNECT_TS = time.time()
def _extract_plain_text(message: UniMsg | None, event: Event) -> str:
if message is not None:
with contextlib.suppress(Exception):
return message.extract_plain_text()
with contextlib.suppress(Exception):
plain = event.get_plaintext()
if plain:
return plain.strip()
return ""
@driver.on_startup
async def _start_auth_runtime_tasks():
await start_auth_runtime_tasks()
@@ -72,37 +66,9 @@ def _skip_auth_for_plugin(matcher: Matcher) -> bool:
return "chat_history" in module_name
def _resolve_actor_user_id(event: Event, fallback_user_id: str) -> str:
"""优先使用事件发起者ID,避免 notice 场景 session.user 指向 bot 自身。"""
event_user_id = getattr(event, "user_id", None)
if event_user_id is None:
return fallback_user_id
event_user_id = str(event_user_id)
return event_user_id or fallback_user_id
def _resolve_event_group_id(event: Event, fallback_group_id: str | None) -> str | None:
"""notice 场景 session.group 可能缺失,回退到事件上的 group_id。"""
event_group_id = getattr(event, "group_id", None)
if event_group_id is None:
return fallback_group_id
resolved = str(event_group_id)
return resolved or fallback_group_id
def _resolve_event_channel_id(
event: Event, fallback_channel_id: str | None
) -> str | None:
"""频道场景回退到事件上的 channel_id。"""
event_channel_id = getattr(event, "channel_id", None)
if event_channel_id is None:
return fallback_channel_id
resolved = str(event_channel_id)
return resolved or fallback_channel_id
@event_preprocessor
async def _drop_message_before_cache_ready(event: Event):
mark_activity()
if event.get_type() != "message":
return
if not is_cache_ready():
@@ -130,46 +96,22 @@ async def _auth_preprocessor(
return
start_time = time.time()
entity = state.get("_zx_entity")
if entity is None:
entity = get_entity_ids(session)
entity.user_id = _resolve_actor_user_id(event, entity.user_id)
entity.group_id = _resolve_event_group_id(event, entity.group_id)
entity.channel_id = _resolve_event_channel_id(event, entity.channel_id)
state["_zx_entity"] = entity
event_cache = state.get("_zx_event_cache")
if event_cache is None:
event_cache = _get_event_cache(event, session, entity)
state["_zx_event_cache"] = event_cache
text = state.get("_zx_plain_text")
if text is None:
text = _extract_plain_text(message, event)
state["_zx_plain_text"] = text
if event_cache is not None:
event_cache["plain_text"] = text
route_modules = state.get("_zx_route_modules")
if route_modules is None:
route_modules = await _get_route_context(text, event_cache)
state["_zx_route_modules"] = route_modules
is_superuser = state.get("_zx_is_superuser")
if is_superuser is None:
is_superuser = entity.user_id in bot.config.superusers
state["_zx_is_superuser"] = is_superuser
if await route_precheck(
matcher,
event_context = get_or_create_event_context(
bot,
event,
session,
message,
entity=entity,
event_cache=event_cache,
text=text,
route_modules=route_modules,
):
state,
message=message,
)
if not event_context.route_modules_loaded:
route_modules = await _get_route_context(
event_context.plain_text,
event_context.event_cache,
)
set_route_modules(state, event_context, route_modules)
if await route_precheck(matcher, event_context):
return
try:
@@ -178,13 +120,9 @@ async def _auth_preprocessor(
event,
bot,
session,
message,
context=event_context,
skip_ban=False,
entity=entity,
event_cache=event_cache,
text=text,
route_modules=route_modules,
is_superuser=is_superuser,
state=state,
)
except IgnoredException:
raise
@@ -203,16 +141,27 @@ async def _auth_preprocessor(
@run_postprocessor
async def _unblock_after_matcher(matcher: Matcher, session: Uninfo, event: Event):
user_id = _resolve_actor_user_id(event, session.user.id)
group_id = _resolve_event_group_id(event, None)
channel_id = _resolve_event_channel_id(event, None)
if session.group:
if session.group.parent:
group_id = session.group.parent.id
channel_id = session.group.id
else:
group_id = session.group.id
async def _unblock_after_matcher(
matcher: Matcher,
session: Uninfo,
event: Event,
state: T_State,
):
context = get_event_context(state)
if context is not None:
user_id = context.user_id
group_id = context.group_id
channel_id = context.channel_id
else:
user_id = resolve_actor_user_id(event, session.user.id)
group_id = resolve_event_group_id(event, None)
channel_id = resolve_event_channel_id(event, None)
if session.group:
if session.group.parent:
group_id = session.group.parent.id
channel_id = session.group.id
else:
group_id = session.group.id
if user_id and matcher.plugin:
module = matcher.plugin.name
LimitManager.unblock(module, user_id, group_id, channel_id)
+24 -14
View File
@@ -16,13 +16,7 @@ from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType
from zhenxun.utils.message import MessageUtils
malicious_check_time = Config.get_config("hook", "MALICIOUS_CHECK_TIME")
malicious_ban_count = Config.get_config("hook", "MALICIOUS_BAN_COUNT")
if not malicious_check_time:
raise ValueError("模块: [hook], 配置项: [MALICIOUS_CHECK_TIME] 为空或小于0")
if not malicious_ban_count:
raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_COUNT] 为空或小于0")
from .auth.context import resolve_actor_user_id, resolve_event_group_id
class BanCheckLimiter:
@@ -36,6 +30,10 @@ class BanCheckLimiter:
self.default_check_time = default_check_time
self.default_count = default_count
def configure(self, check_time: float, count: int) -> None:
self.default_check_time = check_time
self.default_count = count
def add(self, key: str | float):
if self.mint[key] == 1:
self.mtime[key] = time.time()
@@ -59,11 +57,22 @@ class BanCheckLimiter:
_blmt = BanCheckLimiter(
malicious_check_time,
malicious_ban_count,
5,
4,
)
def _get_positive_config(key: str, cast_type: type[int] | type[float]) -> int | float:
value = Config.get_config("hook", key)
try:
parsed_value = cast_type(value)
except (TypeError, ValueError) as e:
raise ValueError(f"模块: [hook], 配置项: [{key}] 不是有效数字") from e
if parsed_value <= 0:
raise ValueError(f"模块: [hook], 配置项: [{key}] 为空或小于0")
return parsed_value
# 恶意触发命令检测
@run_preprocessor
async def _(
@@ -88,11 +97,12 @@ async def _(
else:
return
user_id = session.id1
group_id = session.id3 or session.id2
malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME")
if not malicious_ban_time:
raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_TIME] 为空或小于0")
user_id = resolve_actor_user_id(event, session.id1)
group_id = resolve_event_group_id(event, session.id3 or session.id2)
malicious_check_time = float(_get_positive_config("MALICIOUS_CHECK_TIME", float))
malicious_ban_count = int(_get_positive_config("MALICIOUS_BAN_COUNT", int))
malicious_ban_time = int(_get_positive_config("MALICIOUS_BAN_TIME", int))
_blmt.configure(malicious_check_time, malicious_ban_count)
if user_id and module:
if _blmt.check(f"{user_id}__{module}"):
await BanConsole.ban(