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
-113
View File
@@ -1,113 +0,0 @@
from __future__ import annotations
import asyncio
from collections import deque
import contextlib
import time
from typing import TYPE_CHECKING
from zhenxun.services.log import logger
if TYPE_CHECKING:
from .bot_message_store import BotMessageStore
LOG_COMMAND = "BotMessageStore"
_BUFFER_MAX_RETAIN = 20_000
_FLUSH_TRIGGER_SIZE = 64
_FLUSH_BATCH_SIZE = 500
_FLUSH_INTERVAL_SECONDS = 5.0
_DROP_LOG_INTERVAL_SECONDS = 10.0
_buffer: deque[BotMessageStore] = deque()
_buffer_lock = asyncio.Lock()
_flush_lock = asyncio.Lock()
_flush_task: asyncio.Task[None] | None = None
_dropped = 0
_last_drop_log_at = 0.0
def _ensure_flush_task() -> None:
global _flush_task
if _flush_task is not None and not _flush_task.done():
return
try:
loop = asyncio.get_running_loop()
except RuntimeError:
return
_flush_task = loop.create_task(_flush_loop())
def _record_drop() -> None:
global _dropped, _last_drop_log_at
_dropped += 1
now = time.monotonic()
if now - _last_drop_log_at < _DROP_LOG_INTERVAL_SECONDS:
return
_last_drop_log_at = now
logger.warning(
f"bot_message_store buffer full, dropped {_dropped} records, "
f"backlog={len(_buffer)}",
LOG_COMMAND,
)
async def _flush_loop() -> None:
while True:
await asyncio.sleep(_FLUSH_INTERVAL_SECONDS)
try:
await flush_bot_message_store_buffer("定时")
except asyncio.CancelledError:
raise
except Exception as exc:
logger.warning("定时批量写入 Bot 发送记录失败", LOG_COMMAND, e=exc)
async def append_bot_message_store_record(record: BotMessageStore) -> None:
_ensure_flush_task()
async with _buffer_lock:
if len(_buffer) >= _BUFFER_MAX_RETAIN:
_buffer.popleft()
_record_drop()
_buffer.append(record)
should_flush = len(_buffer) >= _FLUSH_TRIGGER_SIZE and not _flush_lock.locked()
if should_flush:
await flush_bot_message_store_buffer("缓冲区触发")
async def flush_bot_message_store_buffer(reason: str) -> int:
from .bot_message_store import BotMessageStore
async with _flush_lock:
written = 0
while True:
batch: list[BotMessageStore] = []
async with _buffer_lock:
while _buffer and len(batch) < _FLUSH_BATCH_SIZE:
batch.append(_buffer.popleft())
if not batch:
break
try:
await BotMessageStore.bulk_create(batch, batch_size=_FLUSH_BATCH_SIZE)
except Exception as exc:
async with _buffer_lock:
retain_count = max(_BUFFER_MAX_RETAIN - len(_buffer), 0)
for record in reversed(batch[-retain_count:]):
_buffer.appendleft(record)
logger.error(f"{reason}批量写入 Bot 发送记录失败", LOG_COMMAND, e=exc)
return written
written += len(batch)
if written:
logger.debug(f"{reason}批量写入 Bot 发送记录 {written} 条", LOG_COMMAND)
return written
async def stop_bot_message_store_buffer() -> int:
global _flush_task
task = _flush_task
_flush_task = None
if task is not None:
task.cancel()
with contextlib.suppress(BaseException):
await task
return await flush_bot_message_store_buffer("关闭")
-49
View File
@@ -1,49 +0,0 @@
from typing import ClassVar
from tortoise import fields
from zhenxun.services.db_context import Model
class AuthDecisionLog(Model):
id = fields.IntField(pk=True, generated=True, auto_increment=True)
"""自增id"""
bot_id = fields.CharField(255, null=True, description="Bot ID")
"""Bot ID"""
platform = fields.CharField(64, null=True, description="平台")
"""平台"""
group_id = fields.CharField(255, null=True, description="群组id")
"""群组id"""
user_id = fields.CharField(255, null=True, description="用户id")
"""用户id"""
module = fields.CharField(255, null=True, description="插件模块")
"""插件模块"""
effect = fields.CharField(32, description="决策结果")
"""决策结果 allow/deny/skip/defer/error"""
reason = fields.CharField(255, null=True, description="原因")
"""原因"""
shadow_effect = fields.CharField(32, null=True, description="影子决策结果")
"""影子决策结果"""
shadow_reason = fields.CharField(255, null=True, description="影子决策原因")
"""影子决策原因"""
side_effect_state = fields.TextField(null=True, description="副作用状态")
"""副作用状态 JSON 摘要"""
latency_ms = fields.FloatField(default=0, description="耗时毫秒")
"""耗时毫秒"""
overloaded = fields.BooleanField(default=False, description="是否过载")
"""是否过载"""
create_time = fields.DatetimeField(auto_now_add=True, description="创建时间")
"""创建时间"""
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "auth_decision_log"
table_description = "权限决策追加审计日志"
indexes: ClassVar = [
("create_time",),
("module", "create_time"),
("effect", "create_time"),
]
@classmethod
async def _run_script(cls):
return []
+6 -2
View File
@@ -3,6 +3,7 @@ from typing import ClassVar
from typing_extensions import Self
from tortoise import fields
from tortoise.expressions import Q
from zhenxun.services.cache import CacheException, CacheRegistry, CacheRoot
from zhenxun.services.cache.runtime_cache import BanMemoryCache
@@ -89,15 +90,18 @@ class BanConsole(Model):
cls._ensure_cache_registered()
if not user_id and not group_id:
raise UserAndGroupIsNone()
dao = DataAccess(cls)
if user_id:
dao = DataAccess(cls)
return (
await dao.safe_get_or_none(user_id=user_id, group_id=group_id)
if group_id
else await dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
)
else:
return await dao.safe_get_or_none(user_id="", group_id=group_id)
return await cls.safe_get_or_none(
Q(user_id__isnull=True) | Q(user_id=""),
group_id=group_id,
)
@classmethod
async def check_ban_level(
+6 -2
View File
@@ -26,7 +26,7 @@ class BotConsole(Model):
available_plugins = fields.TextField(default="", description="可用插件")
"""可用插件"""
available_tasks = fields.TextField(default="", description="可用被动技能")
"""可用被动技能"""
"""可用被动技能管理镜像,不作为运行白名单。"""
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "bot_console"
@@ -87,7 +87,11 @@ class BotConsole(Model):
@classmethod
async def get_tasks(cls, bot_id: str | None = None, status: bool | None = True):
"""
获取bot被动技能
获取bot被动技能管理镜像。
available_tasks 只服务管理命令和展示,不参与运行白名单判断。
被动运行真源是 TaskInfo.status/load_status、BotConsole.block_tasks、
GroupConsole.block_task/superuser_block_task。
参数:
bot_id (str | None, optional): bot_id. Defaults to None.
-55
View File
@@ -1,55 +0,0 @@
from tortoise import fields
from zhenxun.services.db_context import Model
from zhenxun.utils.enum import BotSentType
from ._bot_message_buffer import append_bot_message_store_record
class BotMessageStore(Model):
id = fields.IntField(pk=True, generated=True, auto_increment=True)
"""自增id"""
bot_id = fields.CharField(255, null=True)
"""bot id"""
user_id = fields.CharField(255, null=True)
"""目标id"""
group_id = fields.CharField(255, null=True)
"""群组id"""
sent_type = fields.CharEnumField(BotSentType)
"""类型"""
text = fields.TextField(null=True)
"""文本内容"""
plain_text = fields.TextField(null=True)
"""纯文本"""
platform = fields.CharField(255, null=True)
"""平台"""
create_time = fields.DatetimeField(auto_now_add=True)
"""创建时间"""
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "bot_message_store"
table_description = "Bot发送消息列表"
@classmethod
async def append_buffered(
cls,
*,
bot_id: str | None = None,
user_id: str | None = None,
group_id: str | None = None,
sent_type: BotSentType,
text: str | None = None,
plain_text: str | None = None,
platform: str | None = None,
) -> None:
await append_bot_message_store_record(
cls(
bot_id=bot_id,
user_id=user_id,
group_id=group_id,
sent_type=sent_type,
text=text,
plain_text=plain_text,
platform=platform,
)
)
+4 -2
View File
@@ -151,8 +151,10 @@ class FgRequest(Model):
"添加好友自动发送BOT自我介绍图片", session=req.user_id
)
else:
await GroupConsole.update_or_create(
group_id=req.group_id, defaults={"group_flag": 1}
await GroupConsole.get_or_create_root_group(
group_id=req.group_id,
defaults={"group_flag": 1},
update_defaults=True,
)
if req.flag == "0":
# 用户手动申请入群,创建群认证后提醒用户拉群
+161 -9
View File
@@ -1,3 +1,4 @@
import asyncio
from typing import TYPE_CHECKING, Any, ClassVar, cast, overload
from typing_extensions import Self
@@ -103,6 +104,8 @@ class GroupConsole(Model):
"""缓存键字段"""
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT]
"""开启锁"""
_root_group_locks: ClassVar[dict[str, asyncio.Lock]] = {}
"""普通群记录应用层锁,规避 channel_id=NULL 唯一键语义差异。"""
@classmethod
async def _get_task_modules(cls, *, default_status: bool) -> list[str]:
@@ -250,6 +253,143 @@ class GroupConsole(Model):
return group, is_create
@classmethod
def _clean_root_group_defaults(cls, defaults: dict | None) -> dict[str, Any]:
cleaned = {}
for field, value in (defaults or {}).items():
if field in {"id", "group_id", "channel_id", "channel_id__isnull"}:
continue
if value is None:
continue
cleaned[field] = value
return cleaned
@classmethod
def _root_group_score(cls, group: Self) -> tuple[int, int]:
score = 0
score += 8 if group.group_name else 0
score += 4 if group.max_member_count else 0
score += 4 if group.member_count else 0
score += 4 if group.group_flag else 0
score += 4 if group.is_super else 0
score += 3 if not group.status else 0
score += 3 if group.level != 5 else 0
score += len(convert_module_format(group.block_plugin))
score += len(convert_module_format(group.superuser_block_plugin))
score += len(convert_module_format(group.block_task))
score += len(convert_module_format(group.superuser_block_task))
return score, int(group.id or 0)
@classmethod
def _merge_module_field(cls, groups: list[Self], field: str) -> str:
modules: list[str] = []
seen = set()
for group in groups:
value = getattr(group, field, "") or ""
for module in cast(list[str], convert_module_format(value)):
if module not in seen:
seen.add(module)
modules.append(module)
return cast(str, convert_module_format(modules))
@classmethod
async def _deduplicate_root_group_records(cls, groups: list[Self]) -> Self:
if len(groups) == 1:
return groups[0]
keep = max(groups, key=cls._root_group_score)
newest_first = sorted(
groups, key=lambda group: int(group.id or 0), reverse=True
)
merged = {
"group_name": next(
(g.group_name for g in newest_first if g.group_name), ""
),
"max_member_count": max(g.max_member_count for g in groups),
"member_count": max(g.member_count for g in groups),
"status": all(g.status for g in groups),
"level": min(g.level for g in groups),
"is_super": any(g.is_super for g in groups),
"group_flag": max(g.group_flag for g in groups),
"block_plugin": cls._merge_module_field(groups, "block_plugin"),
"superuser_block_plugin": cls._merge_module_field(
groups, "superuser_block_plugin"
),
"block_task": cls._merge_module_field(groups, "block_task"),
"superuser_block_task": cls._merge_module_field(
groups, "superuser_block_task"
),
"platform": next((g.platform for g in newest_first if g.platform), "qq"),
}
update_fields = []
for field, value in merged.items():
if getattr(keep, field) != value:
setattr(keep, field, value)
update_fields.append(field)
if update_fields:
await keep.save(update_fields=update_fields)
for group in groups:
if group.id != keep.id:
await group.delete()
await GroupMemoryCache.upsert_from_model(keep)
return keep
@classmethod
async def get_or_create_root_group(
cls,
group_id: str | int,
defaults: dict | None = None,
*,
update_defaults: bool = False,
) -> tuple[Self, bool]:
"""获取或创建普通群记录,并收敛 channel_id=NULL 重复数据。
普通群固定使用 ``channel_id IS NULL``;频道记录必须继续显式传
``channel_id`` 走原有 get_or_create/update_or_create。
"""
gid = str(group_id).strip()
if not gid:
raise ValueError("group_id cannot be empty")
lock = cls._root_group_locks.setdefault(gid, asyncio.Lock())
async with lock:
defaults = cls._clean_root_group_defaults(defaults)
records = await cls.filter(group_id=gid, channel_id__isnull=True).all()
if records:
group = await cls._deduplicate_root_group_records(records)
if update_defaults:
update_fields = []
for field, value in defaults.items():
if hasattr(group, field) and getattr(group, field) != value:
setattr(group, field, value)
update_fields.append(field)
if update_fields:
await group.save(update_fields=update_fields)
await GroupMemoryCache.upsert_from_model(group)
return group, False
group = await cls.create(group_id=gid, channel_id=None, **defaults)
return group, True
@classmethod
async def _get_or_create_group_for_write(
cls,
group_id: str,
channel_id: str | None,
defaults: dict | None = None,
) -> tuple[Self, bool]:
defaults = cls._clean_root_group_defaults(defaults)
if channel_id:
return await cls.get_or_create(
group_id=group_id,
channel_id=channel_id,
defaults=defaults,
)
return await cls.get_or_create_root_group(group_id, defaults=defaults)
async def save(self, *args, **kwargs):
await super().save(*args, **kwargs)
await GroupMemoryCache.upsert_from_model(self)
@@ -326,6 +466,7 @@ class GroupConsole(Model):
module: str,
is_superuser: bool = False,
platform: str | None = None,
channel_id: str | None = None,
):
"""禁用群组插件
@@ -335,8 +476,10 @@ class GroupConsole(Model):
is_superuser: 是否为超级用户
platform: 平台
"""
group, _ = await cls.get_or_create(
group_id=group_id, defaults={"platform": platform}
group, _ = await cls._get_or_create_group_for_write(
group_id=group_id,
channel_id=channel_id,
defaults={"platform": platform},
)
update_fields = []
if is_superuser:
@@ -365,6 +508,7 @@ class GroupConsole(Model):
module: str,
is_superuser: bool = False,
platform: str | None = None,
channel_id: str | None = None,
):
"""禁用群组插件
@@ -374,8 +518,10 @@ class GroupConsole(Model):
is_superuser: 是否为超级用户
platform: 平台
"""
group, _ = await cls.get_or_create(
group_id=group_id, defaults={"platform": platform}
group, _ = await cls._get_or_create_group_for_write(
group_id=group_id,
channel_id=channel_id,
defaults={"platform": platform},
)
update_fields = []
if is_superuser:
@@ -447,6 +593,7 @@ class GroupConsole(Model):
task: str,
is_superuser: bool = False,
platform: str | None = None,
channel_id: str | None = None,
):
"""禁用群组插件
@@ -456,13 +603,15 @@ class GroupConsole(Model):
is_superuser: 是否为超级用户
platform: 平台
"""
group, _ = await cls.get_or_create(
group_id=group_id, defaults={"platform": platform}
group, _ = await cls._get_or_create_group_for_write(
group_id=group_id,
channel_id=channel_id,
defaults={"platform": platform},
)
update_fields = []
if is_superuser:
superuser_block_task = convert_module_format(group.superuser_block_task)
if task not in group.superuser_block_task:
if task not in superuser_block_task:
superuser_block_task.append(task)
group.superuser_block_task = convert_module_format(superuser_block_task)
update_fields.append("superuser_block_task")
@@ -484,6 +633,7 @@ class GroupConsole(Model):
task: str,
is_superuser: bool = False,
platform: str | None = None,
channel_id: str | None = None,
):
"""禁用群组插件
@@ -493,8 +643,10 @@ class GroupConsole(Model):
is_superuser: 是否为超级用户
platform: 平台
"""
group, _ = await cls.get_or_create(
group_id=group_id, defaults={"platform": platform}
group, _ = await cls._get_or_create_group_for_write(
group_id=group_id,
channel_id=channel_id,
defaults={"platform": platform},
)
update_fields = []
if is_superuser:
@@ -1,39 +0,0 @@
from typing import ClassVar
from tortoise import fields
from zhenxun.services.db_context import Model
class RuntimeBackpressureLog(Model):
id = fields.IntField(pk=True, generated=True, auto_increment=True)
"""自增id"""
scope_key = fields.CharField(255, null=True, description="作用域")
"""作用域"""
reason = fields.CharField(255, null=True, description="原因")
"""原因"""
lane = fields.CharField(64, null=True, description="调度通道")
"""调度通道"""
action = fields.CharField(64, description="处理动作")
"""处理动作 execute/skip/defer/signal"""
queue_size = fields.IntField(default=0, description="队列长度")
"""队列长度"""
active_count = fields.IntField(default=0, description="活跃数量")
"""活跃数量"""
duration_ms = fields.FloatField(default=0, description="持续耗时毫秒")
"""持续耗时毫秒"""
create_time = fields.DatetimeField(auto_now_add=True, description="创建时间")
"""创建时间"""
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "runtime_backpressure_log"
table_description = "运行时背压追加审计日志"
indexes: ClassVar = [
("create_time",),
("scope_key", "create_time"),
("lane", "create_time"),
]
@classmethod
async def _run_script(cls):
return []
+1
View File
@@ -219,6 +219,7 @@ class UserConsole(Model):
)
if not updated:
raise InsufficientGold()
await cls.invalidate_user_cache(user_id)
return GoldReservation(
user_id=user_id,
gold=gold,
-547
View File
@@ -1,547 +0,0 @@
from __future__ import annotations
import asyncio
from collections import deque
import contextlib
from dataclasses import dataclass
from datetime import datetime, timedelta
import json
import random
import time
from typing import Any, TypeVar
from tortoise import Tortoise
from zhenxun.builtin_plugins.hooks.auth_runtime_config import (
AUTH_OBSERVABILITY_RUNTIME_CONFIG,
)
from zhenxun.models.auth_decision_log import AuthDecisionLog
from zhenxun.models.runtime_backpressure_log import RuntimeBackpressureLog
from zhenxun.services.log import logger
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
LOG_COMMAND = "AuthObservability"
_BUFFER_MAX_RETAIN = AUTH_OBSERVABILITY_RUNTIME_CONFIG.buffer_max_retain
_FLUSH_TRIGGER_SIZE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.flush_trigger_size
_FLUSH_BATCH_SIZE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.flush_batch_size
_FLUSH_INTERVAL_SECONDS = AUTH_OBSERVABILITY_RUNTIME_CONFIG.flush_interval_seconds
_DROP_LOG_INTERVAL_SECONDS = AUTH_OBSERVABILITY_RUNTIME_CONFIG.drop_log_interval_seconds
_ALLOW_SAMPLE_RATE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.allow_sample_rate
_OVERLOADED_ALLOW_SAMPLE_RATE = (
AUTH_OBSERVABILITY_RUNTIME_CONFIG.overloaded_allow_sample_rate
)
_NON_ALLOW_SAMPLE_RATE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.non_allow_sample_rate
_BACKPRESSURE_SAMPLE_RATE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.backpressure_sample_rate
_BACKPRESSURE_SEVERE_ACTIVE_THRESHOLD = (
AUTH_OBSERVABILITY_RUNTIME_CONFIG.backpressure_severe_active_threshold
)
@dataclass(slots=True)
class AuthDecisionLogRecord:
bot_id: str | None
platform: str | None
group_id: str | None
user_id: str | None
module: str | None
effect: str
reason: str | None = None
shadow_effect: str | None = None
shadow_reason: str | None = None
side_effect_state: dict[str, Any] | None = None
latency_ms: float = 0.0
overloaded: bool = False
def to_model(self) -> AuthDecisionLog:
return AuthDecisionLog(
bot_id=self.bot_id,
platform=self.platform,
group_id=self.group_id,
user_id=self.user_id,
module=self.module,
effect=self.effect,
reason=self.reason,
shadow_effect=self.shadow_effect,
shadow_reason=self.shadow_reason,
side_effect_state=json.dumps(
self.side_effect_state,
ensure_ascii=False,
separators=(",", ":"),
)[:4000]
if self.side_effect_state
else None,
latency_ms=self.latency_ms,
overloaded=self.overloaded,
)
@dataclass(slots=True)
class RuntimeBackpressureLogRecord:
scope_key: str | None
reason: str | None
lane: str | None
action: str
queue_size: int = 0
active_count: int = 0
duration_ms: float = 0.0
def to_model(self) -> RuntimeBackpressureLog:
return RuntimeBackpressureLog(
scope_key=self.scope_key,
reason=self.reason,
lane=self.lane,
action=self.action,
queue_size=self.queue_size,
active_count=self.active_count,
duration_ms=self.duration_ms,
)
_auth_decision_buffer: deque[AuthDecisionLogRecord] = deque()
_backpressure_buffer: deque[RuntimeBackpressureLogRecord] = deque()
_buffer_lock = asyncio.Lock()
_flush_lock = asyncio.Lock()
_flush_task: asyncio.Task[None] | None = None
_dropped = 0
_last_drop_log_at = 0.0
_last_schema_repair_at = 0.0
_SCHEMA_REPAIR_INTERVAL_SECONDS = 300.0
T = TypeVar("T")
def _ensure_flush_task() -> None:
global _flush_task
if _flush_task is not None and not _flush_task.done():
return
_flush_task = asyncio.create_task(_flush_loop())
def _record_drop() -> None:
global _dropped, _last_drop_log_at
_dropped += 1
now = time.monotonic()
if now - _last_drop_log_at < _DROP_LOG_INTERVAL_SECONDS:
return
_last_drop_log_at = now
logger.warning(
"auth observability buffer full, dropped "
f"{_dropped} records, auth_backlog={len(_auth_decision_buffer)}, "
f"backpressure_backlog={len(_backpressure_buffer)}",
LOG_COMMAND,
)
def _sample(rate: float) -> bool:
if rate >= 1:
return True
if rate <= 0:
return False
return random.random() < rate
def _auth_decision_sample_rate(effect: str, overloaded: bool) -> float:
if effect != "allow":
return _NON_ALLOW_SAMPLE_RATE
if overloaded:
return _OVERLOADED_ALLOW_SAMPLE_RATE
return _ALLOW_SAMPLE_RATE
def _backpressure_sample_rate(record: RuntimeBackpressureLogRecord) -> float:
if record.reason and record.reason.startswith("hooks_"):
return 1.0
if record.active_count >= _BACKPRESSURE_SEVERE_ACTIVE_THRESHOLD:
return 1.0
if record.action in {"skip", "defer"}:
return _BACKPRESSURE_SAMPLE_RATE
return min(_BACKPRESSURE_SAMPLE_RATE, 0.02)
async def _append_auth_decision_record(record: AuthDecisionLogRecord) -> None:
_ensure_flush_task()
async with _buffer_lock:
total = len(_auth_decision_buffer) + len(_backpressure_buffer)
if total >= _BUFFER_MAX_RETAIN:
if len(_auth_decision_buffer) >= len(_backpressure_buffer):
with contextlib.suppress(IndexError):
_auth_decision_buffer.popleft()
else:
with contextlib.suppress(IndexError):
_backpressure_buffer.popleft()
_record_drop()
_auth_decision_buffer.append(record)
should_flush = (
len(_auth_decision_buffer) + len(_backpressure_buffer)
>= _FLUSH_TRIGGER_SIZE
and not _flush_lock.locked()
)
if should_flush:
# Fire-and-forget keeps auth hot path independent of database stalls.
asyncio.create_task(flush_auth_observability_buffer("缓冲区触发")) # noqa: RUF006
async def _append_backpressure_record(record: RuntimeBackpressureLogRecord) -> None:
_ensure_flush_task()
async with _buffer_lock:
total = len(_auth_decision_buffer) + len(_backpressure_buffer)
if total >= _BUFFER_MAX_RETAIN:
if len(_auth_decision_buffer) >= len(_backpressure_buffer):
with contextlib.suppress(IndexError):
_auth_decision_buffer.popleft()
else:
with contextlib.suppress(IndexError):
_backpressure_buffer.popleft()
_record_drop()
_backpressure_buffer.append(record)
should_flush = (
len(_auth_decision_buffer) + len(_backpressure_buffer)
>= _FLUSH_TRIGGER_SIZE
and not _flush_lock.locked()
)
if should_flush:
# Fire-and-forget keeps auth hot path independent of database stalls.
asyncio.create_task(flush_auth_observability_buffer("缓冲区触发")) # noqa: RUF006
async def append_auth_decision_log(
*,
bot_id: str | None,
platform: str | None,
group_id: str | None,
user_id: str | None,
module: str | None,
effect: str,
reason: str | None = None,
shadow_effect: str | None = None,
shadow_reason: str | None = None,
side_effect_state: dict[str, Any] | None = None,
latency_ms: float = 0.0,
overloaded: bool = False,
) -> None:
if shadow_effect is None and not _sample(
_auth_decision_sample_rate(effect, overloaded)
):
return
record = AuthDecisionLogRecord(
bot_id=bot_id,
platform=platform,
group_id=group_id,
user_id=user_id,
module=module,
effect=effect,
reason=(reason or "")[:255] or None,
shadow_effect=(shadow_effect or "")[:32] or None,
shadow_reason=(shadow_reason or "")[:255] or None,
side_effect_state=side_effect_state,
latency_ms=latency_ms,
overloaded=overloaded,
)
await _append_auth_decision_record(record)
async def append_runtime_backpressure_log(
*,
scope_key: str | None,
reason: str | None,
lane: str | None,
action: str,
queue_size: int = 0,
active_count: int = 0,
duration_ms: float = 0.0,
) -> None:
record = RuntimeBackpressureLogRecord(
scope_key=(scope_key or "")[:255] or None,
reason=(reason or "")[:255] or None,
lane=(lane or "")[:64] or None,
action=action,
queue_size=queue_size,
active_count=active_count,
duration_ms=duration_ms,
)
if not _sample(_backpressure_sample_rate(record)):
return
await _append_backpressure_record(record)
async def _flush_loop() -> None:
while True:
await asyncio.sleep(_FLUSH_INTERVAL_SECONDS)
try:
await flush_auth_observability_buffer("定时")
except asyncio.CancelledError:
raise
except Exception as exc:
logger.warning("定时批量写入权限观测日志失败", LOG_COMMAND, e=exc)
async def _drain_batch(buffer: deque[T]) -> list[T]:
batch: list[T] = []
async with _buffer_lock:
while buffer and len(batch) < _FLUSH_BATCH_SIZE:
batch.append(buffer.popleft())
return batch
async def _restore_batch(buffer: deque[T], batch: list[T]) -> None:
async with _buffer_lock:
retain_count = max(_BUFFER_MAX_RETAIN - len(buffer), 0)
for record in reversed(batch[-retain_count:]):
buffer.appendleft(record)
def _is_schema_mismatch_error(exc: Exception) -> bool:
message = str(exc).lower()
return any(
marker in message
for marker in (
"no column named",
"unknown column",
"column does not exist",
"no such column",
)
)
async def _try_repair_auth_schema_once() -> bool:
global _last_schema_repair_at
now = time.monotonic()
if now - _last_schema_repair_at < _SCHEMA_REPAIR_INTERVAL_SECONDS:
return False
_last_schema_repair_at = now
try:
from zhenxun.services.db_context.schema_guard import repair_table_schema
await repair_table_schema("auth_decision_log")
await repair_table_schema("runtime_backpressure_log")
return True
except Exception as exc:
logger.warning("权限观测日志表结构自修复失败", LOG_COMMAND, e=exc)
return False
async def flush_auth_observability_buffer(reason: str) -> int:
async with _flush_lock:
written = 0
while True:
auth_batch = await _drain_batch(_auth_decision_buffer)
backpressure_batch = await _drain_batch(_backpressure_buffer)
if not auth_batch and not backpressure_batch:
break
try:
if auth_batch:
await AuthDecisionLog.bulk_create(
[record.to_model() for record in auth_batch],
_FLUSH_BATCH_SIZE,
)
written += len(auth_batch)
if backpressure_batch:
await RuntimeBackpressureLog.bulk_create(
[record.to_model() for record in backpressure_batch],
_FLUSH_BATCH_SIZE,
)
written += len(backpressure_batch)
except Exception as exc:
if _is_schema_mismatch_error(exc):
if await _try_repair_auth_schema_once():
try:
if auth_batch:
await AuthDecisionLog.bulk_create(
[record.to_model() for record in auth_batch],
_FLUSH_BATCH_SIZE,
)
written += len(auth_batch)
if backpressure_batch:
await RuntimeBackpressureLog.bulk_create(
[
record.to_model()
for record in backpressure_batch
],
_FLUSH_BATCH_SIZE,
)
written += len(backpressure_batch)
continue
except Exception as retry_exc:
exc = retry_exc
dropped = len(auth_batch) + len(backpressure_batch)
logger.warning(
f"{reason}批量写入权限观测日志遇到表结构不匹配,"
f"已丢弃低优先级观测日志 {dropped} 条,等待下次启动修复",
LOG_COMMAND,
e=exc,
)
return written
await _restore_batch(_auth_decision_buffer, auth_batch)
await _restore_batch(_backpressure_buffer, backpressure_batch)
logger.error(f"{reason}批量写入权限观测日志失败", LOG_COMMAND, e=exc)
return written
if written:
logger.debug(f"{reason}批量写入权限观测日志 {written} 条", LOG_COMMAND)
return written
async def stop_auth_observability_buffer() -> int:
global _flush_task
task = _flush_task
_flush_task = None
if task is not None:
task.cancel()
with contextlib.suppress(BaseException):
await task
return await flush_auth_observability_buffer("关闭")
def _percentile(values: list[float], ratio: float) -> float:
if not values:
return 0.0
ordered = sorted(values)
index = min(max(round((len(ordered) - 1) * ratio), 0), len(ordered) - 1)
return round(ordered[index], 3)
def _bucket_counts(rows: list[dict[str, Any]], field: str) -> dict[str, int]:
counts: dict[str, int] = {}
for row in rows:
key = str(row.get(field) or "<none>")
counts[key] = counts.get(key, 0) + 1
return counts
def _lane_budget_advice(rows: list[dict[str, Any]]) -> dict[str, dict[str, Any]]:
buckets: dict[str, list[dict[str, Any]]] = {}
for row in rows:
lane = str(row.get("lane") or "<unknown>")
buckets.setdefault(lane, []).append(row)
advice: dict[str, dict[str, Any]] = {}
for lane, items in buckets.items():
if lane == "<unknown>":
continue
durations = [float(item.get("duration_ms") or 0.0) for item in items]
slow_waits = sum(1 for value in durations if value >= 200.0)
active_max = max(
(int(item.get("active_count") or 0) for item in items), default=0
)
total = len(items)
if not total:
continue
pressure_ratio = slow_waits / total
if pressure_ratio >= 0.2 or active_max >= 5:
action = "increase_or_split"
elif pressure_ratio == 0 and active_max <= 1 and total >= 20:
action = "can_reduce"
else:
action = "keep"
advice[lane] = {
"samples": total,
"slow_waits": slow_waits,
"pressure_ratio": round(pressure_ratio, 3),
"active_max": active_max,
"p95_duration_ms": _percentile(durations, 0.95),
"action": action,
}
return advice
def _query_placeholder() -> str:
try:
connection = Tortoise.get_connection("default")
if (
getattr(connection, "capabilities", None)
and getattr(
connection.capabilities,
"dialect",
"",
)
== "postgres"
):
return "$1"
except Exception:
return "?"
return "?"
async def build_auth_observability_report(*, hours: float = 24.0) -> dict[str, Any]:
since = datetime.now() - timedelta(hours=hours)
db = Tortoise.get_connection("default")
placeholder = _query_placeholder()
auth_rows = await db.execute_query_dict(
"SELECT module, effect, reason, shadow_effect, shadow_reason, latency_ms, "
f"overloaded FROM auth_decision_log WHERE create_time >= {placeholder} "
"ORDER BY create_time DESC LIMIT 100000",
[since],
)
backpressure_rows = await db.execute_query_dict(
"SELECT scope_key, lane, reason, action, queue_size, active_count, duration_ms "
f"FROM runtime_backpressure_log WHERE create_time >= {placeholder} "
"ORDER BY create_time DESC LIMIT 100000",
[since],
)
module_buckets: dict[str, list[dict[str, Any]]] = {}
for row in auth_rows:
module_buckets.setdefault(str(row.get("module") or "<unknown>"), []).append(row)
module_stats: list[dict[str, Any]] = []
for module, items in module_buckets.items():
latencies = [float(item.get("latency_ms") or 0.0) for item in items]
module_stats.append(
{
"module": module,
"total": len(items),
"effects": _bucket_counts(items, "effect"),
"shadow_effects": _bucket_counts(items, "shadow_effect"),
"avg_latency_ms": round(sum(latencies) / len(latencies), 3)
if latencies
else 0.0,
"p95_latency_ms": _percentile(latencies, 0.95),
"overloaded": sum(1 for item in items if bool(item.get("overloaded"))),
}
)
backpressure_buckets: dict[str, list[dict[str, Any]]] = {}
for row in backpressure_rows:
key = f"{row.get('lane') or '<unknown>'}:{row.get('reason') or '<none>'}"
backpressure_buckets.setdefault(key, []).append(row)
backpressure_stats: list[dict[str, Any]] = []
for key, items in backpressure_buckets.items():
durations = [float(item.get("duration_ms") or 0.0) for item in items]
backpressure_stats.append(
{
"key": key,
"total": len(items),
"actions": _bucket_counts(items, "action"),
"avg_duration_ms": round(sum(durations) / len(durations), 3)
if durations
else 0.0,
"p95_duration_ms": _percentile(durations, 0.95),
}
)
return {
"created_at": datetime.now().isoformat(timespec="seconds"),
"window_hours": hours,
"auth_decisions": {
"total": len(auth_rows),
"effects": _bucket_counts(auth_rows, "effect"),
"shadow_effects": _bucket_counts(auth_rows, "shadow_effect"),
"top_modules_by_p95": sorted(
module_stats,
key=lambda item: (item["p95_latency_ms"], item["total"]),
reverse=True,
)[:30],
},
"backpressure": {
"total": len(backpressure_rows),
"lane_budget_advice": _lane_budget_advice(backpressure_rows),
"top_reasons": sorted(
backpressure_stats,
key=lambda item: (item["total"], item["p95_duration_ms"]),
reverse=True,
)[:30],
},
}
@PriorityLifecycle.on_shutdown(priority=90)
async def _flush_auth_observability_buffer_on_shutdown() -> None:
await stop_auth_observability_buffer()
+378 -72
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import asyncio
from contextvars import ContextVar
from dataclasses import dataclass, field
import json
import os
@@ -35,6 +36,8 @@ def _coerce_int(value, default: int) -> int:
# RuntimeCache 以模型 save/delete 主动失效为主,周期 refresh 只做兜底。
# 这些默认值避免低压力运行时频繁全量扫表。
# 权限检查热路径以 RuntimeCache/AuthSnapshot 为唯一数据入口;普通业务的
# DataAccess/CacheRoot 缓存不能替代这里的运行态快照。
PLUGININFO_MEM_REFRESH_INTERVAL = 3600 # 60分钟
BAN_MEM_REFRESH_INTERVAL = 300
BAN_MEM_CLEAN_INTERVAL = 60
@@ -52,10 +55,16 @@ LIMIT_MEM_REFRESH_INTERVAL = 900 # 15分钟
LIMIT_MEM_NEGATIVE_TTL = 30
RUNTIME_CACHE_SYNC_ENABLED = True
RUNTIME_CACHE_SYNC_CHANNEL = "ZHENXUN_RUNTIME_CACHE_SYNC"
RUNTIME_CACHE_LOAD_RETRY_SECONDS = 1.0
RUNTIME_CACHE_STARTUP_REFRESH_SKIP_SECONDS = 5.0
INSTANCE_ID = uuid.uuid4().hex
_CACHE_READY_EVENT = asyncio.Event()
_APPLYING_REMOTE_CACHE_EVENT: ContextVar[bool] = ContextVar(
"APPLYING_REMOTE_RUNTIME_CACHE_EVENT",
default=False,
)
def _env_get(name: str, default: str | None = None) -> str | None:
@@ -180,6 +189,65 @@ class PluginInfoSnapshot:
plugin._saved_in_db = True
return plugin
def to_payload(self) -> dict[str, Any]:
return {
"id": self.id,
"module": self.module,
"module_path": self.module_path,
"name": self.name,
"status": self.status,
"block_type": self.block_type.value if self.block_type else None,
"load_status": self.load_status,
"author": self.author,
"version": self.version,
"level": self.level,
"default_status": self.default_status,
"limit_superuser": self.limit_superuser,
"menu_type": self.menu_type,
"plugin_type": self.plugin_type.value if self.plugin_type else None,
"cost_gold": self.cost_gold,
"admin_level": self.admin_level,
"ignore_prompt": self.ignore_prompt,
"is_delete": self.is_delete,
"parent": self.parent,
"is_show": self.is_show,
"ignore_statistics": self.ignore_statistics,
"impression": self.impression,
}
@classmethod
def from_payload(cls, payload: dict[str, Any]) -> "PluginInfoSnapshot":
block_type = payload.get("block_type")
if block_type is not None and not isinstance(block_type, BlockType):
block_type = BlockType(block_type)
plugin_type = payload.get("plugin_type")
if plugin_type is not None and not isinstance(plugin_type, PluginType):
plugin_type = PluginType(plugin_type)
return cls(
id=int(payload.get("id", 0) or 0),
module=str(payload.get("module", "") or ""),
module_path=str(payload.get("module_path", "") or ""),
name=str(payload.get("name", "") or ""),
status=bool(payload.get("status", True)),
block_type=block_type,
load_status=bool(payload.get("load_status", True)),
author=payload.get("author"),
version=payload.get("version"),
level=int(payload.get("level", 0) or 0),
default_status=bool(payload.get("default_status", True)),
limit_superuser=bool(payload.get("limit_superuser", False)),
menu_type=str(payload.get("menu_type", "") or ""),
plugin_type=plugin_type,
cost_gold=int(payload.get("cost_gold", 0) or 0),
admin_level=payload.get("admin_level"),
ignore_prompt=bool(payload.get("ignore_prompt", False)),
is_delete=bool(payload.get("is_delete", False)),
parent=payload.get("parent"),
is_show=bool(payload.get("is_show", True)),
ignore_statistics=bool(payload.get("ignore_statistics", False)),
impression=float(payload.get("impression", 0) or 0),
)
@dataclass(frozen=True)
class BanEntry:
@@ -648,18 +716,87 @@ class RuntimeCacheSync:
cache_type = payload.get("type")
action = payload.get("action")
data = payload.get("data") or {}
if cache_type == "bot":
await BotMemoryCache.apply_sync_event(action, data)
elif cache_type == "group":
await GroupMemoryCache.apply_sync_event(action, data)
elif cache_type == "ban":
await BanMemoryCache.apply_sync_event(action, data)
elif cache_type == "level":
await LevelUserMemoryCache.apply_sync_event(action, data)
elif cache_type == "task":
await TaskInfoMemoryCache.apply_sync_event(action, data)
elif cache_type == "plugin_limit":
await PluginLimitMemoryCache.apply_sync_event(action, data)
token = _APPLYING_REMOTE_CACHE_EVENT.set(True)
try:
if cache_type == "bot":
await BotMemoryCache.apply_sync_event(action, data)
elif cache_type == "group":
await GroupMemoryCache.apply_sync_event(action, data)
elif cache_type == "ban":
await BanMemoryCache.apply_sync_event(action, data)
elif cache_type == "level":
await LevelUserMemoryCache.apply_sync_event(action, data)
elif cache_type == "task":
await TaskInfoMemoryCache.apply_sync_event(action, data)
elif cache_type == "plugin_limit":
await PluginLimitMemoryCache.apply_sync_event(action, data)
elif cache_type == "plugin":
await PluginInfoMemoryCache.apply_sync_event(action, data)
finally:
_APPLYING_REMOTE_CACHE_EVENT.reset(token)
class RuntimeCacheMutation:
"""Small helpers for runtime cache mutation bookkeeping.
Cache classes still own their storage layout. This helper centralizes the
shared mutation side effects: health markers, negative-cache cleanup and
cross-process publish.
"""
_load_locks: ClassVar[dict[str, asyncio.Lock]] = {}
_retry_after: ClassVar[dict[str, float]] = {}
@classmethod
async def ensure_loaded(cls, cache_cls: type, label: str) -> None:
if getattr(cache_cls, "_loaded", False):
return
now = time.monotonic()
if cls._retry_after.get(label, 0.0) > now:
return
lock = cls._load_locks.setdefault(label, asyncio.Lock())
async with lock:
if getattr(cache_cls, "_loaded", False):
return
now = time.monotonic()
if cls._retry_after.get(label, 0.0) > now:
return
try:
await cache_cls.refresh()
except Exception as exc:
cls.mark_error(cache_cls, exc)
cls._retry_after[label] = (
time.monotonic() + RUNTIME_CACHE_LOAD_RETRY_SECONDS
)
raise
@staticmethod
def mark_refreshed(cache_cls: type) -> None:
setattr(cache_cls, "_loaded", True)
setattr(cache_cls, "_last_refresh", time.time())
setattr(cache_cls, "_last_error", None)
@staticmethod
def mark_error(cache_cls: type, exc: Exception) -> None:
setattr(cache_cls, "_last_error", f"{type(exc).__name__}: {exc}")
@staticmethod
def clear_negative_key(cache_cls: type, key: object) -> None:
negative = getattr(cache_cls, "_negative", None)
if isinstance(negative, dict):
negative.pop(key, None)
@staticmethod
def clear_negative_all(cache_cls: type) -> None:
negative = getattr(cache_cls, "_negative", None)
if isinstance(negative, dict):
negative.clear()
@staticmethod
def publish(cache_type: str, action: str, data: dict[str, Any]) -> None:
if _APPLYING_REMOTE_CACHE_EVENT.get():
return
RuntimeCacheSync.publish_event(cache_type, action, data)
class PluginInfoMemoryCache:
@@ -669,6 +806,7 @@ class PluginInfoMemoryCache:
_loaded: ClassVar[bool] = False
_refresh_task: ClassVar[asyncio.Task | None] = None
_last_refresh: ClassVar[float] = 0.0
_last_error: ClassVar[str | None] = None
@classmethod
def _to_model(cls, snapshot: PluginInfoSnapshot | None) -> "PluginInfo | None":
@@ -687,6 +825,10 @@ class PluginInfoMemoryCache:
cls._by_module.pop(old.module, None)
cls._by_module_path[snapshot.module_path] = snapshot
@staticmethod
def _module_snapshot_rank(snapshot: PluginInfoSnapshot) -> tuple[int, int]:
return (1 if snapshot.load_status else 0, snapshot.id)
@classmethod
async def refresh(cls) -> None:
from zhenxun.models.plugin_info import PluginInfo
@@ -698,13 +840,16 @@ class PluginInfoMemoryCache:
for plugin in plugins:
snapshot = PluginInfoSnapshot.from_model(plugin)
if snapshot.module:
by_module[snapshot.module] = snapshot
current = by_module.get(snapshot.module)
if current is None or cls._module_snapshot_rank(
snapshot
) >= cls._module_snapshot_rank(current):
by_module[snapshot.module] = snapshot
if snapshot.module_path:
by_module_path[snapshot.module_path] = snapshot
cls._by_module = by_module
cls._by_module_path = by_module_path
cls._loaded = True
cls._last_refresh = time.time()
RuntimeCacheMutation.mark_refreshed(cls)
logger.debug(
f"plugin cache refreshed: {len(by_module)} entries", LOG_COMMAND
)
@@ -713,7 +858,7 @@ class PluginInfoMemoryCache:
async def ensure_loaded(cls) -> None:
if cls._loaded:
return
await cls.refresh()
await RuntimeCacheMutation.ensure_loaded(cls, "plugin")
@classmethod
def is_loaded(cls) -> bool:
@@ -749,8 +894,6 @@ class PluginInfoMemoryCache:
return
snapshot = PluginInfoSnapshot.from_model(plugin)
cls._store_snapshot(snapshot)
cls._loaded = True
cls._last_refresh = time.time()
@classmethod
def remove_by_module(cls, module: str) -> None:
@@ -765,8 +908,16 @@ class PluginInfoMemoryCache:
async with cls._lock:
snapshot = PluginInfoSnapshot.from_model(plugin)
cls._store_snapshot(snapshot)
cls._loaded = True
cls._last_refresh = time.time()
RuntimeCacheMutation.publish("plugin", "upsert", snapshot.to_payload())
@classmethod
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
try:
snapshot = PluginInfoSnapshot.from_payload(payload)
except Exception:
return
async with cls._lock:
cls._store_snapshot(snapshot)
@classmethod
async def remove(
@@ -783,6 +934,20 @@ class PluginInfoMemoryCache:
snapshot = cls._by_module_path.pop(module_path, None)
if snapshot and snapshot.module:
cls._by_module.pop(snapshot.module, None)
RuntimeCacheMutation.publish(
"plugin",
"delete",
{"module": module, "module_path": module_path},
)
@classmethod
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
if action == "upsert":
await cls.upsert_from_payload(data)
elif action == "delete":
await cls.remove(data.get("module"), data.get("module_path"))
elif action == "refresh":
await cls.refresh()
@classmethod
async def _refresh_loop(cls, interval: int) -> None:
@@ -815,6 +980,8 @@ class BotMemoryCache:
_negative: ClassVar[dict[str, float]] = {}
_loaded: ClassVar[bool] = False
_refresh_task: ClassVar[asyncio.Task | None] = None
_last_refresh: ClassVar[float] = 0.0
_last_error: ClassVar[str | None] = None
@classmethod
def _normalize(cls, bot_id: str | None) -> str | None:
@@ -833,7 +1000,7 @@ class BotMemoryCache:
if not expire_at:
return False
if expire_at <= time.time():
cls._negative.pop(bot_id, None)
RuntimeCacheMutation.clear_negative_key(cls, bot_id)
return False
return True
@@ -851,15 +1018,15 @@ class BotMemoryCache:
async with cls._lock:
records = await BotConsole.all()
cls._by_id = {str(r.bot_id): BotSnapshot.from_model(r) for r in records}
cls._negative = {}
cls._loaded = True
RuntimeCacheMutation.clear_negative_all(cls)
RuntimeCacheMutation.mark_refreshed(cls)
logger.debug(f"bot cache refreshed: {len(cls._by_id)} entries", LOG_COMMAND)
@classmethod
async def ensure_loaded(cls) -> None:
if cls._loaded:
return
await cls.refresh()
await RuntimeCacheMutation.ensure_loaded(cls, "bot")
@classmethod
def is_loaded(cls) -> bool:
@@ -918,15 +1085,16 @@ class BotMemoryCache:
available_tasks=entry.available_tasks,
)
cls._by_id[bot_id] = updated
RuntimeCacheSync.publish_event("bot", "upsert", updated.to_payload())
RuntimeCacheMutation.clear_negative_key(cls, bot_id)
RuntimeCacheMutation.publish("bot", "upsert", updated.to_payload())
@classmethod
async def upsert_from_model(cls, record) -> None:
entry = BotSnapshot.from_model(record)
async with cls._lock:
cls._by_id[entry.bot_id] = entry
cls._negative.pop(entry.bot_id, None)
RuntimeCacheSync.publish_event("bot", "upsert", entry.to_payload())
RuntimeCacheMutation.clear_negative_key(cls, entry.bot_id)
RuntimeCacheMutation.publish("bot", "upsert", entry.to_payload())
@classmethod
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
@@ -935,7 +1103,7 @@ class BotMemoryCache:
return
async with cls._lock:
cls._by_id[entry.bot_id] = entry
cls._negative.pop(entry.bot_id, None)
RuntimeCacheMutation.clear_negative_key(cls, entry.bot_id)
@classmethod
async def remove(cls, bot_id: str | None) -> None:
@@ -944,7 +1112,8 @@ class BotMemoryCache:
return
async with cls._lock:
cls._by_id.pop(bot_id, None)
RuntimeCacheSync.publish_event("bot", "delete", {"bot_id": bot_id})
RuntimeCacheMutation.clear_negative_key(cls, bot_id)
RuntimeCacheMutation.publish("bot", "delete", {"bot_id": bot_id})
@classmethod
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
@@ -986,6 +1155,8 @@ class GroupMemoryCache:
_negative: ClassVar[dict[tuple[str, str], float]] = {}
_loaded: ClassVar[bool] = False
_refresh_task: ClassVar[asyncio.Task | None] = None
_last_refresh: ClassVar[float] = 0.0
_last_error: ClassVar[str | None] = None
@classmethod
def _normalize(cls, value: str | None) -> str | None:
@@ -1014,7 +1185,7 @@ class GroupMemoryCache:
if not expire_at:
return False
if expire_at <= time.time():
cls._negative.pop(key, None)
RuntimeCacheMutation.clear_negative_key(cls, key)
return False
return True
@@ -1038,15 +1209,15 @@ class GroupMemoryCache:
if key:
by_key[key] = entry
cls._by_key = by_key
cls._negative = {}
cls._loaded = True
RuntimeCacheMutation.clear_negative_all(cls)
RuntimeCacheMutation.mark_refreshed(cls)
logger.debug(f"group cache refreshed: {len(by_key)} entries", LOG_COMMAND)
@classmethod
async def ensure_loaded(cls) -> None:
if cls._loaded:
return
await cls.refresh()
await RuntimeCacheMutation.ensure_loaded(cls, "group")
@classmethod
def is_loaded(cls) -> bool:
@@ -1094,8 +1265,8 @@ class GroupMemoryCache:
return
async with cls._lock:
cls._by_key[key] = entry
cls._negative.pop(key, None)
RuntimeCacheSync.publish_event("group", "upsert", entry.to_payload())
RuntimeCacheMutation.clear_negative_key(cls, key)
RuntimeCacheMutation.publish("group", "upsert", entry.to_payload())
@classmethod
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
@@ -1105,7 +1276,7 @@ class GroupMemoryCache:
return
async with cls._lock:
cls._by_key[key] = entry
cls._negative.pop(key, None)
RuntimeCacheMutation.clear_negative_key(cls, key)
@classmethod
async def remove(cls, group_id: str | None, channel_id: str | None = None) -> None:
@@ -1114,7 +1285,8 @@ class GroupMemoryCache:
return
async with cls._lock:
cls._by_key.pop(key, None)
RuntimeCacheSync.publish_event(
RuntimeCacheMutation.clear_negative_key(cls, key)
RuntimeCacheMutation.publish(
"group", "delete", {"group_id": key[0], "channel_id": key[1] or None}
)
@@ -1186,7 +1358,7 @@ class LevelUserMemoryCache:
if not expire_at:
return False
if expire_at <= time.time():
cls._negative.pop(key, None)
RuntimeCacheMutation.clear_negative_key(cls, key)
return False
return True
@@ -1215,16 +1387,15 @@ class LevelUserMemoryCache:
by_user_max[entry.user_id] = entry.user_level
cls._by_key = by_key
cls._by_user_max = by_user_max
cls._negative = {}
cls._loaded = True
cls._last_refresh = time.time()
RuntimeCacheMutation.clear_negative_all(cls)
RuntimeCacheMutation.mark_refreshed(cls)
logger.debug(f"level cache refreshed: {len(by_key)} entries", LOG_COMMAND)
@classmethod
async def ensure_loaded(cls) -> None:
if cls._loaded:
return
await cls.refresh()
await RuntimeCacheMutation.ensure_loaded(cls, "level")
@classmethod
def is_loaded(cls) -> bool:
@@ -1310,13 +1481,13 @@ class LevelUserMemoryCache:
async with cls._lock:
prev = cls._by_key.get(key)
cls._by_key[key] = entry
cls._negative.pop(key, None)
current = cls._by_user_max.get(entry.user_id, 0)
if entry.user_level >= current:
cls._by_user_max[entry.user_id] = entry.user_level
elif prev and prev.user_level == current and entry.user_level < current:
cls._recalc_user_max(entry.user_id)
RuntimeCacheSync.publish_event("level", "upsert", entry.to_payload())
RuntimeCacheMutation.clear_negative_key(cls, key)
RuntimeCacheMutation.publish("level", "upsert", entry.to_payload())
@classmethod
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
@@ -1327,12 +1498,12 @@ class LevelUserMemoryCache:
async with cls._lock:
prev = cls._by_key.get(key)
cls._by_key[key] = entry
cls._negative.pop(key, None)
current = cls._by_user_max.get(entry.user_id, 0)
if entry.user_level >= current:
cls._by_user_max[entry.user_id] = entry.user_level
elif prev and prev.user_level == current and entry.user_level < current:
cls._recalc_user_max(entry.user_id)
RuntimeCacheMutation.clear_negative_key(cls, key)
@classmethod
async def remove(cls, user_id: str | None, group_id: str | None) -> None:
@@ -1343,7 +1514,8 @@ class LevelUserMemoryCache:
removed = cls._by_key.pop(key, None)
if removed and cls._by_user_max.get(removed.user_id) == removed.user_level:
cls._recalc_user_max(removed.user_id)
RuntimeCacheSync.publish_event(
RuntimeCacheMutation.clear_negative_key(cls, key)
RuntimeCacheMutation.publish(
"level", "delete", {"user_id": key[0], "group_id": key[1] or None}
)
@@ -1399,6 +1571,8 @@ class TaskInfoMemoryCache:
_negative: ClassVar[dict[str, float]] = {}
_loaded: ClassVar[bool] = False
_refresh_task: ClassVar[asyncio.Task | None] = None
_last_refresh: ClassVar[float] = 0.0
_last_error: ClassVar[str | None] = None
@classmethod
def _normalize(cls, module: str | None) -> str | None:
@@ -1417,7 +1591,7 @@ class TaskInfoMemoryCache:
if not expire_at:
return False
if expire_at <= time.time():
cls._negative.pop(module, None)
RuntimeCacheMutation.clear_negative_key(cls, module)
return False
return True
@@ -1443,8 +1617,8 @@ class TaskInfoMemoryCache:
by_name[entry.name] = entry
cls._by_module = by_module
cls._by_name = by_name
cls._negative = {}
cls._loaded = True
RuntimeCacheMutation.clear_negative_all(cls)
RuntimeCacheMutation.mark_refreshed(cls)
logger.debug(
f"task info cache refreshed: {len(cls._by_module)} entries",
LOG_COMMAND,
@@ -1454,7 +1628,7 @@ class TaskInfoMemoryCache:
async def ensure_loaded(cls) -> None:
if cls._loaded:
return
await cls.refresh()
await RuntimeCacheMutation.ensure_loaded(cls, "task")
@classmethod
async def get(cls, module: str | None) -> TaskInfoSnapshot | None:
@@ -1488,10 +1662,21 @@ class TaskInfoMemoryCache:
@classmethod
async def is_disabled(cls, module: str | None) -> bool:
"""Backward-compatible runtime disabled check for passive tasks."""
return await cls.is_runtime_disabled(module)
@classmethod
async def is_runtime_disabled(cls, module: str | None) -> bool:
"""Return whether a passive task is unavailable at runtime.
Runtime passive availability is defined by TaskInfo.status and
TaskInfo.load_status. Bot/group scoped block lists are checked by
CommonUtils.task_is_block().
"""
entry = await cls.get(module)
if not entry:
return False
return not entry.status
return not entry.status or not entry.load_status
@classmethod
async def upsert_from_model(cls, record) -> None:
@@ -1500,8 +1685,8 @@ class TaskInfoMemoryCache:
cls._by_module[entry.module] = entry
if entry.name:
cls._by_name[entry.name] = entry
cls._negative.pop(entry.module, None)
RuntimeCacheSync.publish_event("task", "upsert", entry.to_payload())
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
RuntimeCacheMutation.publish("task", "upsert", entry.to_payload())
@classmethod
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
@@ -1512,7 +1697,7 @@ class TaskInfoMemoryCache:
cls._by_module[entry.module] = entry
if entry.name:
cls._by_name[entry.name] = entry
cls._negative.pop(entry.module, None)
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
@classmethod
async def remove(cls, module: str | None) -> None:
@@ -1525,7 +1710,8 @@ class TaskInfoMemoryCache:
current = cls._by_name.get(removed.name)
if current and current.module == removed.module:
cls._by_name.pop(removed.name, None)
RuntimeCacheSync.publish_event("task", "delete", {"module": module})
RuntimeCacheMutation.clear_negative_key(cls, module)
RuntimeCacheMutation.publish("task", "delete", {"module": module})
@classmethod
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
@@ -1568,6 +1754,8 @@ class PluginLimitMemoryCache:
_negative: ClassVar[dict[str, float]] = {}
_loaded: ClassVar[bool] = False
_refresh_task: ClassVar[asyncio.Task | None] = None
_last_refresh: ClassVar[float] = 0.0
_last_error: ClassVar[str | None] = None
@classmethod
def _normalize(cls, value: str | None) -> str | None:
@@ -1586,7 +1774,7 @@ class PluginLimitMemoryCache:
if not expire_at:
return False
if expire_at <= time.time():
cls._negative.pop(module, None)
RuntimeCacheMutation.clear_negative_key(cls, module)
return False
return True
@@ -1611,8 +1799,8 @@ class PluginLimitMemoryCache:
by_module.setdefault(entry.module, []).append(entry)
cls._by_id = by_id
cls._by_module = by_module
cls._negative = {}
cls._loaded = True
RuntimeCacheMutation.clear_negative_all(cls)
RuntimeCacheMutation.mark_refreshed(cls)
logger.debug(
f"plugin limit cache refreshed: {len(by_id)} entries",
LOG_COMMAND,
@@ -1622,7 +1810,7 @@ class PluginLimitMemoryCache:
async def ensure_loaded(cls) -> None:
if cls._loaded:
return
await cls.refresh()
await RuntimeCacheMutation.ensure_loaded(cls, "plugin_limit")
@classmethod
def is_loaded(cls) -> bool:
@@ -1668,7 +1856,7 @@ class PluginLimitMemoryCache:
async def upsert_from_model(cls, record) -> None:
entry = PluginLimitSnapshot.from_model(record)
await cls._upsert_entry(entry)
RuntimeCacheSync.publish_event("plugin_limit", "upsert", entry.to_payload())
RuntimeCacheMutation.publish("plugin_limit", "upsert", entry.to_payload())
@classmethod
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
@@ -1695,6 +1883,7 @@ class PluginLimitMemoryCache:
for item in cls._by_module.get(entry.module, [])
if item.id != entry.id
]
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
return
cls._by_id[entry.id] = entry
module_limits = [
@@ -1704,7 +1893,7 @@ class PluginLimitMemoryCache:
]
module_limits.append(entry)
cls._by_module[entry.module] = module_limits
cls._negative.pop(entry.module, None)
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
@classmethod
async def remove_by_id(cls, limit_id: int | None) -> None:
@@ -1718,7 +1907,8 @@ class PluginLimitMemoryCache:
for item in cls._by_module.get(entry.module, [])
if item.id != entry.id
]
RuntimeCacheSync.publish_event("plugin_limit", "delete", {"id": int(limit_id)})
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
RuntimeCacheMutation.publish("plugin_limit", "delete", {"id": int(limit_id)})
@classmethod
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
@@ -1764,6 +1954,8 @@ class BanMemoryCache:
_refresh_task: ClassVar[asyncio.Task | None] = None
_cleanup_task: ClassVar[asyncio.Task | None] = None
_remove_tasks: ClassVar[set[asyncio.Task]] = set()
_last_refresh: ClassVar[float] = 0.0
_last_error: ClassVar[str | None] = None
@classmethod
def _normalize_id(cls, value: str | None) -> str | None:
@@ -1788,7 +1980,7 @@ class BanMemoryCache:
if not expire_at:
return False
if expire_at <= time.time():
cls._negative.pop(key, None)
RuntimeCacheMutation.clear_negative_key(cls, key)
return False
return True
@@ -1842,8 +2034,8 @@ class BanMemoryCache:
cls._by_user = by_user
cls._by_group = by_group
cls._by_user_group = by_user_group
cls._negative = {}
cls._loaded = True
RuntimeCacheMutation.clear_negative_all(cls)
RuntimeCacheMutation.mark_refreshed(cls)
logger.debug(
"ban cache refreshed: "
f"user={len(by_user)} group={len(by_group)} "
@@ -1855,7 +2047,7 @@ class BanMemoryCache:
async def ensure_loaded(cls) -> None:
if cls._loaded:
return
await cls.refresh()
await RuntimeCacheMutation.ensure_loaded(cls, "ban")
@classmethod
def is_loaded(cls) -> bool:
@@ -1873,13 +2065,13 @@ class BanMemoryCache:
cls._by_user[entry.user_id] = entry
elif entry.group_id:
cls._by_group[entry.group_id] = entry
cls._negative = {}
RuntimeCacheSync.publish_event("ban", "upsert", entry.to_payload())
RuntimeCacheMutation.clear_negative_all(cls)
RuntimeCacheMutation.publish("ban", "upsert", entry.to_payload())
@classmethod
async def remove(cls, user_id: str | None, group_id: str | None) -> None:
await cls._remove_local(user_id, group_id)
RuntimeCacheSync.publish_event(
RuntimeCacheMutation.publish(
"ban", "delete", {"user_id": user_id, "group_id": group_id}
)
@@ -1894,7 +2086,7 @@ class BanMemoryCache:
cls._by_user.pop(user_id, None)
elif group_id:
cls._by_group.pop(group_id, None)
cls._negative = {}
RuntimeCacheMutation.clear_negative_all(cls)
@classmethod
def _get_entry(cls, user_id: str | None, group_id: str | None) -> BanEntry | None:
@@ -1995,7 +2187,7 @@ class BanMemoryCache:
elif entry.group_id:
cls._by_group.pop(entry.group_id, None)
if expired:
cls._negative = {}
RuntimeCacheMutation.clear_negative_all(cls)
if not delete_db or not expired:
return
from tortoise.expressions import Q
@@ -2064,7 +2256,7 @@ class BanMemoryCache:
cls._by_user[entry.user_id] = entry
elif entry.group_id:
cls._by_group[entry.group_id] = entry
cls._negative = {}
RuntimeCacheMutation.clear_negative_all(cls)
@classmethod
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
@@ -2078,12 +2270,126 @@ class BanMemoryCache:
async def _safe_refresh(cache_cls: type, label: str) -> None:
"""安全地刷新单个缓存,异常不影响其他缓存。"""
if getattr(cache_cls, "_loaded", False):
last_refresh = float(getattr(cache_cls, "_last_refresh", 0.0) or 0.0)
if time.time() - last_refresh <= RUNTIME_CACHE_STARTUP_REFRESH_SKIP_SECONDS:
logger.debug(f"{label} cache startup refresh skipped", LOG_COMMAND)
return
try:
await cache_cls.refresh()
except Exception as exc:
RuntimeCacheMutation.mark_error(cache_cls, exc)
logger.error(f"{label} cache init failed", LOG_COMMAND, e=exc)
def _cache_health(
cache_cls: type,
*,
entry_count: int,
negative_count: int = 0,
) -> dict[str, Any]:
return {
"loaded": bool(getattr(cache_cls, "_loaded", False)),
"entry_count": entry_count,
"last_refresh": float(getattr(cache_cls, "_last_refresh", 0.0) or 0.0),
"negative_count": negative_count,
"last_error": getattr(cache_cls, "_last_error", None),
}
def health_snapshot() -> dict[str, dict[str, Any]]:
"""Return in-memory runtime cache health without touching the database."""
return {
"plugin": _cache_health(
PluginInfoMemoryCache,
entry_count=len(PluginInfoMemoryCache._by_module),
),
"bot": _cache_health(
BotMemoryCache,
entry_count=len(BotMemoryCache._by_id),
negative_count=len(BotMemoryCache._negative),
),
"group": _cache_health(
GroupMemoryCache,
entry_count=len(GroupMemoryCache._by_key),
negative_count=len(GroupMemoryCache._negative),
),
"level": _cache_health(
LevelUserMemoryCache,
entry_count=len(LevelUserMemoryCache._by_key),
negative_count=len(LevelUserMemoryCache._negative),
),
"task": _cache_health(
TaskInfoMemoryCache,
entry_count=len(TaskInfoMemoryCache._by_module),
negative_count=len(TaskInfoMemoryCache._negative),
),
"plugin_limit": _cache_health(
PluginLimitMemoryCache,
entry_count=len(PluginLimitMemoryCache._by_id),
negative_count=len(PluginLimitMemoryCache._negative),
),
"ban": _cache_health(
BanMemoryCache,
entry_count=(
len(BanMemoryCache._by_user)
+ len(BanMemoryCache._by_group)
+ len(BanMemoryCache._by_user_group)
),
negative_count=len(BanMemoryCache._negative),
),
}
def passive_status_snapshot(max_modules: int = 50) -> dict[str, Any]:
"""Return passive-task state from in-memory caches only.
This is a local diagnostic helper: it does not query or write the database,
and it is not used by runtime decisions.
"""
tasks = list(TaskInfoMemoryCache._by_module.values())
disabled = sorted(task.module for task in tasks if not task.status)
unloaded = sorted(task.module for task in tasks if not task.load_status)
runtime_enabled = [
task.module for task in tasks if task.status and task.load_status
]
bot_block_total = sum(
len(_parse_block_modules(bot.block_tasks))
for bot in BotMemoryCache._by_id.values()
)
group_block_total = sum(
len(group.block_task_set) + len(group.superuser_block_task_set)
for group in GroupMemoryCache._by_key.values()
)
return {
"cache": health_snapshot(),
"passive_tasks": {
"total": len(tasks),
"status_enabled": sum(1 for task in tasks if task.status),
"load_status_enabled": sum(1 for task in tasks if task.load_status),
"runtime_enabled": len(runtime_enabled),
"disabled_modules": disabled[:max_modules],
"disabled_modules_total": len(disabled),
"unloaded_modules": unloaded[:max_modules],
"unloaded_modules_total": len(unloaded),
},
"scoped_blocks": {
"bot_block_tasks_total": bot_block_total,
"group_block_tasks_total": group_block_total,
},
"semantics": {
"available_tasks": "management_display_mirror_not_runtime_whitelist",
"runtime_truth": [
"TaskInfo.status",
"TaskInfo.load_status",
"BotConsole.block_tasks",
"GroupConsole.block_task",
"GroupConsole.superuser_block_task",
],
},
}
@PriorityLifecycle.on_startup(priority=6)
async def _init_runtime_cache():
await RuntimeCacheSync.start()
+4
View File
@@ -12,6 +12,10 @@ T = TypeVar("T", bound=Model)
class DataAccess(Generic[T]):
"""数据访问兼容层,根据配置保留单点缓存读取和清理能力
边界说明:DataAccess 面向普通业务查询和低频管理链路。权限检查热路径
必须优先使用 RuntimeCache/AuthSnapshot,避免高并发消息处理时触发
DB/cache-aside 读放大。
新的高频运行态路径应优先使用 RuntimeCache 或 BoundedTTLCache。
这里不再把 filter/all/create/update_or_create 结果写入通用缓存,
create/update_or_create 只负责清理旧缓存,避免旧值残留。
+3
View File
@@ -5,6 +5,9 @@ from pydantic import BaseModel
# 数据库操作超时设置(秒)
DB_TIMEOUT_SECONDS = 3.0
# 启动期自动补齐字段/索引可能需要等待数据库锁或扫描较大的表,单独放宽超时
DB_SCHEMA_GUARD_TIMEOUT_SECONDS = 30.0
# 性能监控阈值(秒)
SLOW_QUERY_THRESHOLD = 0.5
+13 -3
View File
@@ -13,7 +13,7 @@ from tortoise.exceptions import OperationalError
from zhenxun.services.log import logger
from .config import DB_TIMEOUT_SECONDS, LOG_COMMAND
from .config import DB_SCHEMA_GUARD_TIMEOUT_SECONDS, LOG_COMMAND
Dialect = Literal["sqlite", "postgres", "mysql", "unknown"]
@@ -466,7 +466,7 @@ async def repair_safe_schema_drift() -> SchemaGuardResult:
try:
await asyncio.wait_for(
connection.execute_query_dict(sql),
timeout=DB_TIMEOUT_SECONDS,
timeout=DB_SCHEMA_GUARD_TIMEOUT_SECONDS,
)
columns[source] = ColumnInfo(name=source, data_type="")
result.repaired_columns += 1
@@ -521,7 +521,7 @@ async def repair_safe_schema_drift() -> SchemaGuardResult:
try:
await asyncio.wait_for(
connection.execute_query_dict(sql),
timeout=DB_TIMEOUT_SECONDS,
timeout=DB_SCHEMA_GUARD_TIMEOUT_SECONDS,
)
existing_indexes.add(index_columns)
result.repaired_indexes += 1
@@ -529,6 +529,16 @@ async def repair_safe_schema_drift() -> SchemaGuardResult:
f"SchemaGuard 已补齐索引: {table}.{index_columns}",
LOG_COMMAND,
)
except TimeoutError as exc:
result.warnings += 1
result.skipped_indexes += 1
logger.warning(
"SchemaGuard 补齐索引超时,已跳过: "
f"{table}.{index_columns} "
f"({DB_SCHEMA_GUARD_TIMEOUT_SECONDS}s)",
LOG_COMMAND,
e=exc,
)
except OperationalError as exc:
err = str(exc).lower()
if any(
-3
View File
@@ -140,9 +140,6 @@ def register_runtime_bootstrap(_driver) -> None:
global _thread_executor
await _stop_launcher_watchdog()
await stop_send_queue()
from zhenxun.models._bot_message_buffer import stop_bot_message_store_buffer
await stop_bot_message_store_buffer()
await stop_memory_governor()
executor = _thread_executor
_thread_executor = None
+41 -13
View File
@@ -27,6 +27,7 @@ from zhenxun.services.log import logger
from zhenxun.services.message_load import should_pause_tasks
from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.decorator.retry import Retry
from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.pydantic_compat import parse_as
from .repository import ScheduleRepository
@@ -37,6 +38,37 @@ SCHEDULE_CONCURRENCY_KEY = "all_groups_concurrency_limit"
_LAST_PRESSURE_SKIP = 0.0
def _resolve_scheduler_bot(bot_id: str | None, log_target: str) -> Bot | None:
if bot_id:
try:
return nonebot.get_bot(bot_id)
except KeyError:
logger.warning(f"{log_target} 需要的 Bot {bot_id} 不在线,本次执行跳过。")
return None
bots = list(nonebot.get_bots().values())
if not bots:
logger.warning(f"{log_target} 当前没有可用 Bot,本次执行跳过。")
return None
if len(bots) == 1:
return bots[0]
qq_client_bots = [
bot for bot in bots if PlatformUtils.get_platform_scope(bot) == "qq_client"
]
if len(qq_client_bots) == 1:
bot = qq_client_bots[0]
logger.warning(
f"{log_target} 未指定 Bot,多 Bot 在线,自动选择 OneBot {bot.self_id}。"
)
return bot
logger.warning(
f"{log_target} 未指定 Bot 且多 Bot 在线,无法安全选择," "本次执行跳过。"
)
return None
class APSchedulerAdapter:
"""封装对 APScheduler 的操作"""
@@ -343,7 +375,12 @@ async def _execute_job(
return
try:
bot = nonebot.get_bot()
bot = _resolve_scheduler_bot(
context_override.bot_id,
f"临时任务 {plugin_name}",
)
if bot is None:
return
logger.info(f"开始执行临时任务: {plugin_name}")
injected_params = {"context": context_override}
state: T_State = {ScheduleContext: context_override}
@@ -380,18 +417,9 @@ async def _execute_job(
logger.warning(f"定时任务 {schedule_id} 不存在或已禁用,跳过执行。")
return
try:
bot = (
nonebot.get_bot(schedule.bot_id)
if schedule.bot_id
else nonebot.get_bot()
)
except (KeyError, ValueError):
logger.warning(
f"任务 {schedule_id} 需要的 Bot {schedule.bot_id} "
f"不在线,本次执行跳过。"
)
raise
bot = _resolve_scheduler_bot(schedule.bot_id, f"任务 {schedule_id}")
if bot is None:
return
resolver = scheduler_manager._target_resolvers.get(schedule.target_type)
if not resolver:
+142 -1
View File
@@ -1,6 +1,7 @@
import asyncio
from collections.abc import Awaitable, Callable
import contextlib
import importlib
from typing import Any, cast
from nonebot.adapters import Bot, Event
@@ -10,6 +11,9 @@ from nonebot.log import logger
_PATCHED = False
_ORIGINAL_FETCH: Callable[..., Awaitable[Any]] | None = None
_ORIGINAL_ONEBOT11_GROUP_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None
_ORIGINAL_QQ_C2C_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None
_ORIGINAL_QQ_GROUP_AT_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None
_ORIGINAL_QQ_GUILD_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None
def _sender_value(sender: Any, key: str, default: Any = None) -> Any:
@@ -81,6 +85,88 @@ async def _fast_onebot11_group_message(bot: Bot, event: Event) -> dict[str, Any]
}
def _qq_bot_app_id(bot: Bot) -> str:
bot_info = getattr(bot, "bot_info", None)
app_id = getattr(bot_info, "id", None)
return str(app_id or getattr(bot, "self_id", ""))
async def _fast_qq_c2c_message(bot: Bot, event: Event) -> dict[str, Any]:
"""Build Uninfo session for QQ official C2C messages from event fields."""
author = _event_value(event, "author")
user_id = str(
_sender_value(author, "user_openid")
or _sender_value(author, "id")
or _event_value(event, "user_id", "")
)
username = str(_sender_value(author, "username", "") or "")
return {
"user_id": user_id,
"name": username,
"nickname": username,
"avatar": f"https://q.qlogo.cn/qqapp/{_qq_bot_app_id(bot)}/{user_id}/100",
}
async def _fast_qq_group_at_message(bot: Bot, event: Event) -> dict[str, Any]:
"""Build Uninfo session for QQ official group-at messages from event fields."""
author = _event_value(event, "author")
user_id = str(
_sender_value(author, "member_openid")
or _sender_value(author, "id")
or _event_value(event, "user_id", "")
)
username = str(_sender_value(author, "username", "") or "")
group_id = str(
_event_value(event, "group_openid") or _event_value(event, "group_id") or ""
)
return {
"user_id": user_id,
"name": username,
"nickname": username,
"avatar": f"https://q.qlogo.cn/qqapp/{_qq_bot_app_id(bot)}/{user_id}/100",
"group_id": group_id,
}
async def _fast_qq_guild_message(bot: Bot, event: Event) -> dict[str, Any]:
"""Build Uninfo session for QQ official guild/channel messages locally.
nonebot-plugin-uninfo enriches guild messages through remote guild/channel
APIs. Runtime auth only needs stable scene/user ids, so avoid remote calls
during matcher fanout.
"""
author = _event_value(event, "author")
member = _event_value(event, "member")
guild_id = str(_event_value(event, "guild_id", "") or "")
channel_id = str(_event_value(event, "channel_id", "") or "")
user_id = str(_sender_value(author, "id", "") or "")
nickname = str(_sender_value(member, "nick", "") or "")
username = str(_sender_value(author, "username", "") or "")
base: dict[str, Any] = {
"user_id": user_id,
"name": username,
"nickname": nickname or username,
"avatar": _sender_value(author, "avatar"),
"guild_id": guild_id,
"channel_id": channel_id,
"guild_name": "",
"guild_avatar": None,
"channel_name": "",
"channel_type": -1,
}
roles = _sender_value(member, "roles")
if roles is not None:
base["roles"] = roles
joined_at = _sender_value(member, "joined_at")
if joined_at is not None:
base["joined_at"] = joined_at
return base
async def _singleflight_fetch(self: Any, bot: Bot, event: Event) -> Any:
original = _ORIGINAL_FETCH
if original is None:
@@ -114,6 +200,8 @@ async def _singleflight_fetch(self: Any, bot: Bot, event: Event) -> Any:
def apply_uninfo_onebot11_patch() -> None:
global _ORIGINAL_FETCH, _ORIGINAL_ONEBOT11_GROUP_MESSAGE, _PATCHED
global _ORIGINAL_QQ_C2C_MESSAGE, _ORIGINAL_QQ_GROUP_AT_MESSAGE
global _ORIGINAL_QQ_GUILD_MESSAGE
if _PATCHED:
return
@@ -129,6 +217,59 @@ def apply_uninfo_onebot11_patch() -> None:
setattr(_fast_onebot11_group_message, "__zhenxun_fast_onebot11__", True)
fetcher.endpoint[GroupMessageEvent] = _fast_onebot11_group_message
with contextlib.suppress(Exception):
qq_event_module = importlib.import_module("nonebot.adapters.qq.event")
AtMessageCreateEvent = getattr(qq_event_module, "AtMessageCreateEvent")
C2CMessageCreateEvent = getattr(qq_event_module, "C2CMessageCreateEvent")
DirectMessageCreateEvent = getattr(qq_event_module, "DirectMessageCreateEvent")
GroupAtMessageCreateEvent = getattr(
qq_event_module,
"GroupAtMessageCreateEvent",
)
GroupMessageCreateEvent = getattr(
qq_event_module,
"GroupMessageCreateEvent",
)
MessageCreateEvent = getattr(qq_event_module, "MessageCreateEvent")
from nonebot_plugin_uninfo.adapters.qq.main import fetcher as qq_fetcher
original_c2c = qq_fetcher.endpoint.get(C2CMessageCreateEvent)
if not getattr(original_c2c, "__zhenxun_fast_qq__", False):
_ORIGINAL_QQ_C2C_MESSAGE = cast(
Callable[..., Awaitable[dict[str, Any]]] | None,
original_c2c,
)
setattr(_fast_qq_c2c_message, "__zhenxun_fast_qq__", True)
qq_fetcher.endpoint[C2CMessageCreateEvent] = _fast_qq_c2c_message
for event_type in (GroupMessageCreateEvent, GroupAtMessageCreateEvent):
original_group_at = qq_fetcher.endpoint.get(event_type)
if getattr(original_group_at, "__zhenxun_fast_qq__", False):
continue
if _ORIGINAL_QQ_GROUP_AT_MESSAGE is None and original_group_at is not None:
_ORIGINAL_QQ_GROUP_AT_MESSAGE = cast(
Callable[..., Awaitable[dict[str, Any]]],
original_group_at,
)
setattr(_fast_qq_group_at_message, "__zhenxun_fast_qq__", True)
qq_fetcher.endpoint[event_type] = _fast_qq_group_at_message
for event_type in (
MessageCreateEvent,
AtMessageCreateEvent,
DirectMessageCreateEvent,
):
original_guild = qq_fetcher.endpoint.get(event_type)
if getattr(original_guild, "__zhenxun_fast_qq__", False):
continue
if _ORIGINAL_QQ_GUILD_MESSAGE is None and original_guild is not None:
_ORIGINAL_QQ_GUILD_MESSAGE = cast(
Callable[..., Awaitable[dict[str, Any]]],
original_guild,
)
setattr(_fast_qq_guild_message, "__zhenxun_fast_qq__", True)
qq_fetcher.endpoint[event_type] = _fast_qq_guild_message
try:
from nonebot_plugin_uninfo.fetch import InfoFetcher
except Exception as e:
@@ -146,4 +287,4 @@ def apply_uninfo_onebot11_patch() -> None:
setattr(_singleflight_fetch, "__zhenxun_singleflight__", True)
setattr(InfoFetcher, "fetch", _singleflight_fetch)
_PATCHED = True
logger.debug("Uninfo OneBot11 fast fetch and singleflight patch applied")
logger.debug("Uninfo fast fetch and singleflight patch applied")
+12 -10
View File
@@ -23,29 +23,31 @@ class CommonUtils:
async def task_is_block(
cls, session: Uninfo | Bot, module: str, group_id: str | None = None
) -> bool:
"""判断被动技能是否可以发送
"""判断被动技能是否被阻断。
运行真源固定为 TaskInfo.status/load_status、BotConsole.block_tasks、
GroupConsole.block_task/superuser_block_task,以及 bot/group ban 状态。
BotConsole.available_tasks 只用于管理展示,不作为运行白名单。
参数:
module: 被动技能模块名
group_id: 群组id
返回:
bool: 是否可以发送
bool: True 表示被动技能应被阻断,False 表示允许继续执行
"""
if isinstance(session, Bot):
if interface := get_interface(session):
info = interface.basic_info()
if info["scope"] == SupportScope.qq_api:
logger.info("q官bot放弃所有被动技能发言...")
"""q官bot放弃所有被动技能发言"""
return False
if session.scene == SupportScope.qq_api:
"""q官bot放弃所有被动技能发言"""
logger.info("q官bot放弃所有被动技能发言...")
return False
logger.debug("q官bot放弃所有被动技能发言...")
return True
if isinstance(session, Session) and session.scope == SupportScope.qq_api:
logger.debug("q官bot放弃所有被动技能发言...")
return True
if not group_id and isinstance(session, Session):
group_id = session.group.id if session.group else None
if await TaskInfoMemoryCache.is_disabled(module):
if await TaskInfoMemoryCache.is_runtime_disabled(module):
"""被动全局状态"""
return True
bot_snapshot = await BotMemoryCache.get(session.self_id)
-5
View File
@@ -13,11 +13,6 @@ class PriorityLifecycleType(StrEnum):
"""关闭"""
class BotSentType(StrEnum):
GROUP = "GROUP"
PRIVATE = "PRIVATE"
class BankHandleType(StrEnum):
DEPOSIT = "DEPOSIT"
"""存款"""
+70 -17
View File
@@ -1,5 +1,6 @@
import asyncio
from collections.abc import Awaitable, Callable
import contextlib
import random
from typing import cast
@@ -37,6 +38,23 @@ def _adapter_name(bot: Bot) -> str:
return adapter.__class__.__name__.lower()
def _scope_name(scope: object) -> str:
"""Normalize uninfo/alconna scope values without changing legacy platform."""
if scope is None:
return ""
raw = getattr(scope, "name", None) or getattr(scope, "value", scope)
text = str(raw or "").strip().lower()
if not text:
text = str(getattr(scope, "name", "") or "").strip().lower()
text = text.replace("-", "_").replace(" ", "_")
compact = "".join(ch for ch in text if ch.isalnum())
if compact.endswith("qqclient"):
return "qq_client"
if compact.endswith("qqapi"):
return "qq_api"
return text
class UserData(BaseModel):
name: str
"""昵称"""
@@ -57,6 +75,31 @@ class UserData(BaseModel):
class PlatformUtils:
@classmethod
def _resolve_unique_qq_client_bot(cls, log_cmd: str | None = None) -> Bot | None:
bots = list(nonebot.get_bots().values())
if not bots:
logger.warning("当前没有可用的 OneBot 协议端 Bot,已跳过。", log_cmd)
return None
qq_client_bots = [
bot for bot in bots if cls.get_platform_scope(bot) == "qq_client"
]
if len(qq_client_bots) == 1:
bot = qq_client_bots[0]
if len(bots) > 1:
logger.warning(
f"多 Bot 在线且未指定 Bot,自动选择 OneBot {bot.self_id}。",
log_cmd,
)
return bot
if not qq_client_bots:
logger.warning("未找到 OneBot 协议端 Bot,已跳过。", log_cmd)
else:
logger.warning(
"存在多个 OneBot 协议端 Bot,无法安全选择,已跳过。", log_cmd
)
return None
@classmethod
def is_qbot(cls, session: Uninfo | Bot) -> bool:
"""判断bot是否为qq官bot
@@ -68,7 +111,11 @@ class PlatformUtils:
bool: 是否为官bot
"""
if isinstance(session, Bot):
if cls.get_platform_scope(session) == "qq_api":
return True
return bool(BotConfig.get_qbot_uid(session.self_id))
if cls.get_platform_scope(session) == "qq_api":
return True
if BotConfig.get_qbot_uid(session.self_id):
return True
return session.scope == SupportScope.qq_api
@@ -83,7 +130,7 @@ class PlatformUtils:
group_id: 群组id
duration: 禁言时长(分钟)
"""
if cls.get_platform(bot) == "qq":
if cls.get_platform_scope(bot) == "qq_client":
await bot.set_group_ban(
group_id=int(group_id),
user_id=int(user_id),
@@ -111,7 +158,9 @@ class PlatformUtils:
Receipt | None: Receipt
"""
if not bot:
bot = nonebot.get_bot()
bot = cls._resolve_unique_qq_client_bot("PlatformUtils:send_superuser")
if bot is None:
return []
superuser_ids = []
if superuser_id:
superuser_ids.append(superuser_id)
@@ -397,6 +446,11 @@ class PlatformUtils:
def get_platform_scope(cls, t: Bot | Uninfo | object) -> str:
"""获取细粒度平台作用域,不改变旧 get_platform 返回值。"""
if isinstance(t, Bot):
if interface := get_interface(t):
with contextlib.suppress(Exception):
scope = _scope_name(interface.basic_info().get("scope"))
if scope:
return scope
adapter_name = _adapter_name(t)
if "onebot" in adapter_name:
return "qq_client"
@@ -406,17 +460,13 @@ class PlatformUtils:
return "qq_api"
return adapter_name or cls.get_platform(t)
scope = str(getattr(t, "scope", "") or "").lower()
scope = _scope_name(getattr(t, "scope", "") or "")
if not scope:
basic = getattr(t, "basic", None)
if isinstance(basic, dict):
scope = str(basic.get("scope") or "").lower()
if "qq_client" in scope:
return "qq_client"
if "qq_api" in scope:
return "qq_api"
if scope.startswith("qq"):
return "qq"
scope = _scope_name(basic.get("scope"))
if scope:
return scope
adapter = getattr(t, "adapter", None)
if adapter is not None:
@@ -499,6 +549,9 @@ class PlatformUtils:
返回:
int: 更新个数
"""
if cls.get_platform_scope(bot) == "qq_api":
logger.warning("QQ 官方适配器不支持旧好友同步,已跳过。", "更新好友信息")
return 0
create_list = []
friend_list, platform = await cls.get_friend_list(bot)
if friend_list:
@@ -521,6 +574,11 @@ class PlatformUtils:
返回:
list[FriendUser]: 好友列表
"""
if cls.get_platform_scope(bot) == "qq_api":
logger.warning(
"QQ 官方适配器不支持旧好友列表查询,已返回空列表。", "好友列表"
)
return [], cls.get_platform(bot)
if interface := get_interface(bot):
user_list = await interface.get_users()
return [
@@ -602,14 +660,9 @@ class BroadcastEngine:
except KeyError:
logger.warning(f"Bot:{i} 对象未连接或不存在", log_cmd)
if not self.bot_list:
try:
bot = nonebot.get_bot()
bot = PlatformUtils._resolve_unique_qq_client_bot(log_cmd)
if bot is not None:
self.bot_list.append(bot)
logger.warning(
f"广播任务未传入Bot对象,使用默认Bot {bot.self_id}", log_cmd
)
except Exception as e:
raise ValueError("当前没有可用的Bot对象...", log_cmd) from e
async def call_check(self, bot: Bot, group_id: str) -> bool:
"""运行发送检测函数