mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-29 00:32:06 +08:00
bugfix:更换数据库初始化超时路径以修复连接超时问题 (#2137)
* bugfix:更换数据库初始化超时路径以修复连接超时问题 * 移除部分观测链路 * 细节修改 * 权限检查细节修改2 * bugfix:修复金币懒加载造成插件金币消耗不了的问题 * bugfix:整理鉴权逻辑 * 完善缓存系统 * 优化官端使用 * bugfix:增加官机和频道的前导统一剥离,新增获取群列表自愈 * bugfix:修复导入问题
This commit is contained in:
@@ -49,14 +49,4 @@ Config.add_plugin_config(
|
||||
type=bool,
|
||||
)
|
||||
|
||||
Config.add_plugin_config(
|
||||
"hook",
|
||||
"RECORD_BOT_SENT_MESSAGES",
|
||||
True,
|
||||
help="记录bot消息发送",
|
||||
default_value=True,
|
||||
type=bool,
|
||||
)
|
||||
|
||||
|
||||
nonebot.load_plugins(str(Path(__file__).parent.resolve()))
|
||||
|
||||
@@ -5,12 +5,12 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.models.level_user import LevelUser
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.cache.runtime_cache import LevelUserMemoryCache, LevelUserSnapshot
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.utils import EntityIDs, get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .data_provider import DEFAULT_PERMISSION_DATA_PROVIDER, LevelUserSnapshot
|
||||
from .exception import SkipPluginException
|
||||
|
||||
|
||||
@@ -50,7 +50,10 @@ async def auth_admin(
|
||||
if cached_levels is not None:
|
||||
global_user, group_users = cached_levels
|
||||
else:
|
||||
global_user, group_users = await LevelUserMemoryCache.get_levels(
|
||||
(
|
||||
global_user,
|
||||
group_users,
|
||||
) = await DEFAULT_PERMISSION_DATA_PROVIDER.get_admin_levels(
|
||||
entity.user_id, entity.group_id
|
||||
)
|
||||
|
||||
|
||||
@@ -7,7 +7,6 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.ban_console import BanConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.cache.runtime_cache import BanMemoryCache
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import PluginType
|
||||
@@ -15,6 +14,7 @@ from zhenxun.utils.utils import EntityIDs, get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .data_provider import DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
from .exception import SkipPluginException
|
||||
from .utils import freq
|
||||
|
||||
@@ -60,9 +60,10 @@ async def is_ban(user_id: str | None, group_id: str | None) -> int:
|
||||
"""
|
||||
if not user_id and not group_id:
|
||||
return 0
|
||||
if not BanMemoryCache.is_loaded():
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
if not provider.ban_cache_loaded():
|
||||
return 0
|
||||
return BanMemoryCache.remaining_time(user_id, group_id)
|
||||
return provider.get_ban_remaining_time(user_id, group_id)
|
||||
|
||||
|
||||
def check_plugin_type(matcher: Matcher) -> bool:
|
||||
|
||||
@@ -2,12 +2,12 @@ import time
|
||||
|
||||
from zhenxun.models.bot_console import BotConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.cache.runtime_cache import BotMemoryCache, BotSnapshot
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .data_provider import DEFAULT_PERMISSION_DATA_PROVIDER, BotSnapshot
|
||||
from .exception import SkipPluginException
|
||||
|
||||
|
||||
@@ -33,12 +33,13 @@ async def auth_bot(
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
if context is not None:
|
||||
bot_id = context.event.bot_id
|
||||
bot_data = context.bot_data
|
||||
bot: BotConsole | BotSnapshot | None = bot_data
|
||||
if bot is None and not skip_fetch:
|
||||
bot = await BotMemoryCache.get(bot_id)
|
||||
bot = await provider.get_bot(bot_id)
|
||||
|
||||
if bot is None:
|
||||
raise SkipPluginException("Bot不存在,阻断权限检测...")
|
||||
|
||||
@@ -10,10 +10,6 @@ from pydantic import BaseModel
|
||||
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.plugin_limit import PluginLimit
|
||||
from zhenxun.services.cache.runtime_cache import (
|
||||
PluginLimitMemoryCache,
|
||||
PluginLimitSnapshot,
|
||||
)
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import LimitWatchType, PluginLimitType
|
||||
@@ -25,6 +21,10 @@ from zhenxun.utils.utils import EntityIDs, get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .data_provider import (
|
||||
DEFAULT_PERMISSION_DATA_PROVIDER,
|
||||
PluginLimitSnapshot,
|
||||
)
|
||||
from .exception import SkipPluginException
|
||||
|
||||
driver = nonebot.get_driver()
|
||||
@@ -106,11 +106,10 @@ class LimitManager:
|
||||
block_limit: ClassVar[dict[str, Limit]] = {}
|
||||
count_limit: ClassVar[dict[str, Limit]] = {}
|
||||
|
||||
# 模块限制缓存,避免频繁查询数据库
|
||||
module_limit_cache: ClassVar[
|
||||
dict[str, tuple[float, list[PluginLimitSnapshot], bool]]
|
||||
# 只缓存异常短路结果;正常 limit 列表统一从 PluginLimitMemoryCache 读取。
|
||||
module_limit_error_cache: ClassVar[
|
||||
dict[str, tuple[float, list[PluginLimitSnapshot]]]
|
||||
] = {}
|
||||
module_cache_ttl: ClassVar[float] = 60 # 模块缓存有效期(秒)
|
||||
module_cache_error_ttl: ClassVar[float] = 5 # 超时缓存有效期(秒)
|
||||
|
||||
@classmethod
|
||||
@@ -132,15 +131,16 @@ class LimitManager:
|
||||
cls.is_updating = True
|
||||
try:
|
||||
start_time = time.time()
|
||||
await PluginLimitMemoryCache.ensure_loaded()
|
||||
limit_list = PluginLimitMemoryCache.get_all_limits()
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
await provider.ensure_module_limits_loaded()
|
||||
limit_list = await provider.get_all_module_limits()
|
||||
|
||||
# 清空旧数据
|
||||
cls.add_module = []
|
||||
cls.cd_limit = {}
|
||||
cls.block_limit = {}
|
||||
cls.count_limit = {}
|
||||
cls.module_limit_cache.clear()
|
||||
cls.module_limit_error_cache.clear()
|
||||
# 添加新数据
|
||||
for limit in limit_list:
|
||||
cls.add_limit(limit)
|
||||
@@ -216,22 +216,21 @@ class LimitManager:
|
||||
"""
|
||||
current_time = time.time()
|
||||
|
||||
# 检查缓存
|
||||
if module in cls.module_limit_cache:
|
||||
cache_time, limits, is_error = cls.module_limit_cache[module]
|
||||
ttl = cls.module_cache_error_ttl if is_error else cls.module_cache_ttl
|
||||
if current_time - cache_time < ttl:
|
||||
# 正常路径不再二次缓存列表,避免与 PluginLimitMemoryCache 形成双真源。
|
||||
if module in cls.module_limit_error_cache:
|
||||
cache_time, limits = cls.module_limit_error_cache[module]
|
||||
if current_time - cache_time < cls.module_cache_error_ttl:
|
||||
return limits
|
||||
cls.module_limit_error_cache.pop(module, None)
|
||||
|
||||
# 缓存不存在或已过期,从内存缓存获取
|
||||
try:
|
||||
await PluginLimitMemoryCache.ensure_loaded()
|
||||
limits = await PluginLimitMemoryCache.get_limits(module)
|
||||
cls.module_limit_cache[module] = (current_time, limits, False)
|
||||
return limits
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
await provider.ensure_module_limits_loaded()
|
||||
return await provider.get_module_limits(module)
|
||||
except Exception as exc:
|
||||
logger.error(f"get module limits failed: {module}", LOGGER_COMMAND, e=exc)
|
||||
cls.module_limit_cache[module] = (current_time, [], True)
|
||||
cls.module_limit_error_cache[module] = (current_time, [])
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -39,6 +39,7 @@ if TYPE_CHECKING:
|
||||
class EventContext:
|
||||
bot_id: str
|
||||
platform: str
|
||||
platform_scope: str
|
||||
event_type: str
|
||||
message_id: str | int | None
|
||||
entity: EntityIDs
|
||||
@@ -170,6 +171,7 @@ def event_cache_key(
|
||||
*,
|
||||
bot_id: str,
|
||||
platform: str,
|
||||
platform_scope: str | None = None,
|
||||
entity: EntityIDs,
|
||||
) -> str:
|
||||
msg_id = _event_message_id(event)
|
||||
@@ -177,7 +179,11 @@ def event_cache_key(
|
||||
msg_id = id(event)
|
||||
group_id = entity.group_id or ""
|
||||
channel_id = entity.channel_id or ""
|
||||
return f"{platform}:{bot_id}:{entity.user_id}:{group_id}:{channel_id}:{msg_id}"
|
||||
scope = platform_scope or platform
|
||||
return (
|
||||
f"{scope}:{platform}:{bot_id}:{entity.user_id}:"
|
||||
f"{group_id}:{channel_id}:{msg_id}"
|
||||
)
|
||||
|
||||
|
||||
def get_event_cache(
|
||||
@@ -185,11 +191,18 @@ def get_event_cache(
|
||||
*,
|
||||
bot_id: str,
|
||||
platform: str,
|
||||
platform_scope: str | None = None,
|
||||
entity: EntityIDs,
|
||||
) -> dict[str, Any] | None:
|
||||
if not EVENT_CACHE:
|
||||
return None
|
||||
key = event_cache_key(event, bot_id=bot_id, platform=platform, entity=entity)
|
||||
key = event_cache_key(
|
||||
event,
|
||||
bot_id=bot_id,
|
||||
platform=platform,
|
||||
platform_scope=platform_scope,
|
||||
entity=entity,
|
||||
)
|
||||
try:
|
||||
return EVENT_CACHE[key]
|
||||
except KeyError:
|
||||
@@ -255,6 +268,7 @@ def get_or_create_event_context(
|
||||
entity = resolve_entity_ids(event, session)
|
||||
|
||||
platform = PlatformUtils.get_platform(session)
|
||||
platform_scope = PlatformUtils.get_platform_scope(session)
|
||||
bot_id = str(bot.self_id)
|
||||
event_cache = state.get(STATE_EVENT_CACHE)
|
||||
if not isinstance(event_cache, dict):
|
||||
@@ -262,6 +276,7 @@ def get_or_create_event_context(
|
||||
event,
|
||||
bot_id=bot_id,
|
||||
platform=platform,
|
||||
platform_scope=platform_scope,
|
||||
entity=entity,
|
||||
)
|
||||
|
||||
@@ -292,6 +307,7 @@ def get_or_create_event_context(
|
||||
context = EventContext(
|
||||
bot_id=bot_id,
|
||||
platform=platform,
|
||||
platform_scope=platform_scope,
|
||||
event_type=event.get_type(),
|
||||
message_id=_event_message_id(event),
|
||||
entity=entity,
|
||||
|
||||
@@ -0,0 +1,151 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from zhenxun.services.cache.runtime_cache import (
|
||||
BanMemoryCache,
|
||||
BotMemoryCache,
|
||||
BotSnapshot,
|
||||
GroupMemoryCache,
|
||||
GroupSnapshot,
|
||||
LevelUserMemoryCache,
|
||||
LevelUserSnapshot,
|
||||
PluginLimitMemoryCache,
|
||||
PluginLimitSnapshot,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
|
||||
|
||||
AdminLevels = tuple[LevelUserSnapshot | None, LevelUserSnapshot | None]
|
||||
|
||||
|
||||
class PermissionDataProvider:
|
||||
"""Auth data facade over runtime caches.
|
||||
|
||||
Permission checks should read stable runtime snapshots through this provider
|
||||
instead of reaching into individual cache classes from multiple auth modules.
|
||||
The provider does not own policy semantics and does not query the database
|
||||
directly.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def plugin_cache_loaded() -> bool:
|
||||
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
|
||||
|
||||
return PluginInfoMemoryCache.is_loaded()
|
||||
|
||||
@staticmethod
|
||||
def get_plugin_if_ready(module: str) -> "PluginInfo | None":
|
||||
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
|
||||
|
||||
return PluginInfoMemoryCache.get_by_module_if_ready(module)
|
||||
|
||||
@staticmethod
|
||||
async def get_plugin(module: str) -> "PluginInfo | None":
|
||||
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
|
||||
|
||||
return await PluginInfoMemoryCache.get_by_module(module)
|
||||
|
||||
@staticmethod
|
||||
def module_limit_cache_loaded() -> bool:
|
||||
return PluginLimitMemoryCache.is_loaded()
|
||||
|
||||
@staticmethod
|
||||
async def ensure_module_limits_loaded() -> None:
|
||||
await PluginLimitMemoryCache.ensure_loaded()
|
||||
|
||||
@staticmethod
|
||||
def get_module_limits_if_ready(
|
||||
module: str,
|
||||
) -> list[PluginLimitSnapshot] | None:
|
||||
return PluginLimitMemoryCache.get_limits_if_ready(module)
|
||||
|
||||
@staticmethod
|
||||
async def get_module_limits(module: str) -> list[PluginLimitSnapshot]:
|
||||
return await PluginLimitMemoryCache.get_limits(module)
|
||||
|
||||
@staticmethod
|
||||
async def get_all_module_limits() -> list[PluginLimitSnapshot]:
|
||||
if not PluginLimitMemoryCache.is_loaded():
|
||||
await PluginLimitMemoryCache.ensure_loaded()
|
||||
return PluginLimitMemoryCache.get_all_limits()
|
||||
|
||||
@staticmethod
|
||||
def bot_cache_loaded() -> bool:
|
||||
return BotMemoryCache.is_loaded()
|
||||
|
||||
@staticmethod
|
||||
def get_bot_if_ready(bot_id: str | None) -> BotSnapshot | None:
|
||||
return BotMemoryCache.get_if_ready(bot_id)
|
||||
|
||||
@staticmethod
|
||||
async def get_bot(bot_id: str | None) -> BotSnapshot | None:
|
||||
return await BotMemoryCache.get(bot_id)
|
||||
|
||||
@staticmethod
|
||||
def group_cache_loaded() -> bool:
|
||||
return GroupMemoryCache.is_loaded()
|
||||
|
||||
@staticmethod
|
||||
def get_group_if_ready(
|
||||
group_id: str | None,
|
||||
channel_id: str | None = None,
|
||||
) -> GroupSnapshot | None:
|
||||
return GroupMemoryCache.get_if_ready(group_id, channel_id)
|
||||
|
||||
@staticmethod
|
||||
async def get_group(
|
||||
group_id: str | None,
|
||||
channel_id: str | None = None,
|
||||
) -> GroupSnapshot | None:
|
||||
return await GroupMemoryCache.get(group_id, channel_id)
|
||||
|
||||
@staticmethod
|
||||
def admin_cache_loaded() -> bool:
|
||||
return LevelUserMemoryCache.is_loaded()
|
||||
|
||||
@staticmethod
|
||||
def get_admin_levels_if_ready(
|
||||
user_id: str | None,
|
||||
group_id: str | None,
|
||||
) -> AdminLevels | None:
|
||||
return LevelUserMemoryCache.get_levels_if_ready(user_id, group_id)
|
||||
|
||||
@staticmethod
|
||||
async def get_admin_levels(
|
||||
user_id: str | None,
|
||||
group_id: str | None,
|
||||
) -> AdminLevels:
|
||||
return await LevelUserMemoryCache.get_levels(user_id, group_id)
|
||||
|
||||
@staticmethod
|
||||
def ban_cache_loaded() -> bool:
|
||||
return BanMemoryCache.is_loaded()
|
||||
|
||||
@staticmethod
|
||||
async def ensure_ban_loaded() -> None:
|
||||
await BanMemoryCache.ensure_loaded()
|
||||
|
||||
@staticmethod
|
||||
def is_banned(user_id: str | None, group_id: str | None) -> bool:
|
||||
return BanMemoryCache.is_banned(user_id, group_id)
|
||||
|
||||
@staticmethod
|
||||
def get_ban_remaining_time(user_id: str | None, group_id: str | None) -> int:
|
||||
return BanMemoryCache.remaining_time(user_id, group_id)
|
||||
|
||||
|
||||
DEFAULT_PERMISSION_DATA_PROVIDER = PermissionDataProvider()
|
||||
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_PERMISSION_DATA_PROVIDER",
|
||||
"AdminLevels",
|
||||
"BotSnapshot",
|
||||
"GroupSnapshot",
|
||||
"LevelUserSnapshot",
|
||||
"PermissionDataProvider",
|
||||
"PluginLimitSnapshot",
|
||||
]
|
||||
@@ -8,6 +8,7 @@ import re
|
||||
from typing import Any, Literal
|
||||
import weakref
|
||||
|
||||
from loguru import logger
|
||||
from nonebot.matcher import Matcher
|
||||
|
||||
ActivationDecision = Literal["match", "miss", "unknown"]
|
||||
@@ -23,6 +24,20 @@ ActivationLane = Literal[
|
||||
"passive_render",
|
||||
]
|
||||
|
||||
KNOWN_SAFE_RULE_NAMES = frozenset(
|
||||
{
|
||||
"CommandRule",
|
||||
"ShellCommandRule",
|
||||
"RegexRule",
|
||||
"StartswithRule",
|
||||
"EndswithRule",
|
||||
"FullmatchRule",
|
||||
"KeywordsRule",
|
||||
"IsTypeRule",
|
||||
"ToMeRule",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ActivationRuleDescriptor:
|
||||
@@ -48,9 +63,96 @@ class HandlerDescriptor:
|
||||
has_custom_rule: bool = False
|
||||
commands: tuple[str, ...] = ()
|
||||
shortcuts: tuple[str, ...] | None = None
|
||||
alconna: tuple[AlconnaDescriptor, ...] = ()
|
||||
rules: tuple[ActivationRuleDescriptor, ...] = ()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AlconnaShortcutDescriptor:
|
||||
pattern: str
|
||||
fuzzy: bool = False
|
||||
flags: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AlconnaDescriptor:
|
||||
command: str = ""
|
||||
aliases: tuple[str, ...] = ()
|
||||
prefixes: tuple[str, ...] = ()
|
||||
shortcuts: tuple[AlconnaShortcutDescriptor, ...] = ()
|
||||
compact: bool = False
|
||||
skip_for_unmatch: bool = True
|
||||
before_rule_count: int = 0
|
||||
after_rule_count: int = 0
|
||||
before_rule_known_safe: bool = True
|
||||
after_rule_known_safe: bool = True
|
||||
input_rewrite_extensions: tuple[str, ...] = ()
|
||||
|
||||
@property
|
||||
def has_reply_merge_extension(self) -> bool:
|
||||
return "ReplyMergeExtension" in self.input_rewrite_extensions
|
||||
|
||||
@property
|
||||
def regex_command(self) -> bool:
|
||||
return self.command.startswith("re:")
|
||||
|
||||
|
||||
class AlconnaActivationIndex:
|
||||
"""Safe, metadata-only Alconna prefilter.
|
||||
|
||||
This index never executes Alconna.parse() or Matcher.check_rule(). It only
|
||||
skips matchers when the command head/shortcut is known to miss. Anything
|
||||
custom or ambiguous stays fail-open to preserve NoneBot compatibility.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._safe = 0
|
||||
self._unknown = 0
|
||||
|
||||
@property
|
||||
def safe_count(self) -> int:
|
||||
return self._safe
|
||||
|
||||
@property
|
||||
def unknown_count(self) -> int:
|
||||
return self._unknown
|
||||
|
||||
def rebuild(self, descriptors: Iterable[HandlerDescriptor]) -> None:
|
||||
safe = 0
|
||||
unknown = 0
|
||||
for descriptor in descriptors:
|
||||
if descriptor.alconna:
|
||||
if any(_alconna_can_prefilter(item) for item in descriptor.alconna):
|
||||
safe += 1
|
||||
else:
|
||||
unknown += 1
|
||||
self._safe = safe
|
||||
self._unknown = unknown
|
||||
if logger.level("DEBUG"):
|
||||
logger.debug(
|
||||
"alconna activation index rebuilt: safe={}, unknown={}",
|
||||
safe,
|
||||
unknown,
|
||||
)
|
||||
|
||||
def select(
|
||||
self,
|
||||
descriptor: HandlerDescriptor,
|
||||
context: ActivationContext,
|
||||
texts: tuple[str, ...],
|
||||
) -> ActivationDecision:
|
||||
if not descriptor.alconna:
|
||||
return "unknown"
|
||||
saw_unknown = False
|
||||
for alconna in descriptor.alconna:
|
||||
decision = matcher_alconna_head_matches(alconna, texts, context)
|
||||
if decision == "match":
|
||||
return "match"
|
||||
if decision == "unknown":
|
||||
saw_unknown = True
|
||||
return "unknown" if saw_unknown else "miss"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ActivationContext:
|
||||
event_type: str
|
||||
@@ -69,37 +171,25 @@ class ActivationContext:
|
||||
@dataclass(slots=True)
|
||||
class ActivationResult:
|
||||
selected: list[type[Matcher]]
|
||||
fallback_required: bool = False
|
||||
selected_by_lane: dict[str, int] = field(default_factory=dict)
|
||||
skipped_by_lane: dict[str, int] = field(default_factory=dict)
|
||||
selected_by_reason: dict[str, int] = field(default_factory=dict)
|
||||
skipped_by_reason: dict[str, int] = field(default_factory=dict)
|
||||
deterministic_selected: set[type[Matcher]] = field(default_factory=set)
|
||||
total_descriptors: int = 0
|
||||
candidate_count: int = 0
|
||||
|
||||
def mark_selected(self, lane: str, reason: str = "selected") -> None:
|
||||
self.selected_by_lane[lane] = self.selected_by_lane.get(lane, 0) + 1
|
||||
self.selected_by_reason[reason] = self.selected_by_reason.get(reason, 0) + 1
|
||||
|
||||
def mark_skipped(self, lane: str, reason: str = "skipped") -> None:
|
||||
self.skipped_by_lane[lane] = self.skipped_by_lane.get(lane, 0) + 1
|
||||
self.skipped_by_reason[reason] = self.skipped_by_reason.get(reason, 0) + 1
|
||||
|
||||
|
||||
class HandlerActivationIndex:
|
||||
"""In-memory matcher activation index.
|
||||
|
||||
The index is intentionally fail-open: only proven misses are rejected before
|
||||
matcher task creation. Unknown custom rules, incomplete command metadata,
|
||||
and Alconna shortcut misses stay selected so plugin compatibility wins over
|
||||
dispatch aggressiveness.
|
||||
The index is intentionally fail-open: only known NoneBot rule misses are
|
||||
rejected before matcher task creation. Custom rules, incomplete command
|
||||
metadata, and Alconna shortcut misses stay selected so plugin compatibility
|
||||
wins over dispatch aggressiveness.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._by_priority: dict[int, list[HandlerDescriptor]] = {}
|
||||
self._matcher_map: dict[type[Matcher], HandlerDescriptor] = {}
|
||||
self._source_keys: set[tuple[int, tuple[int, ...]]] = set()
|
||||
self._alconna_index = AlconnaActivationIndex()
|
||||
self._compiled = False
|
||||
|
||||
@property
|
||||
@@ -122,6 +212,7 @@ class HandlerActivationIndex:
|
||||
(int(priority), tuple(id(matcher) for matcher in priority_matchers))
|
||||
)
|
||||
self._source_keys = source_keys
|
||||
self._alconna_index.rebuild(self._matcher_map.values())
|
||||
self._compiled = True
|
||||
|
||||
def ensure_fresh(self, matchers: dict[int, list[type[Matcher]]]) -> None:
|
||||
@@ -164,32 +255,16 @@ class HandlerActivationIndex:
|
||||
]
|
||||
for descriptor in descriptors:
|
||||
decision = self._select_descriptor(descriptor, context)
|
||||
lane = descriptor.lane
|
||||
if decision == "fallback":
|
||||
result.fallback_required = True
|
||||
if not _consume_uncertain_budget(descriptor, budget):
|
||||
result.mark_skipped(lane, "fallback_budget_exhausted")
|
||||
continue
|
||||
result.selected.append(descriptor.matcher)
|
||||
result.mark_selected(lane, "fallback_budgeted")
|
||||
continue
|
||||
if decision == "miss":
|
||||
result.mark_skipped(lane, _miss_reason(descriptor, context))
|
||||
continue
|
||||
if decision == "deterministic":
|
||||
result.selected.append(descriptor.matcher)
|
||||
result.deterministic_selected.add(descriptor.matcher)
|
||||
result.mark_selected(lane, "deterministic")
|
||||
continue
|
||||
if not _selection_is_guaranteed(descriptor, context):
|
||||
if not _consume_uncertain_budget(descriptor, budget):
|
||||
result.mark_skipped(lane, "unknown_budget_exhausted")
|
||||
if _is_throttleable_broad_passive(descriptor, context):
|
||||
if not _consume_broad_passive_budget(descriptor, budget):
|
||||
continue
|
||||
selected_reason = "unknown_budgeted"
|
||||
else:
|
||||
selected_reason = "guaranteed"
|
||||
result.selected.append(descriptor.matcher)
|
||||
result.mark_selected(lane, selected_reason)
|
||||
result.candidate_count = len(result.selected)
|
||||
return result
|
||||
|
||||
@@ -197,7 +272,7 @@ class HandlerActivationIndex:
|
||||
self,
|
||||
descriptor: HandlerDescriptor,
|
||||
context: ActivationContext,
|
||||
) -> Literal["select", "miss", "fallback", "deterministic"]:
|
||||
) -> Literal["select", "miss", "deterministic"]:
|
||||
if descriptor.temp:
|
||||
return "select"
|
||||
matcher_type = descriptor.matcher_type
|
||||
@@ -224,7 +299,7 @@ class HandlerActivationIndex:
|
||||
self,
|
||||
descriptor: HandlerDescriptor,
|
||||
context: ActivationContext,
|
||||
) -> Literal["select", "miss", "fallback", "deterministic"]:
|
||||
) -> Literal["select", "miss", "deterministic"]:
|
||||
texts = text_match_candidates(
|
||||
context.plain_text,
|
||||
context.raw_text,
|
||||
@@ -232,8 +307,25 @@ class HandlerActivationIndex:
|
||||
)
|
||||
if not texts:
|
||||
return "select"
|
||||
rule_match = matcher_rule_matches_text(
|
||||
descriptor.rules,
|
||||
context.raw_text,
|
||||
context.plain_text,
|
||||
event=context.event,
|
||||
to_me=context.to_me,
|
||||
)
|
||||
if rule_match == "miss":
|
||||
return "miss"
|
||||
command_matched = False
|
||||
if descriptor.commands:
|
||||
if descriptor.alconna:
|
||||
alconna_match = self._alconna_index.select(descriptor, context, texts)
|
||||
if alconna_match == "match":
|
||||
command_matched = True
|
||||
elif alconna_match == "miss":
|
||||
return "miss"
|
||||
else:
|
||||
return "select"
|
||||
elif descriptor.commands:
|
||||
if any(
|
||||
matcher_command_matches(text, command)
|
||||
for text in texts
|
||||
@@ -248,9 +340,7 @@ class HandlerActivationIndex:
|
||||
if shortcut_match == "match":
|
||||
command_matched = True
|
||||
else:
|
||||
# Command extraction for Alconna/custom matchers is incomplete by
|
||||
# design; a miss here is not proof that NoneBot will miss.
|
||||
return "select"
|
||||
return "miss" if not descriptor.has_custom_rule else "select"
|
||||
else:
|
||||
shortcut_match = matcher_alconna_shortcut_matches_any(
|
||||
descriptor.shortcuts,
|
||||
@@ -259,22 +349,19 @@ class HandlerActivationIndex:
|
||||
if shortcut_match == "match":
|
||||
command_matched = True
|
||||
|
||||
rule_match = matcher_rule_matches_text(
|
||||
descriptor.rules,
|
||||
context.raw_text,
|
||||
context.plain_text,
|
||||
event=context.event,
|
||||
to_me=context.to_me,
|
||||
)
|
||||
if rule_match == "match":
|
||||
if (
|
||||
rule_match == "match"
|
||||
and not descriptor.has_custom_rule
|
||||
and descriptor.shortcuts is None
|
||||
and not descriptor.alconna
|
||||
):
|
||||
command_matched = True
|
||||
elif rule_match == "miss":
|
||||
return "miss"
|
||||
elif (
|
||||
if (
|
||||
not (descriptor.commands or descriptor.shortcuts is not None)
|
||||
and not command_matched
|
||||
and not descriptor.alconna
|
||||
):
|
||||
return "fallback"
|
||||
return "select"
|
||||
|
||||
if (
|
||||
context.ai_route_modules
|
||||
@@ -284,7 +371,7 @@ class HandlerActivationIndex:
|
||||
return "select"
|
||||
|
||||
if command_matched:
|
||||
return "deterministic"
|
||||
return "select" if descriptor.has_custom_rule else "deterministic"
|
||||
return "select"
|
||||
|
||||
def _build_descriptor(
|
||||
@@ -298,7 +385,10 @@ class HandlerActivationIndex:
|
||||
if hasattr(matcher, "command"):
|
||||
command_like = True
|
||||
commands = extract_matcher_command_literals(matcher) or ()
|
||||
alconna_descriptors = extract_matcher_alconna_descriptors(matcher)
|
||||
shortcuts = extract_matcher_alconna_shortcuts(matcher)
|
||||
if alconna_descriptors:
|
||||
command_like = True
|
||||
if shortcuts is not None:
|
||||
command_like = True
|
||||
module = matcher_module_name(matcher)
|
||||
@@ -323,6 +413,7 @@ class HandlerActivationIndex:
|
||||
has_custom_rule=matcher_has_custom_rule(matcher),
|
||||
commands=commands,
|
||||
shortcuts=shortcuts,
|
||||
alconna=alconna_descriptors,
|
||||
rules=rules,
|
||||
)
|
||||
|
||||
@@ -393,6 +484,8 @@ def matcher_is_command_like(matcher_cls: type[Matcher]) -> bool:
|
||||
return True
|
||||
if hasattr(matcher_cls, "command"):
|
||||
return True
|
||||
if extract_matcher_alconna_descriptors(matcher_cls):
|
||||
return True
|
||||
return extract_matcher_alconna_shortcuts(matcher_cls) is not None
|
||||
|
||||
|
||||
@@ -420,7 +513,10 @@ def classify_matcher_lane(
|
||||
if hasattr(matcher_cls, "command"):
|
||||
command_like = True
|
||||
commands = extract_matcher_command_literals(matcher_cls) or ()
|
||||
alconna_descriptors = extract_matcher_alconna_descriptors(matcher_cls)
|
||||
shortcuts = extract_matcher_alconna_shortcuts(matcher_cls)
|
||||
if alconna_descriptors:
|
||||
command_like = True
|
||||
if shortcuts is not None:
|
||||
command_like = True
|
||||
return classify_lane(
|
||||
@@ -437,11 +533,6 @@ def extract_matcher_rule_descriptors(
|
||||
matcher_cls: type[Matcher],
|
||||
) -> tuple[ActivationRuleDescriptor, ...]:
|
||||
descriptors: list[ActivationRuleDescriptor] = []
|
||||
if hasattr(matcher_cls, "command"):
|
||||
descriptors.append(
|
||||
ActivationRuleDescriptor("matcher_command", command_like=True)
|
||||
)
|
||||
|
||||
rule = getattr(matcher_cls, "rule", None)
|
||||
checkers = getattr(rule, "checkers", ()) or ()
|
||||
for checker in checkers:
|
||||
@@ -530,48 +621,10 @@ def extract_matcher_rule_descriptors(
|
||||
):
|
||||
descriptors.append(ActivationRuleDescriptor("alconna", command_like=True))
|
||||
else:
|
||||
descriptors.append(_custom_rule_descriptor(call))
|
||||
descriptors.append(ActivationRuleDescriptor("custom"))
|
||||
return tuple(descriptors)
|
||||
|
||||
|
||||
def _custom_rule_descriptor(call: object) -> ActivationRuleDescriptor:
|
||||
keyword_regex = _extract_keyword_regex_pairs(call)
|
||||
if keyword_regex:
|
||||
return ActivationRuleDescriptor(
|
||||
"keyword_regex",
|
||||
keyword_regex,
|
||||
deterministic_text=True,
|
||||
)
|
||||
return ActivationRuleDescriptor("custom")
|
||||
|
||||
|
||||
def _extract_keyword_regex_pairs(
|
||||
call: object,
|
||||
) -> tuple[tuple[str, str, int], ...]:
|
||||
"""Recognize generic keyword + regex custom rules without plugin coupling."""
|
||||
|
||||
source = getattr(call, "key_pattern_list", None)
|
||||
if source is None:
|
||||
source = getattr(call, "keyword_patterns", None)
|
||||
if source is None:
|
||||
source = getattr(call, "patterns", None)
|
||||
if not isinstance(source, Iterable) or isinstance(source, str):
|
||||
return ()
|
||||
|
||||
pairs: list[tuple[str, str, int]] = []
|
||||
for item in source:
|
||||
if not isinstance(item, tuple | list) or len(item) < 2:
|
||||
continue
|
||||
keyword = str(item[0] or "").strip()
|
||||
pattern_obj = item[1]
|
||||
pattern = getattr(pattern_obj, "pattern", pattern_obj)
|
||||
if not keyword or not isinstance(pattern, str) or not pattern:
|
||||
continue
|
||||
flags = int(getattr(pattern_obj, "flags", 0) or 0)
|
||||
pairs.append((keyword, pattern, flags))
|
||||
return tuple(pairs)
|
||||
|
||||
|
||||
def normalize_rule_string_tuple(value: object) -> tuple[str, ...]:
|
||||
if isinstance(value, str):
|
||||
return (value,)
|
||||
@@ -633,13 +686,38 @@ def matcher_rule_matches_text(
|
||||
|
||||
for descriptor in descriptors:
|
||||
kind = descriptor.kind
|
||||
if kind == "regex":
|
||||
if kind in {"custom", "alconna"}:
|
||||
saw_unknown = True
|
||||
continue
|
||||
if kind in {"command", "shell_command"}:
|
||||
saw_deterministic = True
|
||||
commands: set[str] = set()
|
||||
collect_command_literals(descriptor.value, commands)
|
||||
normalized_commands = {
|
||||
normalized
|
||||
for item in commands
|
||||
if (normalized := normalize_command(item))
|
||||
}
|
||||
if any(
|
||||
matcher_command_matches(text, command)
|
||||
for text in plain_candidates
|
||||
for command in normalized_commands
|
||||
):
|
||||
matched_any = True
|
||||
else:
|
||||
return "miss"
|
||||
elif kind in {"regex", "regex_fullmatch"}:
|
||||
saw_deterministic = True
|
||||
pattern = str(descriptor.value or "")
|
||||
if not pattern:
|
||||
continue
|
||||
try:
|
||||
if re.search(pattern, message_text, descriptor.flags):
|
||||
matched = (
|
||||
re.fullmatch(pattern, message_text, descriptor.flags)
|
||||
if kind == "regex_fullmatch"
|
||||
else re.search(pattern, message_text, descriptor.flags)
|
||||
)
|
||||
if matched:
|
||||
matched_any = True
|
||||
else:
|
||||
return "miss"
|
||||
@@ -709,13 +787,6 @@ def matcher_rule_matches_text(
|
||||
matched_any = True
|
||||
else:
|
||||
return "miss"
|
||||
elif kind == "keyword_regex":
|
||||
saw_deterministic = True
|
||||
values = descriptor.value if isinstance(descriptor.value, tuple) else ()
|
||||
if _keyword_regex_matches(values, plain_candidates):
|
||||
matched_any = True
|
||||
else:
|
||||
return "miss"
|
||||
elif kind == "to_me":
|
||||
if not to_me:
|
||||
return "miss"
|
||||
@@ -729,43 +800,15 @@ def matcher_rule_matches_text(
|
||||
elif isinstance(types, tuple) and types:
|
||||
if not isinstance(event, types):
|
||||
return "miss"
|
||||
elif kind in {"custom", "alconna", "matcher_command"}:
|
||||
saw_unknown = True
|
||||
|
||||
if matched_any:
|
||||
return "match"
|
||||
if saw_deterministic:
|
||||
return "miss"
|
||||
return "unknown" if saw_unknown else "match"
|
||||
if saw_unknown:
|
||||
return "unknown"
|
||||
if saw_deterministic:
|
||||
return "miss"
|
||||
return "unknown"
|
||||
|
||||
|
||||
def _keyword_regex_matches(
|
||||
values: object,
|
||||
candidates: tuple[str, ...],
|
||||
) -> bool:
|
||||
if not isinstance(values, tuple):
|
||||
return False
|
||||
for item in values:
|
||||
if not isinstance(item, tuple | list) or len(item) < 3:
|
||||
continue
|
||||
keyword, pattern, flags = item[:3]
|
||||
keyword_text = str(keyword or "")
|
||||
pattern_text = str(pattern or "")
|
||||
if not keyword_text or not pattern_text:
|
||||
continue
|
||||
for text in candidates:
|
||||
if keyword_text not in text:
|
||||
continue
|
||||
try:
|
||||
if re.search(pattern_text, text, int(flags or 0)):
|
||||
return True
|
||||
except re.error:
|
||||
return False
|
||||
return False
|
||||
|
||||
|
||||
def extract_matcher_command_literals(
|
||||
matcher_cls: type[Matcher],
|
||||
) -> tuple[str, ...] | None:
|
||||
@@ -912,6 +955,293 @@ def collect_alconna_shortcuts(value: Any, target: set[str], depth: int = 0) -> N
|
||||
collect_alconna_shortcuts(nested, target, depth + 1)
|
||||
|
||||
|
||||
def extract_matcher_alconna_descriptors(
|
||||
matcher_cls: type[Matcher],
|
||||
) -> tuple[AlconnaDescriptor, ...]:
|
||||
descriptors: list[AlconnaDescriptor] = []
|
||||
rule = getattr(matcher_cls, "rule", None)
|
||||
checkers = getattr(rule, "checkers", ()) or ()
|
||||
for checker in checkers:
|
||||
call = getattr(checker, "call", None)
|
||||
if call is None:
|
||||
continue
|
||||
if call.__class__.__name__ != "AlconnaRule":
|
||||
continue
|
||||
if not call.__class__.__module__.startswith("nonebot_plugin_alconna.rule"):
|
||||
continue
|
||||
command = resolve_maybe_weakref(
|
||||
getattr(call, "command", None) or getattr(call, "alconna", None)
|
||||
)
|
||||
descriptor = _build_alconna_descriptor(call, command)
|
||||
if descriptor is not None:
|
||||
descriptors.append(descriptor)
|
||||
return tuple(descriptors)
|
||||
|
||||
|
||||
def _build_alconna_descriptor(call: Any, command: Any) -> AlconnaDescriptor | None:
|
||||
if command is None:
|
||||
return None
|
||||
command_text = str(getattr(command, "command", "") or "").strip()
|
||||
aliases = tuple(
|
||||
str(item).strip()
|
||||
for item in getattr(command, "aliases", ()) or ()
|
||||
if str(item).strip() and str(item).strip() != command_text
|
||||
)
|
||||
prefixes = tuple(
|
||||
str(item)
|
||||
for item in getattr(command, "prefixes", ()) or ()
|
||||
if isinstance(item, str)
|
||||
)
|
||||
meta = getattr(command, "meta", None)
|
||||
shortcuts = _extract_alconna_shortcut_descriptors(command)
|
||||
return AlconnaDescriptor(
|
||||
command=command_text,
|
||||
aliases=aliases,
|
||||
prefixes=prefixes,
|
||||
shortcuts=shortcuts,
|
||||
compact=bool(getattr(meta, "compact", False)),
|
||||
skip_for_unmatch=bool(getattr(call, "skip", True)),
|
||||
before_rule_count=_rule_checker_count(getattr(call, "before_rules", None)),
|
||||
after_rule_count=_rule_checker_count(getattr(call, "after_rules", None)),
|
||||
before_rule_known_safe=_alconna_rule_is_known_safe(
|
||||
getattr(call, "before_rules", None)
|
||||
),
|
||||
after_rule_known_safe=_alconna_rule_is_known_safe(
|
||||
getattr(call, "after_rules", None)
|
||||
),
|
||||
input_rewrite_extensions=_extract_alconna_input_rewrite_extensions(call),
|
||||
)
|
||||
|
||||
|
||||
def _rule_checker_count(rule: Any) -> int:
|
||||
checkers = getattr(rule, "checkers", ()) or ()
|
||||
with contextlib.suppress(TypeError):
|
||||
return len(checkers)
|
||||
return 1
|
||||
|
||||
|
||||
def _alconna_rule_is_known_safe(rule: Any) -> bool:
|
||||
"""Whether Alconna before/after Rule can be reasoned about statically.
|
||||
|
||||
We still do not execute these rules here. A rule is considered safe only
|
||||
when every checker is an official NoneBot rule whose negative result can be
|
||||
reproduced by the selector. Custom rules stay fail-open.
|
||||
"""
|
||||
|
||||
checkers = getattr(rule, "checkers", ()) or ()
|
||||
for checker in checkers:
|
||||
call = getattr(checker, "call", None)
|
||||
if call is None:
|
||||
return False
|
||||
call_module = call.__class__.__module__
|
||||
call_name = call.__class__.__name__
|
||||
if not call_module.startswith("nonebot.rule"):
|
||||
return False
|
||||
if call_name not in KNOWN_SAFE_RULE_NAMES:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _extract_alconna_shortcut_descriptors(
|
||||
command: Any,
|
||||
) -> tuple[AlconnaShortcutDescriptor, ...]:
|
||||
shortcuts: list[AlconnaShortcutDescriptor] = []
|
||||
with contextlib.suppress(Exception):
|
||||
from arclet.alconna import command_manager
|
||||
|
||||
raw_shortcuts = command_manager.get_shortcut(command) # type: ignore[arg-type]
|
||||
if isinstance(raw_shortcuts, dict):
|
||||
for key, args in raw_shortcuts.items():
|
||||
pattern = str(key or "").strip()
|
||||
if not pattern:
|
||||
continue
|
||||
shortcuts.append(
|
||||
AlconnaShortcutDescriptor(
|
||||
pattern=pattern,
|
||||
fuzzy=bool(getattr(args, "fuzzy", False)),
|
||||
flags=int(getattr(args, "flags", 0) or 0),
|
||||
)
|
||||
)
|
||||
if shortcuts:
|
||||
return tuple(shortcuts)
|
||||
fallback: set[str] = set()
|
||||
collect_alconna_shortcuts(command, fallback)
|
||||
return tuple(
|
||||
AlconnaShortcutDescriptor(pattern=item)
|
||||
for item in sorted(fallback)
|
||||
if item.strip()
|
||||
)
|
||||
|
||||
|
||||
def _extract_alconna_input_rewrite_extensions(call: Any) -> tuple[str, ...]:
|
||||
executor = getattr(call, "executor", None)
|
||||
if executor is None:
|
||||
return ()
|
||||
result: list[str] = []
|
||||
for attr in ("extensions", "_extensions", "exts", "context"):
|
||||
extensions = getattr(executor, attr, None)
|
||||
if not isinstance(extensions, list | tuple | set | frozenset):
|
||||
continue
|
||||
for extension in extensions:
|
||||
overrides = getattr(extension.__class__, "_overrides", None)
|
||||
if not isinstance(overrides, dict):
|
||||
overrides = getattr(extension, "_overrides", None)
|
||||
if not isinstance(overrides, dict):
|
||||
continue
|
||||
if not (
|
||||
bool(overrides.get("message_provider"))
|
||||
or bool(overrides.get("receive_wrapper"))
|
||||
):
|
||||
continue
|
||||
name = extension.__class__.__name__
|
||||
if name not in result:
|
||||
result.append(name)
|
||||
return tuple(result)
|
||||
|
||||
|
||||
def _alconna_can_prefilter(alconna: AlconnaDescriptor) -> bool:
|
||||
if not (alconna.command or alconna.aliases or alconna.shortcuts):
|
||||
return False
|
||||
if not (alconna.before_rule_known_safe and alconna.after_rule_known_safe):
|
||||
return False
|
||||
if any(name != "ReplyMergeExtension" for name in alconna.input_rewrite_extensions):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def matcher_alconna_head_matches(
|
||||
alconna: AlconnaDescriptor,
|
||||
texts: Iterable[str],
|
||||
context: ActivationContext,
|
||||
) -> ActivationDecision:
|
||||
if not _alconna_can_prefilter(alconna):
|
||||
return "unknown"
|
||||
if alconna.has_reply_merge_extension and _event_has_reply(context.event):
|
||||
return "unknown"
|
||||
|
||||
candidates = tuple(text.strip() for text in texts if text and text.strip())
|
||||
if not candidates:
|
||||
return "unknown"
|
||||
|
||||
saw_unknown = False
|
||||
command_heads = (alconna.command, *alconna.aliases)
|
||||
for text in candidates:
|
||||
for command in command_heads:
|
||||
if not command:
|
||||
continue
|
||||
decision = _alconna_command_head_matches(
|
||||
text,
|
||||
command,
|
||||
alconna.prefixes,
|
||||
compact=alconna.compact,
|
||||
)
|
||||
if decision == "match":
|
||||
return "match"
|
||||
if decision == "unknown":
|
||||
saw_unknown = True
|
||||
for shortcut in alconna.shortcuts:
|
||||
decision = _alconna_shortcut_matches(text, shortcut)
|
||||
if decision == "match":
|
||||
return "match"
|
||||
if decision == "unknown":
|
||||
saw_unknown = True
|
||||
return "unknown" if saw_unknown else "miss"
|
||||
|
||||
|
||||
def _alconna_command_head_matches(
|
||||
text: str,
|
||||
command: str,
|
||||
prefixes: tuple[str, ...],
|
||||
*,
|
||||
compact: bool,
|
||||
) -> ActivationDecision:
|
||||
normalized = command.strip()
|
||||
if not normalized:
|
||||
return "unknown"
|
||||
if normalized.startswith("re:"):
|
||||
pattern = normalized.removeprefix("re:").strip()
|
||||
if not pattern:
|
||||
return "unknown"
|
||||
for prefix in prefixes or ("",):
|
||||
try:
|
||||
if re.match(rf"^{re.escape(prefix)}(?:{pattern})", text):
|
||||
return "match"
|
||||
except re.error:
|
||||
return "unknown"
|
||||
return "miss"
|
||||
for prefix in prefixes or ("",):
|
||||
head = f"{prefix}{normalized}"
|
||||
if _alconna_literal_head_matches(text, head, compact=compact):
|
||||
return "match"
|
||||
return "miss"
|
||||
|
||||
|
||||
def _alconna_literal_head_matches(text: str, head: str, *, compact: bool) -> bool:
|
||||
if not text or not head:
|
||||
return False
|
||||
if text == head:
|
||||
return True
|
||||
if not text.startswith(head):
|
||||
return False
|
||||
if len(text) == len(head):
|
||||
return True
|
||||
rest = text[len(head) :]
|
||||
if rest and rest[0].isspace():
|
||||
return True
|
||||
return bool(compact or not head[-1].isascii())
|
||||
|
||||
|
||||
def _alconna_shortcut_matches(
|
||||
text: str,
|
||||
shortcut: AlconnaShortcutDescriptor,
|
||||
) -> ActivationDecision:
|
||||
pattern = shortcut.pattern.strip()
|
||||
if not pattern:
|
||||
return "unknown"
|
||||
normalized = normalize_shortcut_pattern(pattern)
|
||||
placeholder_match = _placeholder_shortcut_decision(text, normalized)
|
||||
if placeholder_match == "match":
|
||||
return "match"
|
||||
if placeholder_match == "unknown":
|
||||
return "unknown"
|
||||
if not is_regex_like_shortcut(normalized) and matcher_command_matches(
|
||||
text,
|
||||
normalized,
|
||||
):
|
||||
return "match"
|
||||
try:
|
||||
if shortcut.fuzzy:
|
||||
return "match" if re.match(f"^{pattern}", text, shortcut.flags) else "miss"
|
||||
return "match" if re.fullmatch(pattern, text, shortcut.flags) else "miss"
|
||||
except re.error:
|
||||
return "unknown"
|
||||
|
||||
|
||||
def _event_has_reply(event: object | None) -> bool:
|
||||
if event is None:
|
||||
return False
|
||||
with contextlib.suppress(Exception):
|
||||
message = event.get_message() # type: ignore[attr-defined]
|
||||
for segment in message:
|
||||
segment_type = getattr(segment, "type", None)
|
||||
if segment_type == "reply":
|
||||
return True
|
||||
if isinstance(segment, dict) and segment.get("type") == "reply":
|
||||
return True
|
||||
for attr in ("reply", "reply_message", "source"):
|
||||
with contextlib.suppress(Exception):
|
||||
if getattr(event, attr, None) is not None:
|
||||
return True
|
||||
raw_text = ""
|
||||
with contextlib.suppress(Exception):
|
||||
raw_text = str(event.get_message()) # type: ignore[attr-defined]
|
||||
lowered = raw_text.casefold()
|
||||
return any(
|
||||
marker in lowered
|
||||
for marker in ("[cq:reply", "type=reply", '"reply"', "'reply'")
|
||||
)
|
||||
|
||||
|
||||
def resolve_maybe_weakref(value: Any) -> Any:
|
||||
if isinstance(value, weakref.ReferenceType):
|
||||
resolved = value()
|
||||
@@ -1026,8 +1356,12 @@ def shortcut_matches_text(text: str, shortcut: str) -> bool:
|
||||
|
||||
|
||||
def placeholder_shortcut_matches(text: str, pattern: str) -> bool:
|
||||
return _placeholder_shortcut_decision(text, pattern) == "match"
|
||||
|
||||
|
||||
def _placeholder_shortcut_decision(text: str, pattern: str) -> ActivationDecision:
|
||||
if "{" not in pattern or "}" not in pattern:
|
||||
return False
|
||||
return "miss"
|
||||
pieces: list[str] = []
|
||||
last = 0
|
||||
for match in re.finditer(r"\{[^{}]+\}", pattern):
|
||||
@@ -1035,12 +1369,16 @@ def placeholder_shortcut_matches(text: str, pattern: str) -> bool:
|
||||
pieces.append(r"\S+")
|
||||
last = match.end()
|
||||
if not pieces:
|
||||
return False
|
||||
return "miss"
|
||||
pieces.append(re.escape(pattern[last:]))
|
||||
try:
|
||||
return re.match(rf"^{''.join(pieces)}(?:\s|$)", text) is not None
|
||||
return (
|
||||
"match"
|
||||
if re.match(rf"^{''.join(pieces)}(?:\s|$)", text) is not None
|
||||
else "miss"
|
||||
)
|
||||
except re.error:
|
||||
return False
|
||||
return "unknown"
|
||||
|
||||
|
||||
def is_regex_like_shortcut(pattern: str) -> bool:
|
||||
@@ -1072,21 +1410,27 @@ def matcher_matches_ai_route_heads(
|
||||
return False
|
||||
|
||||
|
||||
def _selection_is_guaranteed(
|
||||
def _is_throttleable_broad_passive(
|
||||
descriptor: HandlerDescriptor,
|
||||
context: ActivationContext,
|
||||
) -> bool:
|
||||
"""Return True for candidates that must not be budget-throttled."""
|
||||
"""Only broad, no-rule passive message matchers may be budget-throttled."""
|
||||
|
||||
if descriptor.temp or descriptor.lane == "system":
|
||||
return True
|
||||
if context.event_type != "message":
|
||||
return not descriptor.command_like
|
||||
return False
|
||||
if descriptor.temp or descriptor.lane == "system":
|
||||
return False
|
||||
if not descriptor.lane.startswith("passive_"):
|
||||
return False
|
||||
if descriptor.command_like or descriptor.deterministic_text:
|
||||
return False
|
||||
if descriptor.has_custom_rule or descriptor.rules:
|
||||
return False
|
||||
if descriptor.lane == "passive_http" and (
|
||||
context.has_url or _looks_like_rich_message(context.raw_text)
|
||||
):
|
||||
return True
|
||||
return False
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _looks_like_rich_message(text: str) -> bool:
|
||||
@@ -1106,34 +1450,11 @@ def _looks_like_rich_message(text: str) -> bool:
|
||||
)
|
||||
|
||||
|
||||
def _miss_reason(descriptor: HandlerDescriptor, context: ActivationContext) -> str:
|
||||
matcher_type = descriptor.matcher_type
|
||||
if matcher_type and matcher_type != context.event_type:
|
||||
return "type_miss"
|
||||
if context.event_type != "message" and descriptor.command_like:
|
||||
return "non_message_command"
|
||||
if descriptor.command_like:
|
||||
return "command_or_rule_miss"
|
||||
return "rule_miss"
|
||||
|
||||
|
||||
def _uncertain_budget_lane(descriptor: HandlerDescriptor) -> str:
|
||||
lane = descriptor.lane
|
||||
if lane.startswith("passive_"):
|
||||
return lane
|
||||
# Unknown command-like matchers are fail-open for compatibility, but they
|
||||
# should not fan out as unbounded command tasks when no deterministic signal
|
||||
# matched. Put them into the cheapest passive bucket.
|
||||
if lane.startswith("command_"):
|
||||
return "passive_light"
|
||||
return lane
|
||||
|
||||
|
||||
def _consume_uncertain_budget(
|
||||
def _consume_broad_passive_budget(
|
||||
descriptor: HandlerDescriptor,
|
||||
budget: dict[str, int],
|
||||
) -> bool:
|
||||
lane = _uncertain_budget_lane(descriptor)
|
||||
lane = descriptor.lane
|
||||
if lane not in budget:
|
||||
return True
|
||||
if budget[lane] <= 0:
|
||||
|
||||
@@ -15,19 +15,9 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
from zhenxun.configs.utils import PluginExtraData
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.user_console import UserConsole
|
||||
from zhenxun.services.auth_observability import (
|
||||
append_auth_decision_log,
|
||||
append_runtime_backpressure_log,
|
||||
build_auth_observability_report,
|
||||
)
|
||||
from zhenxun.services.cache.cache_containers import CacheDict
|
||||
from zhenxun.services.cache.runtime_cache import (
|
||||
PluginInfoMemoryCache,
|
||||
PluginLimitMemoryCache,
|
||||
)
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import is_overloaded, signal_overload
|
||||
from zhenxun.services.message_load import signal_overload
|
||||
from zhenxun.utils.enum import GoldHandle, PluginType
|
||||
from zhenxun.utils.exception import InsufficientGold
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
@@ -48,6 +38,7 @@ from .auth.context import (
|
||||
set_route_modules,
|
||||
store_permission_context,
|
||||
)
|
||||
from .auth.data_provider import DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
from .auth.exception import (
|
||||
IsSuperuserException,
|
||||
PermissionExemption,
|
||||
@@ -55,7 +46,6 @@ from .auth.exception import (
|
||||
)
|
||||
from .auth_activation import (
|
||||
ActivationContext,
|
||||
ActivationResult,
|
||||
HandlerActivationIndex,
|
||||
classify_matcher_lane,
|
||||
extract_matcher_alconna_shortcuts,
|
||||
@@ -147,9 +137,7 @@ _ROUTE_COMMAND_MAP: dict[str, set[str]] = {}
|
||||
_ROUTE_PREFIX_MAP: dict[str, set[str]] = {}
|
||||
_ROUTE_MODULES_WITH_COMMANDS: set[str] = set()
|
||||
MATCHER_ROUTE_PREFILTER_TTL = AUTH_DISPATCH_RUNTIME_CONFIG.matcher_route_prefilter_ttl
|
||||
PREFILTER_STATS_LOG_INTERVAL = AUTH_DISPATCH_RUNTIME_CONFIG.prefilter_stats_log_interval
|
||||
CACHE_SWEEP_INTERVAL = AUTH_DISPATCH_RUNTIME_CONFIG.cache_sweep_interval
|
||||
DISPATCH_STATS_LOG_INTERVAL = AUTH_DISPATCH_RUNTIME_CONFIG.dispatch_stats_log_interval
|
||||
|
||||
# 全局信号量与计数器
|
||||
HOOKS_ACTIVE_COUNT = 0
|
||||
@@ -182,56 +170,6 @@ _CHECK_MATCHER_ROUTE_CACHE = CacheDict(
|
||||
)
|
||||
|
||||
|
||||
_PREFILTER_STATS = {
|
||||
"checked": 0,
|
||||
"skipped": 0,
|
||||
"before_task_checked": 0,
|
||||
"before_task_skipped": 0,
|
||||
"inside_task_checked": 0,
|
||||
"inside_task_skipped": 0,
|
||||
"type_miss": 0,
|
||||
"route_miss": 0,
|
||||
"command_miss": 0,
|
||||
"empty_text": 0,
|
||||
}
|
||||
_PREFILTER_LAST_LOG = 0.0
|
||||
_DISPATCH_SELECTED = 0
|
||||
_DISPATCH_SKIPPED = 0
|
||||
_DISPATCH_SELECTED_BY_LANE: dict[str, int] = {
|
||||
"command_exact": 0,
|
||||
"command_shortcut": 0,
|
||||
"command_regex": 0,
|
||||
"system": 0,
|
||||
"passive_light": 0,
|
||||
"passive_db": 0,
|
||||
"passive_http": 0,
|
||||
"passive_ai": 0,
|
||||
"passive_render": 0,
|
||||
}
|
||||
_DISPATCH_SKIPPED_BY_LANE: dict[str, int] = {
|
||||
"command_exact": 0,
|
||||
"command_shortcut": 0,
|
||||
"command_regex": 0,
|
||||
"system": 0,
|
||||
"passive_light": 0,
|
||||
"passive_db": 0,
|
||||
"passive_http": 0,
|
||||
"passive_ai": 0,
|
||||
"passive_render": 0,
|
||||
}
|
||||
_DISPATCH_LANE_WAIT_MS: dict[str, float] = {
|
||||
"command_exact": 0.0,
|
||||
"command_shortcut": 0.0,
|
||||
"command_regex": 0.0,
|
||||
"system": 0.0,
|
||||
"passive_light": 0.0,
|
||||
"passive_db": 0.0,
|
||||
"passive_http": 0.0,
|
||||
"passive_ai": 0.0,
|
||||
"passive_render": 0.0,
|
||||
}
|
||||
_DISPATCH_LAST_LOG = 0.0
|
||||
_DISPATCH_SHADOW_LAST_LOG = 0.0
|
||||
_CACHE_SWEEP_TASK: asyncio.Task | None = None
|
||||
_BOT_WAKE_COMMAND_PATTERN = re.compile(r"^bot醒来(?:\s+\S+)?$", re.IGNORECASE)
|
||||
_BOT_WAKE_CANONICAL_PATTERN = re.compile(
|
||||
@@ -367,12 +305,6 @@ def _is_bot_wake_command(module: str, text: str | None) -> bool:
|
||||
)
|
||||
|
||||
|
||||
def _debug_log(message: str, *args, **kwargs) -> None:
|
||||
if is_overloaded():
|
||||
return
|
||||
logger.debug(message, *args, **kwargs)
|
||||
|
||||
|
||||
def _is_command_matcher_class(matcher_cls: type[Matcher]) -> bool:
|
||||
if matcher_cls in _MATCHER_COMMAND_TYPE_CACHE:
|
||||
return _MATCHER_COMMAND_TYPE_CACHE[matcher_cls]
|
||||
@@ -728,41 +660,6 @@ def _activation_context_from_dispatch(
|
||||
)
|
||||
|
||||
|
||||
def _record_dispatch_selection(lane: str, selected: bool, wait_ms: float = 0.0) -> None:
|
||||
global _DISPATCH_LAST_LOG, _DISPATCH_SELECTED, _DISPATCH_SKIPPED
|
||||
if lane == "command":
|
||||
lane = "command_exact"
|
||||
lane = lane if lane in _DISPATCH_SELECTED_BY_LANE else "passive_light"
|
||||
if selected:
|
||||
_DISPATCH_SELECTED += 1
|
||||
_DISPATCH_SELECTED_BY_LANE[lane] += 1
|
||||
_DISPATCH_LANE_WAIT_MS[lane] += wait_ms
|
||||
else:
|
||||
_DISPATCH_SKIPPED += 1
|
||||
_DISPATCH_SKIPPED_BY_LANE[lane] += 1
|
||||
|
||||
now = time.monotonic()
|
||||
if now - _DISPATCH_LAST_LOG < DISPATCH_STATS_LOG_INTERVAL or is_overloaded():
|
||||
return
|
||||
_DISPATCH_LAST_LOG = now
|
||||
wait_snapshot = {
|
||||
lane: round(wait, 2) for lane, wait in _DISPATCH_LANE_WAIT_MS.items()
|
||||
}
|
||||
lane_snapshot = " ".join(
|
||||
f"{lane}={count}" for lane, count in _DISPATCH_SELECTED_BY_LANE.items()
|
||||
)
|
||||
_debug_log(
|
||||
(
|
||||
"dispatch stats: "
|
||||
f"selected={_DISPATCH_SELECTED} "
|
||||
f"skipped={_DISPATCH_SKIPPED} "
|
||||
f"{lane_snapshot} "
|
||||
f"wait_ms={wait_snapshot}"
|
||||
),
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
|
||||
|
||||
def _new_dispatch_budget() -> dict[str, int]:
|
||||
return dict(_DISPATCH_LANE_LIMITS)
|
||||
|
||||
@@ -775,52 +672,6 @@ def _merge_dispatch_budget(
|
||||
target[lane] = source.get(lane, target.get(lane, 0))
|
||||
|
||||
|
||||
def _record_activation_result(activation_result: ActivationResult) -> None:
|
||||
for lane, count in activation_result.skipped_by_lane.items():
|
||||
for _ in range(count):
|
||||
_record_dispatch_selection(lane, False)
|
||||
|
||||
|
||||
def _compact_counter(counter: dict[str, int], limit: int = 6) -> str:
|
||||
if not counter:
|
||||
return "-"
|
||||
items = sorted(counter.items(), key=lambda item: item[1], reverse=True)[:limit]
|
||||
return ",".join(f"{key}={value}" for key, value in items)
|
||||
|
||||
|
||||
def _debug_activation_shadow(
|
||||
*,
|
||||
priority: int,
|
||||
activation_result: ActivationResult,
|
||||
context: EventDispatchContext,
|
||||
) -> None:
|
||||
global _DISPATCH_SHADOW_LAST_LOG
|
||||
if is_overloaded():
|
||||
return
|
||||
now = time.monotonic()
|
||||
if now - _DISPATCH_SHADOW_LAST_LOG < DISPATCH_STATS_LOG_INTERVAL:
|
||||
return
|
||||
_DISPATCH_SHADOW_LAST_LOG = now
|
||||
text_hint = (context.plain_text or context.trie_raw_command or "")[:48]
|
||||
_debug_log(
|
||||
(
|
||||
"dispatch shadow: "
|
||||
f"priority={priority} "
|
||||
f"event={context.event_type} "
|
||||
f"selected={activation_result.candidate_count}/"
|
||||
f"{activation_result.total_descriptors} "
|
||||
f"selected_lane={_compact_counter(activation_result.selected_by_lane)} "
|
||||
f"skipped_lane={_compact_counter(activation_result.skipped_by_lane)} "
|
||||
f"selected_reason="
|
||||
f"{_compact_counter(activation_result.selected_by_reason)} "
|
||||
f"skipped_reason="
|
||||
f"{_compact_counter(activation_result.skipped_by_reason)} "
|
||||
f"text={text_hint!r}"
|
||||
),
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
|
||||
|
||||
def _auth_scope_key(context: EventContext) -> str:
|
||||
group_id = context.group_id or ""
|
||||
channel_id = context.channel_id or ""
|
||||
@@ -876,7 +727,6 @@ async def _dispatch_lane_section(lane: str):
|
||||
wait_ms = (time.perf_counter() - started) * 1000
|
||||
if wait_ms >= AUTH_OVERLOAD_LANE_WAIT_MS:
|
||||
signal_overload(2.0)
|
||||
_record_dispatch_selection(lane, True, wait_ms=wait_ms)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
@@ -891,11 +741,6 @@ def get_dispatch_snapshot() -> dict[str, object]:
|
||||
value = getattr(semaphore, "_value", limit)
|
||||
lane_active[lane] = max(limit - int(value), 0)
|
||||
return {
|
||||
"selected": _DISPATCH_SELECTED,
|
||||
"skipped": _DISPATCH_SKIPPED,
|
||||
"selected_by_lane": dict(_DISPATCH_SELECTED_BY_LANE),
|
||||
"skipped_by_lane": dict(_DISPATCH_SKIPPED_BY_LANE),
|
||||
"lane_wait_ms": dict(_DISPATCH_LANE_WAIT_MS),
|
||||
"lane_active": lane_active,
|
||||
"lane_limits": dict(_DISPATCH_LANE_LIMITS),
|
||||
}
|
||||
@@ -975,58 +820,6 @@ async def _run_selected_matcher(
|
||||
)
|
||||
|
||||
|
||||
def _record_prefilter_stats(
|
||||
skipped: bool,
|
||||
reason: str | None,
|
||||
stage: str = "inside_task",
|
||||
) -> None:
|
||||
global _PREFILTER_LAST_LOG
|
||||
_PREFILTER_STATS["checked"] += 1
|
||||
if skipped:
|
||||
_PREFILTER_STATS["skipped"] += 1
|
||||
if stage == "before_task":
|
||||
_PREFILTER_STATS["before_task_checked"] += 1
|
||||
if skipped:
|
||||
_PREFILTER_STATS["before_task_skipped"] += 1
|
||||
else:
|
||||
_PREFILTER_STATS["inside_task_checked"] += 1
|
||||
if skipped:
|
||||
_PREFILTER_STATS["inside_task_skipped"] += 1
|
||||
if reason == "type_miss":
|
||||
_PREFILTER_STATS["type_miss"] += 1
|
||||
elif reason == "route_miss":
|
||||
_PREFILTER_STATS["route_miss"] += 1
|
||||
elif reason == "command_miss":
|
||||
_PREFILTER_STATS["command_miss"] += 1
|
||||
elif reason == "empty_text":
|
||||
_PREFILTER_STATS["empty_text"] += 1
|
||||
|
||||
if _PREFILTER_STATS["checked"] % 1024 == 0:
|
||||
with contextlib.suppress(Exception):
|
||||
_ = len(_CHECK_MATCHER_ROUTE_CACHE)
|
||||
|
||||
now = time.monotonic()
|
||||
if now - _PREFILTER_LAST_LOG < PREFILTER_STATS_LOG_INTERVAL or is_overloaded():
|
||||
return
|
||||
_PREFILTER_LAST_LOG = now
|
||||
_debug_log(
|
||||
(
|
||||
"matcher prefilter stats: "
|
||||
f"checked={_PREFILTER_STATS['checked']} "
|
||||
f"skipped={_PREFILTER_STATS['skipped']} "
|
||||
f"before_task={_PREFILTER_STATS['before_task_skipped']}/"
|
||||
f"{_PREFILTER_STATS['before_task_checked']} "
|
||||
f"inside_task={_PREFILTER_STATS['inside_task_skipped']}/"
|
||||
f"{_PREFILTER_STATS['inside_task_checked']} "
|
||||
f"type_miss={_PREFILTER_STATS['type_miss']} "
|
||||
f"route_miss={_PREFILTER_STATS['route_miss']} "
|
||||
f"command_miss={_PREFILTER_STATS['command_miss']} "
|
||||
f"empty_text={_PREFILTER_STATS['empty_text']}"
|
||||
),
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
|
||||
|
||||
_MAX_MATCHER_CACHE = 512
|
||||
|
||||
|
||||
@@ -1038,8 +831,6 @@ _SELECTOR_DEPS = HandleEventSelectorDependencies(
|
||||
activation_context_from_dispatch=_activation_context_from_dispatch,
|
||||
new_dispatch_budget=_new_dispatch_budget,
|
||||
dispatch_lane_for_matcher=_dispatch_lane_for_matcher,
|
||||
record_activation_result=_record_activation_result,
|
||||
debug_activation_shadow=_debug_activation_shadow,
|
||||
merge_dispatch_budget=_merge_dispatch_budget,
|
||||
build_matcher_state=_build_matcher_state,
|
||||
run_selected_matcher=_run_selected_matcher,
|
||||
@@ -1068,15 +859,6 @@ async def _get_route_context(text: str, event_cache: dict | None) -> set[str]:
|
||||
return matched
|
||||
|
||||
|
||||
def _get_auth_route_precheck_deps() -> dict:
|
||||
return {
|
||||
"route_modules_with_commands": _ROUTE_MODULES_WITH_COMMANDS,
|
||||
"get_route_context": _get_route_context,
|
||||
"is_command_matcher_class": _is_command_matcher_class,
|
||||
"matcher_has_alconna_shortcuts": _matcher_has_alconna_shortcuts,
|
||||
}
|
||||
|
||||
|
||||
async def _cache_sweep_loop() -> None:
|
||||
while True:
|
||||
await asyncio.sleep(CACHE_SWEEP_INTERVAL)
|
||||
@@ -1145,7 +927,8 @@ async def _has_limits_cached(
|
||||
if event_cache is not None:
|
||||
event_cache.setdefault("module_limits_ready", {})[module] = True
|
||||
return has_limits
|
||||
limit_entries = PluginLimitMemoryCache.get_limits_if_ready(module)
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
limit_entries = provider.get_module_limits_if_ready(module)
|
||||
if limit_entries is not None:
|
||||
has_limits = bool(limit_entries)
|
||||
module_limit_cache[module] = has_limits
|
||||
@@ -1278,16 +1061,17 @@ async def _get_plugin_cache_first(
|
||||
*,
|
||||
allow_cache_load: bool,
|
||||
) -> tuple[PluginInfo | None, bool]:
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
plugin = None
|
||||
if event_cache is not None:
|
||||
plugin_cache = event_cache.setdefault("plugin_cache", {})
|
||||
if module in plugin_cache:
|
||||
return cast(PluginInfo | None, plugin_cache[module]), False
|
||||
|
||||
plugin = PluginInfoMemoryCache.get_by_module_if_ready(module)
|
||||
cache_miss = plugin is None and not PluginInfoMemoryCache.is_loaded()
|
||||
plugin = provider.get_plugin_if_ready(module)
|
||||
cache_miss = plugin is None and not provider.plugin_cache_loaded()
|
||||
if plugin is None and allow_cache_load:
|
||||
plugin = await PluginInfoMemoryCache.get_by_module(module)
|
||||
plugin = await provider.get_plugin(module)
|
||||
cache_miss = False
|
||||
if event_cache is not None:
|
||||
event_cache.setdefault("plugin_cache", {})[module] = plugin
|
||||
@@ -1351,7 +1135,6 @@ async def reserve_gold(
|
||||
)
|
||||
except InsufficientGold:
|
||||
raise
|
||||
await DataAccess(UserConsole).clear_cache(user_id=user_id)
|
||||
logger.debug(f"预扣功能花费金币: {cost_gold}", LOGGER_COMMAND, session=session)
|
||||
return reservation
|
||||
|
||||
@@ -1384,14 +1167,12 @@ async def _record_backpressure(
|
||||
action: str,
|
||||
duration_ms: float = 0.0,
|
||||
) -> None:
|
||||
await append_runtime_backpressure_log(
|
||||
scope_key=lane_context.scope_key,
|
||||
reason=reason,
|
||||
lane=lane_context.lane,
|
||||
action=action,
|
||||
queue_size=lane_context.queue_size,
|
||||
active_count=HOOKS_ACTIVE_COUNT,
|
||||
duration_ms=duration_ms,
|
||||
logger.debug(
|
||||
"auth backpressure: "
|
||||
f"scope={lane_context.scope_key}, lane={lane_context.lane}, "
|
||||
f"reason={reason}, action={action}, queue={lane_context.queue_size}, "
|
||||
f"active={HOOKS_ACTIVE_COUNT}, duration_ms={duration_ms:.1f}",
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
|
||||
|
||||
@@ -1437,7 +1218,6 @@ async def _prepare_auth_state(
|
||||
context: EventContext,
|
||||
bot: Bot,
|
||||
event_cache: dict | None,
|
||||
route_skip_checks: bool,
|
||||
skip_ban: bool,
|
||||
hook_recorder: HookTraceRecorder,
|
||||
state: dict | None,
|
||||
@@ -1502,7 +1282,6 @@ async def _prepare_auth_state(
|
||||
|
||||
policy_context = PolicyContext(
|
||||
snapshot=snapshot,
|
||||
route_skip_checks=route_skip_checks,
|
||||
allow_sleep_bypass=_is_bot_wake_command(module, context.plain_text),
|
||||
allow_group_sleep_bypass=_is_group_wake_command(plugin, context.plain_text),
|
||||
)
|
||||
@@ -1522,7 +1301,6 @@ async def _prepare_auth_state_with_fallback(
|
||||
context: EventContext,
|
||||
bot: Bot,
|
||||
event_cache: dict | None,
|
||||
route_skip_checks: bool,
|
||||
skip_ban: bool,
|
||||
hook_recorder: HookTraceRecorder,
|
||||
state: dict | None,
|
||||
@@ -1533,7 +1311,6 @@ async def _prepare_auth_state_with_fallback(
|
||||
context=context,
|
||||
bot=bot,
|
||||
event_cache=event_cache,
|
||||
route_skip_checks=route_skip_checks,
|
||||
skip_ban=skip_ban,
|
||||
hook_recorder=hook_recorder,
|
||||
state=state,
|
||||
@@ -1548,7 +1325,6 @@ async def _prepare_auth_state_with_fallback(
|
||||
context=context,
|
||||
bot=bot,
|
||||
event_cache=event_cache,
|
||||
route_skip_checks=route_skip_checks,
|
||||
skip_ban=skip_ban,
|
||||
hook_recorder=hook_recorder,
|
||||
state=state,
|
||||
@@ -1614,16 +1390,26 @@ async def _reserve_limit_side_effect(
|
||||
async def _resolve_cost_gold(
|
||||
*,
|
||||
prep: AuthPreparation,
|
||||
route_skip_checks: bool,
|
||||
hook_recorder: HookTraceRecorder,
|
||||
session: Uninfo,
|
||||
) -> int:
|
||||
plugin = prep.plugin
|
||||
if route_skip_checks or prep.profile.cost_gold <= 0:
|
||||
if prep.profile.cost_gold <= 0:
|
||||
hook_recorder.set("cost_gold", "skipped")
|
||||
return 0
|
||||
cost_start = time.time()
|
||||
try:
|
||||
if prep.user is None:
|
||||
user_start = time.time()
|
||||
prep.user = await with_timeout(
|
||||
UserConsole.get_user(
|
||||
prep.permission_context.user_id,
|
||||
PlatformUtils.get_platform(session),
|
||||
),
|
||||
name="get_cost_user",
|
||||
)
|
||||
prep.permission_context.user = prep.user
|
||||
hook_recorder.set("get_cost_user", f"{time.time() - user_start:.3f}s")
|
||||
cost_gold = await with_timeout(
|
||||
get_plugin_cost(
|
||||
prep.user,
|
||||
@@ -1649,7 +1435,6 @@ async def _run_auth_hooks(
|
||||
prep: AuthPreparation,
|
||||
session: Uninfo,
|
||||
event_cache: dict | None,
|
||||
route_skip_checks: bool,
|
||||
lane_context: AuthLaneContext,
|
||||
hook_recorder: HookTraceRecorder,
|
||||
side_effect_commit: SideEffectCommit,
|
||||
@@ -1660,26 +1445,23 @@ async def _run_auth_hooks(
|
||||
await _enter_hooks_section(lane_context)
|
||||
hook_tasks = []
|
||||
try:
|
||||
if not route_skip_checks:
|
||||
has_limits = await _has_limits_cached(
|
||||
profile.module,
|
||||
event_cache,
|
||||
known=profile.has_limit,
|
||||
)
|
||||
if has_limits:
|
||||
hook_tasks.append(
|
||||
time_hook(
|
||||
_reserve_limit_side_effect(
|
||||
prep=prep,
|
||||
session=session,
|
||||
side_effect_commit=side_effect_commit,
|
||||
),
|
||||
"auth_limit",
|
||||
hook_recorder,
|
||||
)
|
||||
has_limits = await _has_limits_cached(
|
||||
profile.module,
|
||||
event_cache,
|
||||
known=profile.has_limit,
|
||||
)
|
||||
if has_limits:
|
||||
hook_tasks.append(
|
||||
time_hook(
|
||||
_reserve_limit_side_effect(
|
||||
prep=prep,
|
||||
session=session,
|
||||
side_effect_commit=side_effect_commit,
|
||||
),
|
||||
"auth_limit",
|
||||
hook_recorder,
|
||||
)
|
||||
else:
|
||||
hook_recorder.set("auth_limit", "skipped")
|
||||
)
|
||||
else:
|
||||
hook_recorder.set("auth_limit", "skipped")
|
||||
|
||||
@@ -1703,13 +1485,6 @@ async def _run_auth_hooks(
|
||||
return time.time() - hooks_start
|
||||
|
||||
|
||||
async def build_auth_decision_backpressure_report(
|
||||
*,
|
||||
hours: float = 24.0,
|
||||
) -> dict:
|
||||
return await build_auth_observability_report(hours=hours)
|
||||
|
||||
|
||||
_AUTH_PIPELINE_DEPS = AuthPipelineDependencies(
|
||||
route_modules_with_commands=_ROUTE_MODULES_WITH_COMMANDS,
|
||||
get_route_context=_get_route_context,
|
||||
@@ -1726,7 +1501,6 @@ _AUTH_PIPELINE_DEPS = AuthPipelineDependencies(
|
||||
run_auth_hooks=_run_auth_hooks,
|
||||
bot_filter=bot_filter,
|
||||
reserve_gold=reserve_gold,
|
||||
append_auth_decision_log=append_auth_decision_log,
|
||||
insufficient_gold_error=InsufficientGold,
|
||||
logger=logger,
|
||||
log_command=LOGGER_COMMAND,
|
||||
|
||||
@@ -31,8 +31,6 @@ class HandleEventSelectorDependencies:
|
||||
activation_context_from_dispatch: Callable[[EventDispatchContext, Event], Any]
|
||||
new_dispatch_budget: Callable[[], dict[str, int]]
|
||||
dispatch_lane_for_matcher: Callable[[type[Matcher], EventDispatchContext], str]
|
||||
record_activation_result: Callable[[Any], None]
|
||||
debug_activation_shadow: Callable[..., None]
|
||||
merge_dispatch_budget: Callable[[dict[str, int], dict[str, int]], None]
|
||||
build_matcher_state: Callable[[dict], dict]
|
||||
run_selected_matcher: Callable[..., Awaitable[None]]
|
||||
@@ -43,11 +41,81 @@ _ORIGINAL_HANDLE_EVENT: Callable[..., Awaitable[None]] | None = None
|
||||
_ORIGINAL_ADAPTER_HANDLE_EVENTS: dict[object, object] = {}
|
||||
|
||||
|
||||
def _trim_leading_text(message: Any) -> None:
|
||||
if not message:
|
||||
return
|
||||
segment = message[0]
|
||||
if getattr(segment, "type", None) != "text":
|
||||
return
|
||||
data = getattr(segment, "data", None)
|
||||
if not isinstance(data, dict):
|
||||
return
|
||||
data["text"] = str(data.get("text", "")).lstrip("\xa0").lstrip()
|
||||
if not data["text"]:
|
||||
del message[0]
|
||||
|
||||
|
||||
def _is_self_mention_segment(segment: Any, bot: Bot) -> bool:
|
||||
segment_type = getattr(segment, "type", None)
|
||||
if segment_type not in {"mention_user", "group_mention_user"}:
|
||||
return False
|
||||
data = getattr(segment, "data", None)
|
||||
if not isinstance(data, dict):
|
||||
return False
|
||||
if data.get("is_you") or data.get("is_bot"):
|
||||
return True
|
||||
user_id = data.get("user_id")
|
||||
return user_id is not None and str(user_id) == str(bot.self_id)
|
||||
|
||||
|
||||
def _ensure_nonempty_qq_message(message: Any) -> None:
|
||||
if message:
|
||||
return
|
||||
with contextlib.suppress(Exception):
|
||||
message_module = importlib.import_module("nonebot.adapters.qq.message")
|
||||
MessageSegment = getattr(message_module, "MessageSegment")
|
||||
message.append(MessageSegment.text(""))
|
||||
|
||||
|
||||
def _normalize_qq_self_at_message(bot: Bot, event: Event) -> None:
|
||||
"""Remove the leading bot mention left by QQ official @ events.
|
||||
|
||||
nonebot-adapter-qq's @ event branches can mark ``to_me`` but keep the
|
||||
synthetic leading mention segment. Alconna command heads then see
|
||||
``<@bot>命令`` and fail to match, while regular ``event.get_plaintext()``
|
||||
still looks correct. Normalizing here keeps the runtime behavior aligned
|
||||
with OneBot/standard to_me preprocessing without changing plugin code or
|
||||
database state.
|
||||
"""
|
||||
if event.__class__.__name__ not in {
|
||||
"AtMessageCreateEvent",
|
||||
"GroupAtMessageCreateEvent",
|
||||
}:
|
||||
return
|
||||
adapter = getattr(bot, "adapter", None)
|
||||
adapter_name = ""
|
||||
get_name = getattr(adapter, "get_name", None)
|
||||
if callable(get_name):
|
||||
with contextlib.suppress(Exception):
|
||||
adapter_name = str(get_name()).lower()
|
||||
if adapter_name != "qq":
|
||||
return
|
||||
with contextlib.suppress(Exception):
|
||||
message = event.get_message()
|
||||
if not message or not _is_self_mention_segment(message[0], bot):
|
||||
return
|
||||
message.pop(0)
|
||||
setattr(event, "to_me", True)
|
||||
_trim_leading_text(message)
|
||||
_ensure_nonempty_qq_message(message)
|
||||
|
||||
|
||||
async def patched_handle_event(
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
deps: HandleEventSelectorDependencies,
|
||||
) -> None:
|
||||
_normalize_qq_self_at_message(bot, event)
|
||||
show_log = True
|
||||
escape_tag = getattr(nb_message, "escape_tag")
|
||||
logger_ = getattr(nb_message, "logger")
|
||||
@@ -153,12 +221,6 @@ async def patched_handle_event(
|
||||
|
||||
if activation_result is not None:
|
||||
selected_matchers = activation_result.selected
|
||||
deps.record_activation_result(activation_result)
|
||||
deps.debug_activation_shadow(
|
||||
priority=priority,
|
||||
activation_result=activation_result,
|
||||
context=dispatch_context,
|
||||
)
|
||||
if (
|
||||
activation_result.candidate_count
|
||||
> deps.overload_selected_threshold
|
||||
@@ -186,12 +248,6 @@ async def patched_handle_event(
|
||||
except Exception:
|
||||
single_result = None
|
||||
if single_result is not None:
|
||||
deps.record_activation_result(single_result)
|
||||
deps.debug_activation_shadow(
|
||||
priority=priority,
|
||||
activation_result=single_result,
|
||||
context=dispatch_context,
|
||||
)
|
||||
deps.merge_dispatch_budget(
|
||||
priority_budget,
|
||||
single_budget,
|
||||
@@ -238,6 +294,7 @@ def install_handle_event_selector(deps: HandleEventSelectorDependencies) -> None
|
||||
for module_name in (
|
||||
"nonebot.adapters.onebot.v11.bot",
|
||||
"nonebot.adapters.onebot.v12.bot",
|
||||
"nonebot.adapters.qq.bot",
|
||||
"onebug.mixin.process",
|
||||
):
|
||||
with contextlib.suppress(Exception):
|
||||
|
||||
@@ -26,13 +26,11 @@ from .auth.context import (
|
||||
)
|
||||
from .auth_checker import (
|
||||
LimitManager,
|
||||
_get_auth_route_precheck_deps,
|
||||
_get_route_context,
|
||||
auth,
|
||||
start_auth_runtime_tasks,
|
||||
stop_auth_runtime_tasks,
|
||||
)
|
||||
from .auth_route import route_precheck
|
||||
|
||||
_SKIP_AUTH_PLUGINS = {"chat_history", "chat_message"}
|
||||
_BOT_CONNECT_TS: float | None = None
|
||||
@@ -113,9 +111,6 @@ async def _auth_preprocessor(
|
||||
)
|
||||
set_route_modules(state, event_context, route_modules)
|
||||
|
||||
if await route_precheck(matcher, event_context, **_get_auth_route_precheck_deps()):
|
||||
return
|
||||
|
||||
try:
|
||||
await auth(
|
||||
matcher,
|
||||
|
||||
@@ -10,7 +10,6 @@ from nonebot.adapters import Bot, Event
|
||||
from nonebot.matcher import Matcher
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.services.message_load import is_overloaded
|
||||
from zhenxun.utils.utils import EntityIDs
|
||||
|
||||
from .auth.context import (
|
||||
@@ -86,7 +85,6 @@ class AuthPipelineContext:
|
||||
event_cache: dict | None = None
|
||||
text: str = ""
|
||||
route_modules: set[str] | None = None
|
||||
route_skip_checks: bool = False
|
||||
is_command_matcher: bool = False
|
||||
lane_context: AuthLaneContext | None = None
|
||||
side_effect_cache: PermissionSideEffectCache | None = None
|
||||
@@ -149,7 +147,6 @@ class AuthPipelineDependencies:
|
||||
run_auth_hooks: Callable[..., Awaitable[float]]
|
||||
bot_filter: Callable[..., None]
|
||||
reserve_gold: Callable[..., Awaitable[Any]]
|
||||
append_auth_decision_log: Callable[..., Awaitable[None]]
|
||||
insufficient_gold_error: type[Exception]
|
||||
logger: Any
|
||||
log_command: str
|
||||
@@ -173,10 +170,7 @@ def apply_policy_precheck(
|
||||
hook_recorder.set("auth_core", f"policy:{decision.reason}")
|
||||
if decision.denied:
|
||||
raise_for_policy(decision, deps.policy_skip_message(decision.reason))
|
||||
if decision.allowed and decision.reason in {
|
||||
"hidden_plugin_skip_auth",
|
||||
"route_miss_skip_checks",
|
||||
}:
|
||||
if decision.allowed and decision.reason in {"hidden_plugin_skip_auth"}:
|
||||
flags.should_return_allowed = True
|
||||
return flags
|
||||
|
||||
@@ -260,17 +254,16 @@ async def route_gate_stage(
|
||||
if ctx.route_modules is None:
|
||||
ctx.route_modules = await deps.get_route_context(ctx.text, ctx.event_cache)
|
||||
set_route_modules(ctx.state, ctx.event_context, ctx.route_modules)
|
||||
ctx.route_skip_checks = (
|
||||
route_missed = (
|
||||
ctx.is_command_matcher
|
||||
and ctx.module in deps.route_modules_with_commands
|
||||
and ctx.module not in ctx.route_modules
|
||||
and not deps.matcher_has_alconna_shortcuts(type(ctx.matcher))
|
||||
)
|
||||
if ctx.route_skip_checks:
|
||||
if route_missed:
|
||||
if ctx.event_cache is not None:
|
||||
ctx.event_cache["route_skip"] = True
|
||||
ctx.event_cache["route_miss_after_native_match"] = True
|
||||
_recorder(ctx).set("route", "miss")
|
||||
ctx.stop(allowed=True, effect="allow", reason="route_miss_skip_checks")
|
||||
|
||||
|
||||
async def prepare_snapshot_stage(
|
||||
@@ -282,7 +275,6 @@ async def prepare_snapshot_stage(
|
||||
context=ctx.event_context,
|
||||
bot=ctx.bot,
|
||||
event_cache=ctx.event_cache,
|
||||
route_skip_checks=ctx.route_skip_checks,
|
||||
skip_ban=ctx.skip_ban,
|
||||
hook_recorder=ctx.hook_recorder,
|
||||
state=ctx.state,
|
||||
@@ -305,7 +297,6 @@ async def policy_precheck_stage(
|
||||
context=ctx.event_context,
|
||||
bot=ctx.bot,
|
||||
event_cache=ctx.event_cache,
|
||||
route_skip_checks=ctx.route_skip_checks,
|
||||
skip_ban=ctx.skip_ban,
|
||||
hook_recorder=ctx.hook_recorder,
|
||||
state=ctx.state,
|
||||
@@ -340,7 +331,6 @@ async def policy_precheck_stage(
|
||||
)
|
||||
ctx.cost_gold = await deps.resolve_cost_gold(
|
||||
prep=ctx.prep,
|
||||
route_skip_checks=ctx.route_skip_checks,
|
||||
hook_recorder=ctx.hook_recorder,
|
||||
session=ctx.session,
|
||||
)
|
||||
@@ -356,7 +346,6 @@ async def legacy_hook_adapter_stage(
|
||||
prep=prep,
|
||||
session=ctx.session,
|
||||
event_cache=ctx.event_cache,
|
||||
route_skip_checks=ctx.route_skip_checks,
|
||||
lane_context=_lane_context(ctx),
|
||||
hook_recorder=_recorder(ctx),
|
||||
side_effect_commit=_side_effect_commit(ctx),
|
||||
@@ -425,35 +414,12 @@ async def decision_log_stage(
|
||||
ctx.auth_allowed,
|
||||
None if ctx.auth_allowed else ctx.decision_reason,
|
||||
)
|
||||
side_effect_state = commit.snapshot() if commit is not None else None
|
||||
shadow_effect = None
|
||||
shadow_reason = None
|
||||
if has_deferred_commit:
|
||||
shadow_effect = "defer"
|
||||
shadow_reason = "side_effect_pending:" + ",".join(
|
||||
commit.pending_kinds if commit is not None else ()
|
||||
)
|
||||
if ctx.entered_side_effect_lock and ctx.side_effect_lock is not None:
|
||||
try:
|
||||
ctx.side_effect_lock.release()
|
||||
except Exception:
|
||||
pass
|
||||
ctx.entered_side_effect_lock = False
|
||||
latency_ms = (time.time() - ctx.start_time) * 1000
|
||||
await deps.append_auth_decision_log(
|
||||
bot_id=ctx.event_context.bot_id,
|
||||
platform=ctx.event_context.platform,
|
||||
group_id=_entity(ctx).group_id,
|
||||
user_id=_entity(ctx).user_id,
|
||||
module=ctx.module,
|
||||
effect=ctx.decision_effect or "error",
|
||||
reason=ctx.decision_reason,
|
||||
shadow_effect=shadow_effect,
|
||||
shadow_reason=shadow_reason,
|
||||
side_effect_state=side_effect_state,
|
||||
latency_ms=latency_ms,
|
||||
overloaded=is_overloaded(),
|
||||
)
|
||||
|
||||
|
||||
def build_auth_pipeline(deps: AuthPipelineDependencies) -> AuthPipeline:
|
||||
|
||||
@@ -60,7 +60,6 @@ class PolicyResource:
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PolicyContext:
|
||||
snapshot: AuthSnapshot
|
||||
route_skip_checks: bool = False
|
||||
allow_sleep_bypass: bool = False
|
||||
allow_group_sleep_bypass: bool = False
|
||||
|
||||
@@ -101,8 +100,6 @@ class PolicyDecisionPoint:
|
||||
profile = resource.profile
|
||||
if profile.hidden:
|
||||
return PolicyDecision("allow", "hidden_plugin_skip_auth")
|
||||
if context.route_skip_checks:
|
||||
return PolicyDecision("allow", "route_miss_skip_checks")
|
||||
if snapshot.ban_state is True and not principal.is_superuser:
|
||||
return PolicyDecision("deny", "user_or_group_banned")
|
||||
if profile.superuser_only and not principal.is_superuser:
|
||||
|
||||
@@ -2,11 +2,13 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from zhenxun.services.cache.runtime_cache import (
|
||||
PluginLimitMemoryCache,
|
||||
from zhenxun.utils.enum import BlockType, PluginType
|
||||
|
||||
from .auth.data_provider import (
|
||||
DEFAULT_PERMISSION_DATA_PROVIDER,
|
||||
PermissionDataProvider,
|
||||
PluginLimitSnapshot,
|
||||
)
|
||||
from zhenxun.utils.enum import BlockType, PluginType
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
@@ -81,6 +83,7 @@ async def get_plugin_auth_profile(
|
||||
*,
|
||||
event_cache: dict | None = None,
|
||||
allow_cache_load: bool = True,
|
||||
provider: PermissionDataProvider = DEFAULT_PERMISSION_DATA_PROVIDER,
|
||||
) -> PluginAuthProfile:
|
||||
module = str(getattr(plugin, "module", "") or "")
|
||||
profile_cache: dict[str, PluginAuthProfile] = {}
|
||||
@@ -98,10 +101,10 @@ async def get_plugin_auth_profile(
|
||||
limits = limit_cache[module]
|
||||
limits_ready = True
|
||||
if limits is None:
|
||||
limits = PluginLimitMemoryCache.get_limits_if_ready(module)
|
||||
limits = provider.get_module_limits_if_ready(module)
|
||||
limits_ready = limits is not None
|
||||
if limits is None and allow_cache_load:
|
||||
limits = await PluginLimitMemoryCache.get_limits(module)
|
||||
limits = await provider.get_module_limits(module)
|
||||
limits_ready = True
|
||||
if limits is None:
|
||||
limits = []
|
||||
|
||||
@@ -1,60 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
from nonebot.matcher import Matcher
|
||||
|
||||
from zhenxun.utils.enum import PluginType
|
||||
|
||||
from .auth.context import EventContext, set_route_modules
|
||||
|
||||
RouteContextGetter = Callable[[str, dict | None], Awaitable[set[str]]]
|
||||
CommandMatcherChecker = Callable[[type[Matcher]], bool]
|
||||
AlconnaShortcutChecker = Callable[[type[Matcher]], bool]
|
||||
|
||||
|
||||
async def route_precheck(
|
||||
matcher: Matcher,
|
||||
context: EventContext,
|
||||
*,
|
||||
route_modules_with_commands: set[str],
|
||||
get_route_context: RouteContextGetter,
|
||||
is_command_matcher_class: CommandMatcherChecker,
|
||||
matcher_has_alconna_shortcuts: AlconnaShortcutChecker,
|
||||
) -> bool:
|
||||
"""Skip expensive auth checks for command matchers proven to be off-route."""
|
||||
|
||||
module = matcher.plugin_name or ""
|
||||
if not module:
|
||||
return False
|
||||
if _is_hidden_plugin(matcher):
|
||||
return False
|
||||
if not is_command_matcher_class(type(matcher)):
|
||||
return False
|
||||
|
||||
route_modules = context.route_modules if context.route_modules_loaded else None
|
||||
if route_modules is None:
|
||||
route_modules = await get_route_context(
|
||||
context.plain_text,
|
||||
context.event_cache,
|
||||
)
|
||||
set_route_modules(None, context, route_modules)
|
||||
|
||||
if module in route_modules_with_commands and module not in route_modules:
|
||||
if matcher_has_alconna_shortcuts(type(matcher)):
|
||||
return False
|
||||
if context.event_cache is not None:
|
||||
context.event_cache["route_skip"] = True
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _is_hidden_plugin(matcher: Matcher) -> bool:
|
||||
plugin = matcher.plugin
|
||||
if not plugin or not plugin.metadata:
|
||||
return False
|
||||
extra = plugin.metadata.extra or {}
|
||||
return extra.get("plugin_type") == PluginType.HIDDEN
|
||||
|
||||
|
||||
__all__ = ["route_precheck"]
|
||||
@@ -1,24 +1,120 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from dataclasses import dataclass, field
|
||||
import time
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from zhenxun.services.cache.runtime_cache import (
|
||||
BanMemoryCache,
|
||||
BotMemoryCache,
|
||||
BotSnapshot,
|
||||
GroupMemoryCache,
|
||||
GroupSnapshot,
|
||||
LevelUserMemoryCache,
|
||||
LevelUserSnapshot,
|
||||
)
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .auth.config import LOGGER_COMMAND
|
||||
from .auth.context import EventContext
|
||||
from .auth.data_provider import (
|
||||
DEFAULT_PERMISSION_DATA_PROVIDER,
|
||||
PermissionDataProvider,
|
||||
)
|
||||
from .auth_profile import PluginAuthProfile
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nonebot.adapters import Bot
|
||||
|
||||
QQ_CLIENT_GROUP_REPAIR_TTL = 60
|
||||
_QQ_CLIENT_GROUP_REPAIR_FAILURES: dict[tuple[str, str], float] = {}
|
||||
_QQ_CLIENT_GROUP_REPAIR_LOCKS: dict[tuple[str, str], asyncio.Lock] = {}
|
||||
|
||||
|
||||
def _build_runtime_group_snapshot(context: EventContext) -> GroupSnapshot | None:
|
||||
"""Provide a non-persistent default group for QQ official runtime auth."""
|
||||
if context.platform_scope != "qq_api" or not context.group_id:
|
||||
return None
|
||||
return GroupSnapshot(
|
||||
group_id=context.group_id,
|
||||
channel_id=context.channel_id,
|
||||
group_name="",
|
||||
max_member_count=0,
|
||||
member_count=0,
|
||||
status=True,
|
||||
level=5,
|
||||
is_super=False,
|
||||
group_flag=0,
|
||||
block_plugin="",
|
||||
superuser_block_plugin="",
|
||||
block_task="",
|
||||
superuser_block_task="",
|
||||
platform=context.platform,
|
||||
)
|
||||
|
||||
|
||||
def _qq_client_group_repair_key(context: EventContext) -> tuple[str, str] | None:
|
||||
if context.platform_scope != "qq_client" or not context.group_id:
|
||||
return None
|
||||
return (context.group_id, context.channel_id or "")
|
||||
|
||||
|
||||
def _qq_client_group_repair_on_cooldown(key: tuple[str, str]) -> bool:
|
||||
expire_at = _QQ_CLIENT_GROUP_REPAIR_FAILURES.get(key)
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
_QQ_CLIENT_GROUP_REPAIR_FAILURES.pop(key, None)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
async def _repair_missing_qq_client_group(
|
||||
context: EventContext,
|
||||
*,
|
||||
provider: PermissionDataProvider,
|
||||
) -> GroupSnapshot | None:
|
||||
"""Persist a minimal OneBot group when startup group sync returned empty."""
|
||||
key = _qq_client_group_repair_key(context)
|
||||
if key is None or not provider.group_cache_loaded():
|
||||
return None
|
||||
group_id, _ = key
|
||||
if _qq_client_group_repair_on_cooldown(key):
|
||||
return None
|
||||
|
||||
try:
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
|
||||
lock = _QQ_CLIENT_GROUP_REPAIR_LOCKS.setdefault(key, asyncio.Lock())
|
||||
async with lock:
|
||||
existing = provider.get_group_if_ready(
|
||||
group_id,
|
||||
context.channel_id,
|
||||
)
|
||||
if existing is not None:
|
||||
return existing
|
||||
defaults = {
|
||||
"group_name": "",
|
||||
"max_member_count": 0,
|
||||
"member_count": 0,
|
||||
"group_flag": 1,
|
||||
"platform": context.platform,
|
||||
}
|
||||
group, _ = await GroupConsole.get_or_create_root_group(
|
||||
group_id=group_id,
|
||||
defaults=defaults,
|
||||
)
|
||||
from zhenxun.services.cache.runtime_cache import GroupMemoryCache
|
||||
|
||||
await GroupMemoryCache.upsert_from_model(group)
|
||||
return GroupSnapshot.from_model(group)
|
||||
except Exception as exc:
|
||||
_QQ_CLIENT_GROUP_REPAIR_FAILURES[key] = time.time() + QQ_CLIENT_GROUP_REPAIR_TTL
|
||||
logger.warning(
|
||||
"协议端群记录缺失自愈失败,已短期跳过重复修复",
|
||||
LOGGER_COMMAND,
|
||||
group_id=context.group_id,
|
||||
e=exc,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class AuthSnapshot:
|
||||
@@ -72,6 +168,7 @@ async def build_auth_snapshot(
|
||||
bot: "Bot",
|
||||
skip_ban: bool = False,
|
||||
allow_cache_load: bool = False,
|
||||
provider: PermissionDataProvider = DEFAULT_PERMISSION_DATA_PROVIDER,
|
||||
) -> AuthSnapshot:
|
||||
event_cache = context.event_cache
|
||||
entity = context.entity
|
||||
@@ -85,15 +182,15 @@ async def build_auth_snapshot(
|
||||
):
|
||||
bot_data = event_cache.get("bot_data")
|
||||
else:
|
||||
bot_data = BotMemoryCache.get_if_ready(bot.self_id)
|
||||
bot_data = provider.get_bot_if_ready(bot.self_id)
|
||||
if bot_data is None:
|
||||
if allow_cache_load:
|
||||
bot_data = await BotMemoryCache.get(bot.self_id)
|
||||
elif not BotMemoryCache.is_loaded():
|
||||
bot_data = await provider.get_bot(bot.self_id)
|
||||
elif not provider.bot_cache_loaded():
|
||||
cache_misses.add("bot")
|
||||
if event_cache is not None:
|
||||
event_cache["bot_data"] = bot_data
|
||||
event_cache["bot_cache_ready"] = BotMemoryCache.is_loaded()
|
||||
event_cache["bot_cache_ready"] = provider.bot_cache_loaded()
|
||||
|
||||
group = None
|
||||
if entity.group_id:
|
||||
@@ -104,14 +201,31 @@ async def build_auth_snapshot(
|
||||
):
|
||||
group = event_cache.get("group")
|
||||
else:
|
||||
group = GroupMemoryCache.get_if_ready(entity.group_id, entity.channel_id)
|
||||
if group is None and not GroupMemoryCache.is_loaded():
|
||||
group = provider.get_group_if_ready(entity.group_id, entity.channel_id)
|
||||
if group is None and not provider.group_cache_loaded():
|
||||
cache_misses.add("group")
|
||||
elif group is None and allow_cache_load:
|
||||
group = await GroupMemoryCache.get(entity.group_id, entity.channel_id)
|
||||
group = await provider.get_group(entity.group_id, entity.channel_id)
|
||||
if event_cache is not None:
|
||||
event_cache["group"] = group
|
||||
event_cache["group_cache_ready"] = GroupMemoryCache.is_loaded()
|
||||
event_cache["group_cache_ready"] = provider.group_cache_loaded()
|
||||
if group is None:
|
||||
group = await _repair_missing_qq_client_group(
|
||||
context,
|
||||
provider=provider,
|
||||
)
|
||||
if group is None and (runtime_group := _build_runtime_group_snapshot(context)):
|
||||
group = runtime_group
|
||||
cache_misses.discard("group")
|
||||
if event_cache is not None:
|
||||
event_cache["group"] = group
|
||||
event_cache["group_cache_ready"] = True
|
||||
event_cache["group_runtime_virtual"] = True
|
||||
elif group is not None:
|
||||
cache_misses.discard("group")
|
||||
if event_cache is not None:
|
||||
event_cache["group"] = group
|
||||
event_cache["group_cache_ready"] = True
|
||||
|
||||
admin_levels = None
|
||||
if profile.need_admin:
|
||||
@@ -122,13 +236,13 @@ async def build_auth_snapshot(
|
||||
):
|
||||
admin_levels = event_cache.get("admin_levels")
|
||||
else:
|
||||
admin_levels = LevelUserMemoryCache.get_levels_if_ready(
|
||||
admin_levels = provider.get_admin_levels_if_ready(
|
||||
entity.user_id,
|
||||
entity.group_id,
|
||||
)
|
||||
if admin_levels is None:
|
||||
if allow_cache_load:
|
||||
admin_levels = await LevelUserMemoryCache.get_levels(
|
||||
admin_levels = await provider.get_admin_levels(
|
||||
entity.user_id,
|
||||
entity.group_id,
|
||||
)
|
||||
@@ -136,19 +250,19 @@ async def build_auth_snapshot(
|
||||
cache_misses.add("admin_levels")
|
||||
if event_cache is not None:
|
||||
event_cache["admin_levels"] = admin_levels
|
||||
event_cache["admin_cache_ready"] = LevelUserMemoryCache.is_loaded()
|
||||
event_cache["admin_cache_ready"] = provider.admin_cache_loaded()
|
||||
|
||||
ban_state = None
|
||||
if not skip_ban:
|
||||
if event_cache is not None and "ban_state" in event_cache:
|
||||
ban_state = event_cache.get("ban_state")
|
||||
elif BanMemoryCache.is_loaded():
|
||||
ban_state = BanMemoryCache.is_banned(entity.user_id, entity.group_id)
|
||||
elif provider.ban_cache_loaded():
|
||||
ban_state = provider.is_banned(entity.user_id, entity.group_id)
|
||||
if event_cache is not None:
|
||||
event_cache["ban_state"] = ban_state
|
||||
elif allow_cache_load:
|
||||
await BanMemoryCache.ensure_loaded()
|
||||
ban_state = BanMemoryCache.is_banned(entity.user_id, entity.group_id)
|
||||
await provider.ensure_ban_loaded()
|
||||
ban_state = provider.is_banned(entity.user_id, entity.group_id)
|
||||
if event_cache is not None:
|
||||
event_cache["ban_state"] = ban_state
|
||||
else:
|
||||
@@ -174,6 +288,7 @@ async def get_or_build_auth_snapshot(
|
||||
bot: "Bot",
|
||||
skip_ban: bool = False,
|
||||
allow_cache_load: bool = False,
|
||||
provider: PermissionDataProvider = DEFAULT_PERMISSION_DATA_PROVIDER,
|
||||
) -> AuthSnapshot:
|
||||
event_cache = context.event_cache
|
||||
module = profile.module
|
||||
@@ -190,6 +305,7 @@ async def get_or_build_auth_snapshot(
|
||||
bot=bot,
|
||||
skip_ban=skip_ban,
|
||||
allow_cache_load=allow_cache_load,
|
||||
provider=provider,
|
||||
)
|
||||
if event_cache is not None:
|
||||
event_cache.setdefault("auth_snapshots", {})[module] = snapshot
|
||||
|
||||
@@ -1,11 +1,9 @@
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
from nonebot.adapters import Bot, Message
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.bot_message_store import BotMessageStore
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import BotSentType
|
||||
from zhenxun.utils.log_sanitizer import sanitize_for_logging
|
||||
from zhenxun.utils.manager.message_manager import MessageManager
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
@@ -45,13 +43,15 @@ def replace_message(message: Message) -> str:
|
||||
async def handle_api_result(
|
||||
bot: Bot, exception: Exception | None, api: str, data: dict[str, Any], result: Any
|
||||
):
|
||||
if exception or api != "send_msg":
|
||||
if (
|
||||
exception
|
||||
or api != "send_msg"
|
||||
or PlatformUtils.get_platform_scope(bot) != "qq_client"
|
||||
):
|
||||
return
|
||||
user_id = data.get("user_id")
|
||||
group_id = data.get("group_id")
|
||||
message_id = result.get("message_id")
|
||||
message_id = result.get("message_id") if isinstance(result, Mapping) else None
|
||||
message: Message = data.get("message", "")
|
||||
message_type = data.get("message_type")
|
||||
try:
|
||||
if user_id and message_id:
|
||||
MessageManager.add(str(user_id), str(message_id))
|
||||
@@ -62,27 +62,5 @@ async def handle_api_result(
|
||||
logger.warning(
|
||||
f"收集消息id发生错误...data: {data}, result: {result}", LOG_COMMAND, e=e
|
||||
)
|
||||
if not Config.get_config("hook", "RECORD_BOT_SENT_MESSAGES"):
|
||||
return
|
||||
try:
|
||||
await BotMessageStore.append_buffered(
|
||||
bot_id=bot.self_id,
|
||||
user_id=user_id,
|
||||
group_id=group_id,
|
||||
sent_type=BotSentType.GROUP
|
||||
if message_type == "group"
|
||||
else BotSentType.PRIVATE,
|
||||
text=replace_message(message),
|
||||
plain_text=message.extract_plain_text()
|
||||
if isinstance(message, Message)
|
||||
else replace_message(message),
|
||||
platform=PlatformUtils.get_platform(bot),
|
||||
)
|
||||
sanitized_message = sanitize_for_logging(message, context="nonebot_message")
|
||||
logger.debug(f"消息发送记录,message: {sanitized_message}")
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"消息发送记录发生错误...data: {data}, result: {result}",
|
||||
LOG_COMMAND,
|
||||
e=e,
|
||||
)
|
||||
sanitized_message = sanitize_for_logging(message, context="nonebot_message")
|
||||
logger.debug(f"消息发送记录,message: {sanitized_message}")
|
||||
|
||||
Reference in New Issue
Block a user