bugfix:更换数据库初始化超时路径以修复连接超时问题 (#2137)

* bugfix:更换数据库初始化超时路径以修复连接超时问题

* 移除部分观测链路

* 细节修改

* 权限检查细节修改2

* bugfix:修复金币懒加载造成插件金币消耗不了的问题

* bugfix:整理鉴权逻辑

* 完善缓存系统

* 优化官端使用

* bugfix:增加官机和频道的前导统一剥离,新增获取群列表自愈

* bugfix:修复导入问题
This commit is contained in:
Copaan
2026-06-07 18:14:01 +08:00
committed by GitHub
parent 381d497c6d
commit 8afc8f8673
56 changed files with 1899 additions and 1654 deletions
@@ -109,7 +109,7 @@ class MemberUpdateManage:
members = await interface.get_members(SceneType.GROUP, group_scene.id)
try:
group_console, _ = await GroupConsole.get_or_create(
group_console, _ = await GroupConsole.get_or_create_root_group(
group_id=group_id, defaults={"platform": platform}
)
group_console.member_count = len(members)
@@ -39,6 +39,11 @@ _matcher = on_alconna(
async def _(bot: Bot, session: EventSession, arparma: Arparma):
logger.info("更新群组信息", arparma.header_result, session=session)
try:
if PlatformUtils.get_platform_scope(bot) != "qq_client":
await MessageUtils.build_message(
"当前平台不支持旧群组信息同步,仅 OneBot 协议端可用。"
).send(reply_to=True)
return
await PlatformUtils.update_group(bot)
await MessageUtils.build_message("已经成功更新了群组信息!").send(reply_to=True)
except Exception:
@@ -87,7 +87,7 @@ class PluginManager:
for gid in groups_to_open | groups_to_close:
platform = bot.adapter.get_name() if bot else "qq"
await GroupConsole.get_or_create(
await GroupConsole.get_or_create_root_group(
group_id=gid, defaults={"platform": platform}
)
@@ -173,10 +173,12 @@ class TaskStrategy(SwitchStrategy):
async def set_all_default_status(self, status: bool) -> None:
await TaskInfo.all().update(default_status=status)
# Bulk updates bypass model save hooks; keep TaskInfoMemoryCache in sync.
await self.refresh_cache()
async def set_all_global_status(self, status: bool) -> None:
await TaskInfo.all().update(status=status)
# Bulk updates bypass model save hooks; keep TaskInfoMemoryCache in sync.
await self.refresh_cache()
async def refresh_cache(self) -> None:
+3
View File
@@ -4,6 +4,7 @@ from nonebot.adapters import Bot
from zhenxun.configs.config import Config
from zhenxun.services.log import logger
from zhenxun.utils.platform import PlatformUtils
Config.add_plugin_config(
"catchphrase",
@@ -16,6 +17,8 @@ Config.add_plugin_config(
@Bot.on_calling_api
async def handle_api_call(bot: Bot, api: str, data: dict[str, Any]):
if PlatformUtils.get_platform_scope(bot) != "qq_client":
return
if api == "send_msg":
catchphrase = Config.get_config("catchphrase", "CATCHPHRASE")
if catchphrase and (message := data.get("message")):
-10
View File
@@ -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
+18 -2
View File
@@ -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",
]
+495 -174
View File
@@ -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:
+42 -268
View File
@@ -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,
+4 -38
View File
@@ -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"]
+135 -19
View File
@@ -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
+9 -31
View File
@@ -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}")
+8 -1
View File
@@ -30,7 +30,7 @@ async def _(bot: Bot):
参数:
bot: Bot
"""
if PlatformUtils.get_platform(bot) != "qq":
if PlatformUtils.get_platform_scope(bot) != "qq_client":
return
logger.debug(f"更新Bot: {bot.self_id} 的群认证...", "群认证同步")
@@ -44,6 +44,13 @@ async def _(bot: Bot):
)
return
if not current_group_list:
logger.warning(
f"Bot: {bot.self_id} 未获取到任何群组,"
"本次不会创建群认证;后续群消息将尝试按事件自愈。",
"群认证同步",
)
db_group_list: list[str] = await GroupConsole.all().values_list(
"group_id", flat=True
) # pyright: ignore[reportAssignmentType]
+7 -3
View File
@@ -8,7 +8,7 @@ from nonebot_plugin_apscheduler import scheduler
from zhenxun.configs.utils import PluginExtraData, Task
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.task_info import TaskInfo
from zhenxun.services.cache.runtime_cache import TaskInfoMemoryCache
from zhenxun.services.cache.runtime_cache import GroupMemoryCache, TaskInfoMemoryCache
from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
@@ -63,6 +63,8 @@ async def update_to_group(create_list: list[tuple[bool, TaskInfo]]):
)
group.block_task = CommonUtils.convert_module_format(block_tasks)
await GroupConsole.bulk_update(group_list, ["block_task"], 10)
for group in group_list:
await GroupMemoryCache.upsert_from_model(group)
async def to_db(
@@ -146,8 +148,10 @@ async def _():
for plugin in get_loaded_plugins():
await _handle_setting(plugin, task_info_list, task_list)
if not task_info_list:
await TaskInfo.all().update(load_status=False)
await TaskInfoMemoryCache.refresh()
logger.warning(
"未扫描到任何被动技能,跳过 TaskInfo.load_status 全量关闭,"
"避免插件加载异常时误关闭全部被动技能。",
)
return
module_dict = {t[1]: t[0] for t in await TaskInfo.all().values_list("id", "module")}
load_task = []
+9 -7
View File
@@ -2,6 +2,7 @@ from pathlib import Path
import nonebot
from zhenxun.configs.config import BotConfig
from zhenxun.services.log import logger
path = Path(__file__).parent
@@ -15,11 +16,12 @@ except ImportError:
logger.warning("未安装 onebot-adapter,无法加载QQ平台专用插件...")
try:
from nonebot.adapters.qq import ( # noqa: F401 # pyright: ignore [reportMissingImports]
Bot,
)
if BotConfig.qq_adapter_load:
try:
from nonebot.adapters.qq import ( # noqa: F401 # pyright: ignore [reportMissingImports]
Bot,
)
nonebot.load_plugins(str((path / "qq_api").resolve()))
except ImportError:
logger.warning("未安装 qq-adapter,无法加载QQ官平台专用插件...")
nonebot.load_plugins(str((path / "qq_api").resolve()))
except ImportError:
logger.warning("未安装 qq-adapter,无法加载QQ官平台专用插件...")
@@ -108,9 +108,7 @@ async def _(
):
if session.user.id == bot.self_id:
"""新成员为bot本身"""
group, _ = await GroupConsole.get_or_create(
group_id=str(event.group_id), channel_id__isnull=True
)
group, _ = await GroupConsole.get_or_create_root_group(str(event.group_id))
try:
await GroupManager.add_bot(
bot, str(event.operator_id), str(event.group_id), group
@@ -124,7 +124,7 @@ class GroupManager:
group_id=group_id,
)
return
await GroupConsole.update_or_create(
await GroupConsole.get_or_create_root_group(
group_id=group_info["group_id"],
defaults={
"group_name": group_info["group_name"],
@@ -134,6 +134,7 @@ class GroupManager:
"block_plugin": block_plugin,
"platform": "qq",
},
update_defaults=True,
)
@classmethod
@@ -1,14 +1,13 @@
"""QQ official platform observer.
Official QQ identifiers are not in the same namespace as OneBot QQ numbers.
This observer intentionally avoids writing legacy identity tables; runtime auth
uses a non-persistent group snapshot when needed.
"""
from nonebot import on_message
from nonebot_plugin_uninfo import Uninfo
from zhenxun.models.friend_user import FriendUser
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.group_member_info import GroupInfoUser
from zhenxun.services.hot_query_cache import (
invalidate_group_members,
invalidate_member_names,
)
from zhenxun.services.log import logger
from zhenxun.utils.platform import PlatformUtils
@@ -20,21 +19,5 @@ _matcher = on_message(priority=999, block=False, rule=rule)
@_matcher.handle()
async def _(session: Uninfo):
platform = PlatformUtils.get_platform(session)
if session.group:
if not await GroupConsole.exists(group_id=session.group.id):
await GroupConsole.create(group_id=session.group.id)
logger.info("添加当前群组ID信息", session=session)
await GroupInfoUser.update_or_create(
user_id=session.user.id,
group_id=session.group.id,
platform=PlatformUtils.get_platform(session),
)
await invalidate_group_members(session.group.id, [session.user.id])
await invalidate_member_names([session.user.id])
elif not await FriendUser.exists(user_id=session.user.id, platform=platform):
await FriendUser.create(
user_id=session.user.id, platform=PlatformUtils.get_platform(session)
)
logger.info("添加当前好友用户信息", "", session=session)
async def _():
return
+2 -1
View File
@@ -203,7 +203,7 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio
session=event.user_id,
target=event.group_id,
)
group, _ = await GroupConsole.update_or_create(
group, _ = await GroupConsole.get_or_create_root_group(
group_id=str(event.group_id),
defaults={
"group_name": "",
@@ -211,6 +211,7 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio
"member_count": 0,
"group_flag": 1,
},
update_defaults=True,
)
await bot.set_group_add_request(
flag=event.flag, sub_type="invite", approve=True
@@ -18,6 +18,8 @@ async def _():
return
bots = nonebot.get_bots()
for bot in bots.values():
if PlatformUtils.get_platform_scope(bot) != "qq_client":
continue
try:
await PlatformUtils.update_group(bot)
except Exception as e:
@@ -36,6 +38,8 @@ async def _():
return
bots = nonebot.get_bots()
for bot in bots.values():
if PlatformUtils.get_platform_scope(bot) != "qq_client":
continue
try:
await PlatformUtils.update_friend(bot)
except Exception as e:
@@ -52,8 +52,8 @@ async def _():
if last_message:
now = datetime.now(pytz.timezone("Asia/Shanghai"))
if now - timedelta(days=2) > last_message.create_time:
_group, _ = await GroupConsole.get_or_create(
group_id=group.group_id, channel_id__isnull=True
_group, _ = await GroupConsole.get_or_create_root_group(
group.group_id
)
modules = [f"<{module}" for module in modules]
_group.block_task = ",".join(modules) + "," # type: ignore
@@ -147,7 +147,7 @@ def CheckGroupId():
@_matcher.assign("modify-level", parameterless=[CheckGroupId()])
async def _(session: EventSession, arparma: Arparma, state: T_State, level: int):
gid = state["group_id"]
group, _ = await GroupConsole.get_or_create(group_id=gid)
group, _ = await GroupConsole.get_or_create_root_group(gid)
old_level = group.level
group.level = level
await group.save(update_fields=["level"])
@@ -176,10 +176,10 @@ async def _(session: EventSession, arparma: Arparma, state: T_State):
@_matcher.assign("auth-handle", parameterless=[CheckGroupId()])
async def _(session: EventSession, arparma: Arparma, state: T_State):
gid = state["group_id"]
await GroupConsole.update_or_create(
await GroupConsole.get_or_create_root_group(
group_id=gid,
channel_id__isnull=True,
defaults={"group_flag": 0 if arparma.find("delete") else 1},
update_defaults=True,
)
s = "删除" if arparma.find("delete") else "添加"
await MessageUtils.build_message(f"{s}群认证成功!").send(reply_to=True)
@@ -54,6 +54,11 @@ async def _(
arparma: Arparma,
):
try:
if PlatformUtils.get_platform_scope(bot) != "qq_client":
await MessageUtils.build_message(
"当前平台不支持旧群组信息同步,仅 OneBot 协议端可用。"
).send()
return
num = await PlatformUtils.update_group(bot)
logger.info(
f"更新群聊信息完成,共更新了 {num} 个群组的信息!",
@@ -75,6 +80,11 @@ async def _(
arparma: Arparma,
):
try:
if PlatformUtils.get_platform_scope(bot) != "qq_client":
await MessageUtils.build_message(
"当前平台不支持旧好友信息同步,仅 OneBot 协议端可用。"
).send()
return
num = await PlatformUtils.update_friend(bot)
logger.info(
f"更新好友信息完成,共更新了 {num} 个好友的信息!",
@@ -213,9 +213,10 @@ async def _(param: HandleRequest) -> Result:
group.group_flag = 1
await group.save(update_fields=["group_flag"])
else:
await GroupConsole.update_or_create(
await GroupConsole.get_or_create_root_group(
group_id=req.group_id,
defaults={"group_flag": 1},
update_defaults=True,
)
try:
await FgRequest.approve(bot, param.id)
@@ -240,8 +241,7 @@ async def _(param: HandleRequest) -> Result:
async def _(param: LeaveGroup) -> Result:
try:
bot = nonebot.get_bot(param.bot_id)
platform = PlatformUtils.get_platform(bot)
if platform != "qq":
if PlatformUtils.get_platform_scope(bot) != "qq_client":
return Result.warning_("该平台不支持退群操作...")
group_list, _ = await PlatformUtils.get_group_list(bot)
if param.group_id not in [g.group_id for g in group_list]:
@@ -265,8 +265,7 @@ async def _(param: LeaveGroup) -> Result:
async def _(param: DeleteFriend) -> Result:
try:
bot = nonebot.get_bot(param.bot_id)
platform = PlatformUtils.get_platform(bot)
if platform != "qq":
if PlatformUtils.get_platform_scope(bot) != "qq_client":
return Result.warning_("该平台不支持删除好友操作...")
friend_list, _ = await PlatformUtils.get_friend_list(bot)
if param.user_id not in [f.user_id for f in friend_list]:
+1 -1
View File
@@ -40,7 +40,7 @@ def reply_check() -> Rule:
if event.get_type() == "message":
return (
bool(await reply_fetch(event, bot))
and PlatformUtils.get_platform(session) == "qq"
and PlatformUtils.get_platform_scope(session) == "qq_client"
)
return False