diff --git a/.env.example b/.env.example index 45b3b0be..f8a9c81e 100644 --- a/.env.example +++ b/.env.example @@ -1,4 +1,4 @@ -SUPERUSERS=[""] +SUPERUSERS=[""] COMMAND_START=[""] @@ -31,7 +31,7 @@ QBOT_ID_DATA = '{ DB_URL = "" # NONE: 不使用缓存, MEMORY: 使用内存缓存, REDIS: 使用Redis缓存 -CACHE_MODE = NONE +CACHE_MODE = MEMORY # REDIS配置,使用REDIS替换Cache内存缓存 # REDIS地址 @@ -64,6 +64,9 @@ PORT = 8080 # 第三方插件路径,如果多个目录用, 隔开 # EXT_PATH=[""] +# qq adapter load = True +QQ_ADAPTER_LOAD=False + # kook adapter toekn # kaiheila_bots =[{"token": ""}] @@ -94,3 +97,4 @@ PORT = 8080 # application_commands的{"*": ["*"]}代表将全部应用命令注册为全局应用命令 # {"admin": ["123", "456"]}则代表将admin命令注册为id是123、456服务器的局部命令,其余命令不注册 + diff --git a/scripts/auth_stage3_perf_summary.py b/scripts/auth_stage3_perf_summary.py new file mode 100644 index 00000000..8e0cf55a --- /dev/null +++ b/scripts/auth_stage3_perf_summary.py @@ -0,0 +1,83 @@ +from __future__ import annotations + +import argparse +import json +from pathlib import Path +import sys +from typing import Any + + +def _load_json(path: Path) -> dict[str, Any]: + return json.loads(path.read_text(encoding="utf-8")) + + +def _num(value: Any) -> float: + try: + return float(value) + except (TypeError, ValueError): + return 0.0 + + +def _safe_div(left: float, right: float) -> float: + return round(left / right, 4) if right else 0.0 + + +def _extract(path: Path) -> dict[str, Any]: + payload = _load_json(path) + summary = payload.get("summary") or {} + trace = summary.get("db_trace") or {} + events = _num(summary.get("events_sent_total")) + commands = _num(summary.get("commands_sent_total")) + return { + "path": str(path), + "status": payload.get("status"), + "elapsed_seconds": payload.get("elapsed_seconds"), + "events": int(events), + "commands": int(commands), + "throughput_eps": summary.get("throughput_events_per_sec"), + "command_success_rate": summary.get("command_success_rate"), + "latency_avg_ms": summary.get("latency_avg_ms"), + "latency_p50_ms": summary.get("latency_p50_ms"), + "latency_p95_ms": summary.get("latency_p95_ms"), + "latency_p99_ms": summary.get("latency_p99_ms"), + "db_timeouts": summary.get("db_timeouts"), + "db_slow_queries": summary.get("db_slow_queries"), + "chat_history_failures": summary.get("chat_history_failures"), + "statistics_flush_failures": summary.get("statistics_flush_failures"), + "db_calls": int(_num(trace.get("calls"))), + "db_reads": int(_num(trace.get("reads"))), + "db_writes": int(_num(trace.get("writes"))), + "db_scripts": int(_num(trace.get("scripts"))), + "db_calls_per_event": _safe_div(_num(trace.get("calls")), events), + "db_reads_per_event": _safe_div(_num(trace.get("reads")), events), + "db_writes_per_event": _safe_div(_num(trace.get("writes")), events), + "db_calls_per_command": _safe_div(_num(trace.get("calls")), commands), + "db_writes_per_command": _safe_div(_num(trace.get("writes")), commands), + "db_avg_elapsed_ms": trace.get("avg_elapsed_ms"), + "db_avg_wait_ms": trace.get("avg_wait_ms"), + "db_max_elapsed_ms": trace.get("max_elapsed_ms"), + "db_max_wait_ms": trace.get("max_wait_ms"), + "db_max_active": trace.get("max_active"), + "db_max_waiting": trace.get("max_waiting"), + "db_connection_creates": trace.get("connection_creates"), + "db_top_tables": trace.get("top_tables", [])[:12], + } + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("reports", nargs="+") + parser.add_argument("--output") + args = parser.parse_args() + rows = [_extract(Path(item).resolve()) for item in args.reports] + payload = {"reports": rows} + text = json.dumps(payload, ensure_ascii=False, indent=2) + if args.output: + output = Path(args.output).resolve() + output.parent.mkdir(parents=True, exist_ok=True) + output.write_text(text, encoding="utf-8") + sys.stdout.write(text + "\n") + + +if __name__ == "__main__": + main() diff --git a/zhenxun/builtin_plugins/admin/group_member_update/_data_source.py b/zhenxun/builtin_plugins/admin/group_member_update/_data_source.py index eeac663a..bb1c12f5 100644 --- a/zhenxun/builtin_plugins/admin/group_member_update/_data_source.py +++ b/zhenxun/builtin_plugins/admin/group_member_update/_data_source.py @@ -9,6 +9,10 @@ from zhenxun.configs.config import Config from zhenxun.models.group_console import GroupConsole from zhenxun.models.group_member_info import GroupInfoUser from zhenxun.models.level_user import LevelUser +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 @@ -181,4 +185,12 @@ class MemberUpdateManage: group_id=group_id, platform="qq", ) + changed_user_ids = ( + {user.user_id for user in data_list[0]} + | {user.user_id for user in data_list[1]} + | delete_member_ids + ) + if data_list[0] or data_list[1] or data_list[2] or delete_member_ids: + await invalidate_group_members(group_id, changed_user_ids) + await invalidate_member_names(changed_user_ids) return "群组成员信息更新完成!" diff --git a/zhenxun/builtin_plugins/chat_history/chat_message_handle.py b/zhenxun/builtin_plugins/chat_history/chat_message_handle.py index fa692de1..acbf2c2f 100644 --- a/zhenxun/builtin_plugins/chat_history/chat_message_handle.py +++ b/zhenxun/builtin_plugins/chat_history/chat_message_handle.py @@ -19,8 +19,9 @@ from zhenxun import ui from zhenxun.configs.config import Config from zhenxun.configs.utils import Command, PluginExtraData, RegisterConfig from zhenxun.models.chat_history import ChatHistory -from zhenxun.models.group_member_info import GroupInfoUser +from zhenxun.models.friend_user import FriendUser from zhenxun.services import avatar_service +from zhenxun.services.hot_query_cache import get_group_member_map, get_member_names from zhenxun.services.log import logger from zhenxun.ui.models import ImageCell, TextCell from zhenxun.utils.enum import PluginType @@ -118,7 +119,8 @@ async def _( show_quit_member = Config.get_config("chat_history", "SHOW_QUIT_MEMBER", True) fetch_count = count.result - if not show_quit_member: + has_group_context = bool(group_id) + if has_group_context and not show_quit_member: fetch_count = count.result * 2 raw_rank_data = await ChatHistory.get_group_msg_rank( @@ -128,27 +130,37 @@ async def _( if raw_rank_data: rank_data = cast(list[tuple[str, int]], raw_rank_data) rows_data = [] - platform = "qq" + platform = getattr(session, "platform", None) or "qq" user_ids_in_rank = [str(uid) for uid, _ in rank_data] - users_in_group_query = GroupInfoUser.filter( - user_id__in=user_ids_in_rank, group_id=group_id - ) - users_in_group = {u.user_id: u for u in await users_in_group_query} + users_in_group = {} + user_names: dict[str, str] = {} + if has_group_context: + users_in_group = await get_group_member_map(group_id, user_ids_in_rank) + else: + friend_users = await FriendUser.filter( + user_id__in=user_ids_in_rank + ).values_list("user_id", "user_name") + user_names.update(dict(friend_users)) + group_user_names = await get_member_names(user_ids_in_rank) + for user_id, user_name in group_user_names.items(): + if user_name and user_id not in user_names: + user_names[user_id] = user_name for idx, (uid, num) in enumerate(rank_data): if len(rows_data) >= count.result: break uid_str = str(uid) - user_in_group = users_in_group.get(uid_str) - - if not user_in_group and not show_quit_member: - continue - - user_name = ( - user_in_group.user_name if user_in_group else f"{uid_str}(已退群)" - ) + if has_group_context: + user_in_group = users_in_group.get(uid_str) + if not user_in_group and not show_quit_member: + continue + user_name = ( + user_in_group.user_name if user_in_group else f"{uid_str}(已退群)" + ) + else: + user_name = user_names.get(uid_str) or uid_str avatar_path = await avatar_service.get_avatar_path(platform, uid_str) diff --git a/zhenxun/builtin_plugins/hooks/auth/auth_limit.py b/zhenxun/builtin_plugins/hooks/auth/auth_limit.py index e16eb0f2..b18d15f5 100644 --- a/zhenxun/builtin_plugins/hooks/auth/auth_limit.py +++ b/zhenxun/builtin_plugins/hooks/auth/auth_limit.py @@ -1,4 +1,6 @@ import asyncio +from collections.abc import Callable +from dataclasses import dataclass, field import time from typing import Any, ClassVar @@ -46,6 +48,26 @@ class Limit(BaseModel): arbitrary_types_allowed = True +@dataclass(slots=True) +class LimitReservation: + module: str + releases: list[Callable[[], None]] = field(default_factory=list) + should_auto_unblock: bool = False + active: bool = True + + def commit(self) -> None: + self.active = False + self.releases.clear() + + def release(self) -> None: + if not self.active: + return + for release in reversed(self.releases): + release() + self.active = False + self.releases.clear() + + def _limit_notice_key( limit: PluginLimit | PluginLimitSnapshot, user_id: str, @@ -247,14 +269,9 @@ class LimitManager: for limit in limits: cls.add_limit(limit) - # 检查各种限制 try: - if limit_model := cls.cd_limit.get(module): - await cls.__check(limit_model, user_id, group_id, channel_id) - if limit_model := cls.block_limit.get(module): - await cls.__check(limit_model, user_id, group_id, channel_id) - if limit_model := cls.count_limit.get(module): - await cls.__check(limit_model, user_id, group_id, channel_id) + reservation = await cls.reserve(module, user_id, group_id, channel_id) + reservation.commit() finally: # 记录总执行时间 elapsed = time.time() - start_time @@ -267,13 +284,53 @@ class LimitManager: ) @classmethod - async def __check( + async def reserve( + cls, + module: str, + user_id: str, + group_id: str | None, + channel_id: str | None, + ) -> LimitReservation: + """检查并预留限制状态;调用方失败时可 release 回滚内存限制。""" + if ( + time.time() - cls.last_update_time > cls.update_interval + and not cls.is_updating + ): + asyncio.create_task(cls.update_limits()) # noqa: RUF006 + + if module not in cls.add_module: + limits = await cls.get_module_limits(module) + for limit in limits: + cls.add_limit(limit) + + reservation = LimitReservation(module=module) + try: + if limit_model := cls.cd_limit.get(module): + reservation.releases.append( + await cls.__reserve(limit_model, user_id, group_id, channel_id) + ) + if limit_model := cls.block_limit.get(module): + reservation.should_auto_unblock = True + reservation.releases.append( + await cls.__reserve(limit_model, user_id, group_id, channel_id) + ) + if limit_model := cls.count_limit.get(module): + reservation.releases.append( + await cls.__reserve(limit_model, user_id, group_id, channel_id) + ) + except Exception: + reservation.release() + raise + return reservation + + @classmethod + async def __reserve( cls, limit_model: Limit | None, user_id: str, group_id: str | None, channel_id: str | None, - ): + ) -> Callable[[], None]: """检测限制 参数: @@ -286,7 +343,7 @@ class LimitManager: IgnoredException: IgnoredException """ if not limit_model: - return + return lambda: None limit = limit_model.limit limiter = limit_model.limiter is_limit = ( @@ -317,12 +374,40 @@ class LimitManager: group_id=group_id, ) if isinstance(limiter, FreqLimiter): + had_next_time = key_type in limiter.next_time + old_next_time = limiter.next_time.get(key_type, 0.0) limiter.start_cd(key_type) + + def release_freq() -> None: + if had_next_time: + limiter.next_time[key_type] = old_next_time + else: + limiter.next_time.pop(key_type, None) + + return release_freq if isinstance(limiter, UserBlockLimiter): + old_flag = limiter.flag_data.get(key_type, False) + old_time = limiter.time.get(key_type, 0.0) limiter.set_true(key_type) + + def release_block() -> None: + limiter.flag_data[key_type] = old_flag + if old_time: + limiter.time[key_type] = old_time + else: + limiter.time.pop(key_type, None) + + return release_block if isinstance(limiter, CountLimiter): + old_count = limiter.count.get(key_type, 0) limiter.increase(key_type) + def release_count() -> None: + limiter.count[key_type] = old_count + + return release_count + return lambda: None + async def auth_limit( plugin: PluginInfo, @@ -343,11 +428,42 @@ async def auth_limit( entity = get_entity_ids(session) try: await asyncio.wait_for( - LimitManager.check( - plugin.module, entity.user_id, entity.group_id, entity.channel_id - ), + _reserve_and_commit_limit(plugin.module, entity), timeout=DB_TIMEOUT_SECONDS * 2, # 给予更长的超时时间 ) except asyncio.TimeoutError: logger.error(f"检查插件限制超时: {plugin.module}", LOGGER_COMMAND) # 超时时不抛出异常,允许继续执行 + + +async def reserve_auth_limit( + plugin: PluginInfo, + session: Uninfo, + *, + context: PermissionContext | None = None, + entity: EntityIDs | None = None, +) -> LimitReservation: + del session + if context is not None: + entity = context.entity + if entity is None: + raise RuntimeError("reserve_auth_limit requires entity or context") + return await LimitManager.reserve( + plugin.module, + entity.user_id, + entity.group_id, + entity.channel_id, + ) + + +async def _reserve_and_commit_limit( + module: str, + entity: EntityIDs, +) -> None: + reservation = await LimitManager.reserve( + module, + entity.user_id, + entity.group_id, + entity.channel_id, + ) + reservation.commit() diff --git a/zhenxun/builtin_plugins/hooks/auth/context.py b/zhenxun/builtin_plugins/hooks/auth/context.py index 3e270ffb..e491398a 100644 --- a/zhenxun/builtin_plugins/hooks/auth/context.py +++ b/zhenxun/builtin_plugins/hooks/auth/context.py @@ -3,7 +3,7 @@ from __future__ import annotations import asyncio import contextlib from dataclasses import dataclass, field -from typing import Any +from typing import TYPE_CHECKING, Any from nonebot.adapters import Bot, Event from nonebot_plugin_alconna import UniMsg @@ -31,6 +31,9 @@ EVENT_CACHE = ( else None ) +if TYPE_CHECKING: + from zhenxun.builtin_plugins.hooks.auth_side_effect import SideEffectCommit + @dataclass class EventContext: @@ -62,6 +65,7 @@ class EventContext: class PermissionSideEffectCache: auth_results: dict[str, tuple[bool, str | None]] = field(default_factory=dict) module_locks: dict[str, asyncio.Lock] = field(default_factory=dict) + commits: dict[str, "SideEffectCommit"] = field(default_factory=dict) def lock_for(self, module: str) -> asyncio.Lock: lock = self.module_locks.get(module) diff --git a/zhenxun/builtin_plugins/hooks/auth_activation.py b/zhenxun/builtin_plugins/hooks/auth_activation.py new file mode 100644 index 00000000..0255ccb9 --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth_activation.py @@ -0,0 +1,1214 @@ +from __future__ import annotations + +from collections.abc import Iterable +import contextlib +from dataclasses import dataclass, field +import html +import re +from typing import Any, Literal +import weakref + +from nonebot.matcher import Matcher + +ActivationDecision = Literal["match", "miss", "unknown"] +ActivationLane = Literal[ + "command_exact", + "command_shortcut", + "command_regex", + "system", + "passive_light", + "passive_db", + "passive_http", + "passive_ai", + "passive_render", +] + + +@dataclass(frozen=True, slots=True) +class ActivationRuleDescriptor: + kind: str + value: object | None = None + flags: int = 0 + ignorecase: bool = False + deterministic_text: bool = False + command_like: bool = False + + +@dataclass(slots=True) +class HandlerDescriptor: + matcher: type[Matcher] + module: str + matcher_type: str + priority: int + lane: ActivationLane + temp: bool = False + block: bool = False + command_like: bool = False + deterministic_text: bool = False + has_custom_rule: bool = False + commands: tuple[str, ...] = () + shortcuts: tuple[str, ...] | None = None + rules: tuple[ActivationRuleDescriptor, ...] = () + + +@dataclass(slots=True) +class ActivationContext: + event_type: str + event: object | None = None + plain_text: str = "" + raw_text: str = "" + to_me: bool = False + has_url: bool = False + has_image: bool = False + is_command_like: bool = False + route_modules: set[str] = field(default_factory=set) + ai_route_modules: set[str] = field(default_factory=set) + ai_route_heads: set[str] = field(default_factory=set) + + +@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. + """ + + 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._compiled = False + + @property + def compiled(self) -> bool: + return self._compiled + + def rebuild(self, matchers: dict[int, list[type[Matcher]]]) -> None: + self._by_priority.clear() + self._matcher_map.clear() + source_keys: set[tuple[int, tuple[int, ...]]] = set() + for priority, priority_matchers in matchers.items(): + descriptors = [ + self._build_descriptor(matcher, priority) + for matcher in priority_matchers + ] + self._by_priority[priority] = descriptors + for descriptor in descriptors: + self._matcher_map[descriptor.matcher] = descriptor + source_keys.add( + (int(priority), tuple(id(matcher) for matcher in priority_matchers)) + ) + self._source_keys = source_keys + self._compiled = True + + def ensure_fresh(self, matchers: dict[int, list[type[Matcher]]]) -> None: + source_keys = { + (int(priority), tuple(id(matcher) for matcher in items)) + for priority, items in matchers.items() + } + if not self._compiled or source_keys != self._source_keys: + self.rebuild(matchers) + + def descriptors_for_priority( + self, + priority: int, + priority_matchers: Iterable[type[Matcher]], + ) -> list[HandlerDescriptor]: + priority_matchers_list = list(priority_matchers) + descriptors = self._by_priority.get(priority) + if descriptors is not None and len(descriptors) == len(priority_matchers_list): + return descriptors + # Fallback for dynamic matcher list changes inside a priority bucket. + return [ + self._matcher_map.get(matcher) or self._build_descriptor(matcher, priority) + for matcher in priority_matchers_list + ] + + def descriptor_for(self, matcher: type[Matcher]) -> HandlerDescriptor | None: + return self._matcher_map.get(matcher) + + def select_priority( + self, + priority: int, + priority_matchers: list[type[Matcher]], + context: ActivationContext, + budget: dict[str, int], + ) -> ActivationResult: + result = ActivationResult(selected=[], total_descriptors=len(priority_matchers)) + descriptors = [ + self._matcher_map.get(matcher) or self._build_descriptor(matcher, priority) + for matcher in priority_matchers + ] + 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") + 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 + + def _select_descriptor( + self, + descriptor: HandlerDescriptor, + context: ActivationContext, + ) -> Literal["select", "miss", "fallback", "deterministic"]: + if descriptor.temp: + return "select" + matcher_type = descriptor.matcher_type + if matcher_type and matcher_type != context.event_type: + return "miss" + if context.event_type != "message": + return "miss" if descriptor.command_like else "select" + if descriptor.command_like: + return self._select_command_descriptor(descriptor, context) + 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": + return "deterministic" + if rule_match == "miss": + return "miss" + return "select" + + def _select_command_descriptor( + self, + descriptor: HandlerDescriptor, + context: ActivationContext, + ) -> Literal["select", "miss", "fallback", "deterministic"]: + texts = text_match_candidates( + context.plain_text, + context.raw_text, + context.event, + ) + if not texts: + return "select" + command_matched = False + if descriptor.commands: + if any( + matcher_command_matches(text, command) + for text in texts + for command in descriptor.commands + ): + command_matched = True + else: + shortcut_match = matcher_alconna_shortcut_matches_any( + descriptor.shortcuts, + texts, + ) + 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" + else: + shortcut_match = matcher_alconna_shortcut_matches_any( + descriptor.shortcuts, + texts, + ) + 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": + command_matched = True + elif rule_match == "miss": + return "miss" + elif ( + not (descriptor.commands or descriptor.shortcuts is not None) + and not command_matched + ): + return "fallback" + + if ( + context.ai_route_modules + and descriptor.module not in context.ai_route_modules + ): + if not matcher_matches_ai_route_heads(descriptor, context.ai_route_heads): + return "select" + + if command_matched: + return "deterministic" + return "select" + + def _build_descriptor( + self, + matcher: type[Matcher], + priority: int, + ) -> HandlerDescriptor: + rules = extract_matcher_rule_descriptors(matcher) + command_like = any(rule.command_like for rule in rules) + deterministic = any(rule.deterministic_text for rule in rules) + if hasattr(matcher, "command"): + command_like = True + commands = extract_matcher_command_literals(matcher) or () + shortcuts = extract_matcher_alconna_shortcuts(matcher) + if shortcuts is not None: + command_like = True + module = matcher_module_name(matcher) + lane = classify_lane( + matcher, + module=module, + command_like=command_like, + deterministic_text=deterministic, + shortcuts=shortcuts, + commands=commands, + ) + return HandlerDescriptor( + matcher=matcher, + module=module, + matcher_type=getattr(matcher, "type", "") or "", + priority=priority, + lane=lane, + temp=bool(getattr(matcher, "temp", False)), + block=bool(getattr(matcher, "block", False)), + command_like=command_like, + deterministic_text=deterministic, + has_custom_rule=matcher_has_custom_rule(matcher), + commands=commands, + shortcuts=shortcuts, + rules=rules, + ) + + +def matcher_module_name(matcher_cls: type[Matcher]) -> str: + module = getattr(matcher_cls, "plugin_name", "") or "" + if module: + return module + plugin = getattr(matcher_cls, "plugin", None) + if not plugin: + return "" + return (getattr(plugin, "name", "") or "").strip() + + +def classify_lane( + matcher_cls: type[Matcher], + *, + module: str, + command_like: bool, + deterministic_text: bool, + shortcuts: tuple[str, ...] | None, + commands: tuple[str, ...], +) -> ActivationLane: + if getattr(matcher_cls, "temp", False): + return "system" + if command_like: + if shortcuts: + return "command_shortcut" + has_regex_command = any( + is_regex_like_command_literal(item) for item in commands + ) + if deterministic_text or has_regex_command: + return "command_regex" + return "command_exact" + module_l = (module or "").casefold() + if any(hint in module_l for hint in PASSIVE_AI_HINTS): + return "passive_ai" + if any(hint in module_l for hint in PASSIVE_RENDER_HINTS): + return "passive_render" + if any(hint in module_l for hint in PASSIVE_HTTP_HINTS): + return "passive_http" + if matcher_has_custom_rule(matcher_cls): + return "passive_light" + if any(hint in module_l for hint in PASSIVE_DB_HINTS): + return "passive_db" + return "passive_light" + + +def matcher_has_custom_rule(matcher_cls: type[Matcher]) -> bool: + 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 + call_module = call.__class__.__module__ + if call_module.startswith("nonebot.rule") or call_module.startswith( + "nonebot_plugin_alconna.rule" + ): + continue + return True + return False + + +def matcher_is_command_like(matcher_cls: type[Matcher]) -> bool: + rules = extract_matcher_rule_descriptors(matcher_cls) + if any(rule.command_like for rule in rules): + return True + if hasattr(matcher_cls, "command"): + return True + return extract_matcher_alconna_shortcuts(matcher_cls) is not None + + +def matcher_has_deterministic_text_rule(matcher_cls: type[Matcher]) -> bool: + return any( + rule.deterministic_text + for rule in extract_matcher_rule_descriptors(matcher_cls) + ) + + +def classify_matcher_lane( + matcher_cls: type[Matcher], + *, + ai_route_modules: set[str] | None = None, +) -> ActivationLane: + module = matcher_module_name(matcher_cls) + if ai_route_modules and any( + module.casefold() == route_module.casefold() + for route_module in ai_route_modules + ): + return "passive_ai" + rules = extract_matcher_rule_descriptors(matcher_cls) + command_like = any(rule.command_like for rule in rules) + deterministic = any(rule.deterministic_text for rule in rules) + if hasattr(matcher_cls, "command"): + command_like = True + commands = extract_matcher_command_literals(matcher_cls) or () + shortcuts = extract_matcher_alconna_shortcuts(matcher_cls) + if shortcuts is not None: + command_like = True + return classify_lane( + matcher_cls, + module=module, + command_like=command_like, + deterministic_text=deterministic, + shortcuts=shortcuts, + commands=commands, + ) + + +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: + call = getattr(checker, "call", None) + if call is None: + continue + call_module = call.__class__.__module__ + call_name = call.__class__.__name__ + if call_module.startswith("nonebot.rule"): + if call_name == "CommandRule": + descriptors.append( + ActivationRuleDescriptor( + "command", + getattr(call, "cmds", ()), + command_like=True, + ) + ) + elif call_name == "ShellCommandRule": + descriptors.append( + ActivationRuleDescriptor( + "shell_command", + getattr(call, "cmds", ()), + command_like=True, + ) + ) + elif call_name == "RegexRule": + descriptors.append( + ActivationRuleDescriptor( + "regex", + getattr(call, "regex", ""), + flags=int(getattr(call, "flags", 0) or 0), + deterministic_text=True, + command_like=True, + ) + ) + elif call_name == "StartswithRule": + descriptors.append( + ActivationRuleDescriptor( + "startswith", + normalize_rule_string_tuple(getattr(call, "msg", ())), + ignorecase=bool(getattr(call, "ignorecase", False)), + deterministic_text=True, + command_like=True, + ) + ) + elif call_name == "EndswithRule": + descriptors.append( + ActivationRuleDescriptor( + "endswith", + normalize_rule_string_tuple(getattr(call, "msg", ())), + ignorecase=bool(getattr(call, "ignorecase", False)), + deterministic_text=True, + command_like=True, + ) + ) + elif call_name == "FullmatchRule": + descriptors.append( + ActivationRuleDescriptor( + "fullmatch", + normalize_rule_string_tuple(getattr(call, "msg", ())), + ignorecase=bool(getattr(call, "ignorecase", False)), + deterministic_text=True, + command_like=True, + ) + ) + elif call_name == "KeywordsRule": + descriptors.append( + ActivationRuleDescriptor( + "keywords", + normalize_rule_string_tuple(getattr(call, "keywords", ())), + deterministic_text=True, + command_like=True, + ) + ) + elif call_name == "IsTypeRule": + descriptors.append( + ActivationRuleDescriptor("is_type", getattr(call, "types", ())) + ) + elif call_name == "ToMeRule": + descriptors.append(ActivationRuleDescriptor("to_me")) + else: + descriptors.append(ActivationRuleDescriptor("custom")) + elif ( + call_module.startswith("nonebot_plugin_alconna.rule") + and call_name == "AlconnaRule" + ): + descriptors.append(ActivationRuleDescriptor("alconna", command_like=True)) + else: + descriptors.append(_custom_rule_descriptor(call)) + 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,) + if isinstance(value, list | tuple | set | frozenset): + return tuple(str(item) for item in value if str(item)) + return () + + +def text_match_candidates( + plain_text: str, + raw_text: str = "", + event: object | None = None, +) -> tuple[str, ...]: + """Return text variants visible to different matcher rule providers.""" + + candidates: list[str] = [] + + def add(text: object) -> None: + if not isinstance(text, str): + return + normalized = text.strip() + if normalized and normalized not in candidates: + candidates.append(normalized) + + unescaped = _unescape_message_text(normalized) + if unescaped and unescaped not in candidates: + candidates.append(unescaped) + + add(plain_text) + if event is not None: + with contextlib.suppress(Exception): + getter = getattr(event, "get_plaintext", None) + if callable(getter): + add(getter()) + add(raw_text) + return tuple(candidates) + + +def _unescape_message_text(text: str) -> str: + if not text: + return "" + unescaped = html.unescape(text) + return unescaped.replace("\\/", "/").replace("\\u002F", "/").replace("\\u002f", "/") + + +def matcher_rule_matches_text( + descriptors: tuple[ActivationRuleDescriptor, ...], + raw_text: str, + plain_text: str, + *, + event: object | None = None, + to_me: bool = False, +) -> ActivationDecision: + matched_any = False + saw_deterministic = False + saw_unknown = False + message_text = raw_text or plain_text + plain_candidates = text_match_candidates(plain_text, raw_text, event) + + for descriptor in descriptors: + kind = descriptor.kind + if kind == "regex": + saw_deterministic = True + pattern = str(descriptor.value or "") + if not pattern: + continue + try: + if re.search(pattern, message_text, descriptor.flags): + matched_any = True + else: + return "miss" + except re.error: + return "unknown" + elif kind == "startswith": + saw_deterministic = True + values = descriptor.value if isinstance(descriptor.value, tuple) else () + candidates = ( + tuple(item.casefold() for item in values) + if descriptor.ignorecase + else values + ) + texts = ( + tuple(item.casefold() for item in plain_candidates) + if descriptor.ignorecase + else plain_candidates + ) + if any( + text.startswith(item) for text in texts for item in candidates if item + ): + matched_any = True + else: + return "miss" + elif kind == "endswith": + saw_deterministic = True + values = descriptor.value if isinstance(descriptor.value, tuple) else () + candidates = ( + tuple(item.casefold() for item in values) + if descriptor.ignorecase + else values + ) + texts = ( + tuple(item.casefold() for item in plain_candidates) + if descriptor.ignorecase + else plain_candidates + ) + if any( + text.endswith(item) for text in texts for item in candidates if item + ): + matched_any = True + else: + return "miss" + elif kind == "fullmatch": + saw_deterministic = True + values = descriptor.value if isinstance(descriptor.value, tuple) else () + candidates = ( + tuple(item.casefold() for item in values) + if descriptor.ignorecase + else values + ) + texts = ( + tuple(item.casefold() for item in plain_candidates) + if descriptor.ignorecase + else plain_candidates + ) + if any(text in candidates for text in texts): + matched_any = True + else: + return "miss" + elif kind == "keywords": + saw_deterministic = True + values = descriptor.value if isinstance(descriptor.value, tuple) else () + if any( + item and item in text for text in plain_candidates for item in values + ): + 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" + elif kind == "is_type": + if event is None: + return "unknown" + types = descriptor.value + if isinstance(types, type): + if not isinstance(event, types): + return "miss" + 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" + if saw_unknown: + return "unknown" + 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: + commands: set[str] = set() + collect_command_literals(getattr(matcher_cls, "command", None), commands) + 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 + for attr in ("cmds", "command", "commands", "cmd"): + collect_command_literals(getattr(call, attr, None), commands) + normalized_commands = { + normalized for item in commands if (normalized := normalize_command(item)) + } + normalized = tuple(sorted(normalized_commands)) + return normalized or None + + +def collect_command_literals(value: Any, target: set[str], depth: int = 0) -> None: + if depth > 3 or value is None: + return + if isinstance(value, str): + text = value.strip() + if text: + target.add(text) + return + if isinstance(value, weakref.ReferenceType): + resolved = value() + if resolved is not None and resolved is not value: + collect_command_literals(resolved, target, depth + 1) + return + if isinstance(value, list | tuple | set | frozenset): + if all(isinstance(item, str) for item in value): + parts = tuple(str(item).strip() for item in value if str(item).strip()) + if parts: + target.add(" ".join(parts)) + target.add("".join(parts)) + return + for item in value: + collect_command_literals(item, target, depth + 1) + return + if callable(value) and getattr(value, "__self__", None) is not None: + with contextlib.suppress(TypeError, RuntimeError, ReferenceError): + resolved = value() + if resolved is not None and resolved is not value: + collect_command_literals(resolved, target, depth + 1) + return + for attr in ( + "name", + "path", + "aliases", + "header_display", + "command", + "commands", + "cmd", + "cmds", + ): + nested = getattr(value, attr, None) + if nested is not None and nested is not value: + collect_command_literals(nested, target, depth + 1) + + +def extract_matcher_alconna_shortcuts( + matcher_cls: type[Matcher], +) -> tuple[str, ...] | None: + shortcuts: set[str] = set() + for attr in ("command", "_rule", "rule"): + collect_alconna_shortcuts(getattr(matcher_cls, attr, None), shortcuts) + 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 + command = resolve_maybe_weakref( + getattr(call, "command", None) or getattr(call, "alconna", None) + ) + collect_alconna_shortcuts(command, shortcuts) + normalized_shortcuts = { + normalized + for item in shortcuts + if item and (normalized := normalize_shortcut_pattern(item)) + } + normalized = tuple(sorted(normalized_shortcuts)) + return normalized if normalized else None + + +def collect_alconna_shortcuts(value: Any, target: set[str], depth: int = 0) -> None: + if depth > 4 or value is None: + return + if isinstance(value, weakref.ReferenceType): + resolved = value() + if resolved is not None and resolved is not value: + collect_alconna_shortcuts(resolved, target, depth + 1) + return + get_shortcuts = getattr(value, "get_shortcuts", None) + if callable(get_shortcuts): + with contextlib.suppress(Exception): + raw_shortcuts = get_shortcuts() + if isinstance(raw_shortcuts, list | tuple | set | frozenset): + for shortcut in raw_shortcuts: + if isinstance(shortcut, str) and shortcut.strip(): + target.add(shortcut.strip()) + elif callable(value): + with contextlib.suppress(Exception): + resolved = value() + if resolved is not None and resolved is not value: + collect_alconna_shortcuts(resolved, target, depth + 1) + return + formatter = getattr(value, "formatter", None) + data = getattr(formatter, "data", None) + if isinstance(data, dict): + for trace in data.values(): + trace_shortcuts = getattr(trace, "shortcuts", None) + if not isinstance(trace_shortcuts, dict): + continue + for shortcut in trace_shortcuts: + if isinstance(shortcut, str) and shortcut.strip(): + target.add(shortcut.strip()) + for attr in ("shortcut", "shortcuts"): + shortcuts = getattr(value, attr, None) + if isinstance(shortcuts, dict): + for key in shortcuts: + if isinstance(key, str) and key.strip(): + target.add(key.strip()) + elif isinstance(shortcuts, list | tuple | set | frozenset): + for item in shortcuts: + if isinstance(item, str) and item.strip(): + target.add(item.strip()) + with contextlib.suppress(Exception): + from arclet.alconna import command_manager + + for shortcut_map in command_manager.get_shortcut(value).values(): # type: ignore[arg-type] + origin_key = getattr(shortcut_map, "origin_key", None) + if isinstance(origin_key, str) and origin_key.strip(): + target.add(origin_key.strip()) + for attr in ("command", "commands", "base", "formatter", "source"): + nested = getattr(value, attr, None) + if nested is not None and nested is not value: + collect_alconna_shortcuts(nested, target, depth + 1) + + +def resolve_maybe_weakref(value: Any) -> Any: + if isinstance(value, weakref.ReferenceType): + resolved = value() + return resolved if resolved is not None else value + if callable(value) and getattr(value, "__self__", None) is not None: + with contextlib.suppress(TypeError, RuntimeError, ReferenceError): + resolved = value() + if resolved is not None and resolved is not value: + return resolved + return value + + +def normalize_command(command: str) -> str: + text = command.strip() + if not text: + return "" + text = re.sub(r"^(?:\s*(?:\[[^\]]*]|\<[^>]*>))+\s*", "", text) + cut_points = [idx for idx in (text.find("["), text.find("<")) if idx >= 0] + if cut_points: + text = text[: min(cut_points)] + text = re.sub(r"\s+", " ", text).strip() + text = re.sub(r"(?:\s+[?*]+|[?*]+)$", "", text).strip() + return text + + +def matcher_command_matches(text: str, command: str) -> bool: + normalized = command.strip() + if not normalized: + return False + if normalized.startswith("re:"): + pattern = normalized.removeprefix("re:").strip() + if not pattern: + return False + try: + return re.search(pattern, text) is not None + except re.error: + return False + if command_matches(text, normalized): + return True + return text.startswith(normalized) and not normalized[-1].isascii() + + +def command_matches(text: str, command: str) -> bool: + if not text or not command: + return False + if text == command: + return True + if text.startswith(command): + if len(text) == len(command): + return True + return text[len(command)].isspace() + return False + + +def is_regex_like_command_literal(command: str) -> bool: + text = command.strip() + if not text: + return False + if text.startswith("re:"): + return True + return any(token in text for token in ("\\", "(", ")", "[", "]", "|", "^", "$")) + + +def normalize_shortcut_pattern(pattern: str) -> str: + text = str(pattern or "").strip() + if not text: + return "" + text = re.sub(r"^\[(?:[^\]]*)\]\s*", "", text) + text = re.sub(r"\s*\.\.\.args?$", "", text).strip() + text = re.sub(r"\s+", " ", text) + text = re.sub(r"\s*\.\.\.$", "", text).strip() + text = re.sub(r"^\^", "", text) + return text + + +def matcher_alconna_shortcut_matches( + shortcuts: tuple[str, ...] | None, + text: str, +) -> ActivationDecision: + if shortcuts is None: + return "unknown" + for shortcut in shortcuts: + if shortcut_matches_text(text, shortcut): + return "match" + return "unknown" + + +def matcher_alconna_shortcut_matches_any( + shortcuts: tuple[str, ...] | None, + texts: Iterable[str], +) -> ActivationDecision: + if shortcuts is None: + return "unknown" + for text in texts: + if matcher_alconna_shortcut_matches(shortcuts, text) == "match": + return "match" + return "unknown" + + +def shortcut_matches_text(text: str, shortcut: str) -> bool: + pattern = normalize_shortcut_pattern(shortcut) + if not pattern: + return False + if placeholder_shortcut_matches(text, pattern): + return True + if is_regex_like_shortcut(pattern): + try: + return re.match(pattern, text) is not None + except re.error: + return False + return matcher_command_matches(text, pattern) + + +def placeholder_shortcut_matches(text: str, pattern: str) -> bool: + if "{" not in pattern or "}" not in pattern: + return False + pieces: list[str] = [] + last = 0 + for match in re.finditer(r"\{[^{}]+\}", pattern): + pieces.append(re.escape(pattern[last : match.start()])) + pieces.append(r"\S+") + last = match.end() + if not pieces: + return False + pieces.append(re.escape(pattern[last:])) + try: + return re.match(rf"^{''.join(pieces)}(?:\s|$)", text) is not None + except re.error: + return False + + +def is_regex_like_shortcut(pattern: str) -> bool: + return any(token in pattern for token in ("\\", "(", ")", "[", "]", "|", "^", "$")) + + +def matcher_matches_ai_route_heads( + descriptor: HandlerDescriptor, + ai_route_heads: set[str], +) -> bool: + if not ai_route_heads: + return False + for command in descriptor.commands: + normalized_command = command.strip().casefold() + if not normalized_command: + continue + for head in ai_route_heads: + if not head: + continue + if matcher_command_matches(head, normalized_command) or command_matches( + normalized_command, + head, + ): + return True + for shortcut in descriptor.shortcuts or (): + for head in ai_route_heads: + if head and shortcut_matches_text(head, shortcut): + return True + return False + + +def _selection_is_guaranteed( + descriptor: HandlerDescriptor, + context: ActivationContext, +) -> bool: + """Return True for candidates that must not be budget-throttled.""" + + if descriptor.temp or descriptor.lane == "system": + return True + if context.event_type != "message": + return not descriptor.command_like + if descriptor.lane == "passive_http" and ( + context.has_url or _looks_like_rich_message(context.raw_text) + ): + return True + return False + + +def _looks_like_rich_message(text: str) -> bool: + lowered = (text or "").casefold() + return any( + marker in lowered + for marker in ( + "[cq:json", + "[json:", + "[cq:xml", + "[xml:", + "qqdocurl", + "jumpurl", + "miniapp", + "com.tencent", + ) + ) + + +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( + descriptor: HandlerDescriptor, + budget: dict[str, int], +) -> bool: + lane = _uncertain_budget_lane(descriptor) + if lane not in budget: + return True + if budget[lane] <= 0: + return False + budget[lane] -= 1 + return True + + +PASSIVE_DB_HINTS = ( + "word_bank", + "black_word", + "history", + "statistics", + "sign", + "gold", + "redbag", + "mute", + "group", + "user", + "admin", + "ban", + "limit", + "check", +) +PASSIVE_HTTP_HINTS = ( + "http", + "translate", + "bilibili", + "music", + "comment", + "nbnhhsh", + "quote", + "search", + "jitang", + "poetry", + "anime", + "cover", +) +PASSIVE_AI_HINTS = ( + "chatinter", + "dialogue", + "ai", + "llm", + "fudu", + "bym_ai", +) +PASSIVE_RENDER_HINTS = ( + "render", + "image", + "meme", + "memes", + "word_cloud", + "wordcloud", + "pic", + "picture", + "coser", + "luxun", +) + + +__all__ = [ + "ActivationContext", + "ActivationDecision", + "ActivationResult", + "ActivationRuleDescriptor", + "HandlerActivationIndex", + "HandlerDescriptor", + "classify_matcher_lane", + "command_matches", + "extract_matcher_alconna_shortcuts", + "extract_matcher_command_literals", + "extract_matcher_rule_descriptors", + "matcher_command_matches", + "matcher_has_custom_rule", + "matcher_has_deterministic_text_rule", + "matcher_is_command_like", + "matcher_rule_matches_text", +] diff --git a/zhenxun/builtin_plugins/hooks/auth_checker.py b/zhenxun/builtin_plugins/hooks/auth_checker.py index 5907a7f1..df6f1df1 100644 --- a/zhenxun/builtin_plugins/hooks/auth_checker.py +++ b/zhenxun/builtin_plugins/hooks/auth_checker.py @@ -1,13 +1,12 @@ import asyncio -from collections.abc import Awaitable, Callable import contextlib -import importlib import re import time from typing import cast from nonebot import get_loaded_plugins from nonebot.adapters import Bot, Event +from nonebot.consts import CMD_ARG_KEY, CMD_KEY, PREFIX_KEY, RAW_CMD_KEY from nonebot.exception import IgnoredException from nonebot.matcher import Matcher import nonebot.message as nb_message @@ -16,30 +15,27 @@ 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 ( - BotMemoryCache, - BotSnapshot, - GroupMemoryCache, - GroupSnapshot, - LevelUserMemoryCache, - LevelUserSnapshot, PluginInfoMemoryCache, + PluginLimitMemoryCache, ) from zhenxun.services.data_access import DataAccess from zhenxun.services.log import logger -from zhenxun.services.message_load import is_overloaded -from zhenxun.utils.enum import BlockType, GoldHandle, PluginType +from zhenxun.services.message_load import is_overloaded, signal_overload +from zhenxun.utils.enum import GoldHandle, PluginType from zhenxun.utils.exception import InsufficientGold from zhenxun.utils.platform import PlatformUtils -from .auth.auth_admin import auth_admin from .auth.auth_ban import auth_ban -from .auth.auth_bot import auth_bot from .auth.auth_cost import auth_cost -from .auth.auth_group import auth_group -from .auth.auth_limit import LimitManager, auth_limit -from .auth.auth_plugin import auth_plugin +from .auth.auth_group import _is_group_wake_command +from .auth.auth_limit import LimitManager, reserve_auth_limit from .auth.bot_filter import bot_filter from .auth.config import LOGGER_COMMAND, WARNING_THRESHOLD from .auth.context import ( @@ -57,25 +53,88 @@ from .auth.exception import ( PermissionExemption, SkipPluginException, ) -from .auth.utils import send_message +from .auth_activation import ( + ActivationContext, + ActivationResult, + HandlerActivationIndex, + classify_matcher_lane, + extract_matcher_alconna_shortcuts, + text_match_candidates, +) +from .auth_event_selector import ( + HandleEventSelectorDependencies, + install_handle_event_selector, + uninstall_handle_event_selector, +) +from .auth_legacy_fallback import legacy_pure_auth_fallback +from .auth_pipeline import ( + AuthPipelineContext, + AuthPipelineDependencies, + build_auth_pipeline, + decision_log_stage, +) +from .auth_policy import ( + PolicyContext, + PolicyDecisionPoint, +) +from .auth_profile import get_plugin_auth_profile +from .auth_runtime_config import AUTH_DISPATCH_RUNTIME_CONFIG +from .auth_side_effect import SideEffectCommit +from .auth_snapshot import get_or_build_auth_snapshot +from .auth_trace import HookTraceRecorder +from .auth_types import ( + AuthLaneContext, + AuthPreparation, + EventDispatchContext, +) -AUTH_HOOKS_CONCURRENCY_LIMIT = 5 -AUTH_DB_CONCURRENCY_LIMIT = 6 +AUTH_HOOKS_CONCURRENCY_LIMIT = AUTH_DISPATCH_RUNTIME_CONFIG.hooks_concurrency_limit +AUTH_DB_CONCURRENCY_LIMIT = AUTH_DISPATCH_RUNTIME_CONFIG.db_concurrency_limit +AUTH_DISPATCH_COMMAND_EXACT_LIMIT = AUTH_DISPATCH_RUNTIME_CONFIG.command_exact_limit +AUTH_DISPATCH_COMMAND_SHORTCUT_LIMIT = ( + AUTH_DISPATCH_RUNTIME_CONFIG.command_shortcut_limit +) +AUTH_DISPATCH_COMMAND_REGEX_LIMIT = AUTH_DISPATCH_RUNTIME_CONFIG.command_regex_limit +AUTH_DISPATCH_SYSTEM_LIMIT = AUTH_DISPATCH_RUNTIME_CONFIG.system_limit +AUTH_DISPATCH_PASSIVE_LIGHT_LIMIT = AUTH_DISPATCH_RUNTIME_CONFIG.passive_light_limit +AUTH_DISPATCH_PASSIVE_DB_LIMIT = AUTH_DISPATCH_RUNTIME_CONFIG.passive_db_limit +AUTH_DISPATCH_PASSIVE_HTTP_LIMIT = AUTH_DISPATCH_RUNTIME_CONFIG.passive_http_limit +AUTH_DISPATCH_PASSIVE_AI_LIMIT = AUTH_DISPATCH_RUNTIME_CONFIG.passive_ai_limit +AUTH_DISPATCH_PASSIVE_RENDER_LIMIT = AUTH_DISPATCH_RUNTIME_CONFIG.passive_render_limit +AUTH_OVERLOAD_SELECTED_THRESHOLD = ( + AUTH_DISPATCH_RUNTIME_CONFIG.overload_selected_threshold +) +AUTH_OVERLOAD_LANE_WAIT_MS = AUTH_DISPATCH_RUNTIME_CONFIG.overload_lane_wait_ms # 超时设置(秒) -TIMEOUT_SECONDS = 5.0 +TIMEOUT_SECONDS = AUTH_DISPATCH_RUNTIME_CONFIG.timeout_seconds # 熔断计数器 CIRCUIT_BREAKERS = { "auth_ban": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0}, - "auth_bot": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0}, - "auth_group": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0}, - "auth_admin": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0}, - "auth_plugin": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0}, "auth_limit": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0}, + "auth_hooks_gather": { + "failures": 0, + "threshold": 3, + "active": False, + "reset_time": 0, + }, + "get_plugin_cost": { + "failures": 0, + "threshold": 3, + "active": False, + "reset_time": 0, + }, + "get_plugin_and_user": { + "failures": 0, + "threshold": 3, + "active": False, + "reset_time": 0, + }, + "reserve_gold": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0}, } # 熔断重置时间(秒) -CIRCUIT_RESET_TIME = 300 # 5分钟 +CIRCUIT_RESET_TIME = AUTH_DISPATCH_RUNTIME_CONFIG.circuit_reset_time # 并发控制:限制同时进入 hooks 并行检查的协程数 HOOKS_CONCURRENCY_LIMIT = AUTH_HOOKS_CONCURRENCY_LIMIT @@ -87,9 +146,10 @@ _ROUTE_INDEX_READY = False _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 = 2 -PREFILTER_STATS_LOG_INTERVAL = 10.0 -CACHE_SWEEP_INTERVAL = 1.0 +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 @@ -97,17 +157,31 @@ HOOKS_SEMAPHORE = asyncio.Semaphore(HOOKS_CONCURRENCY_LIMIT) DB_SEMAPHORE = asyncio.Semaphore(DB_CONCURRENCY_LIMIT) DB_ACTIVE_COUNT = 0 -_CHECK_MATCHER_PATCHED = False -_ORIGINAL_CHECK_AND_RUN_MATCHER: Callable[..., Awaitable[None]] | None = None -_HANDLE_EVENT_PATCHED = False -_ORIGINAL_HANDLE_EVENT: Callable[..., Awaitable[None]] | None = None -_ORIGINAL_ADAPTER_HANDLE_EVENTS: dict[object, object] = {} +_DISPATCH_LANE_LIMITS: dict[str, int] = { + "command_exact": AUTH_DISPATCH_COMMAND_EXACT_LIMIT, + "command_shortcut": AUTH_DISPATCH_COMMAND_SHORTCUT_LIMIT, + "command_regex": AUTH_DISPATCH_COMMAND_REGEX_LIMIT, + "system": AUTH_DISPATCH_SYSTEM_LIMIT, + "passive_light": AUTH_DISPATCH_PASSIVE_LIGHT_LIMIT, + "passive_db": AUTH_DISPATCH_PASSIVE_DB_LIMIT, + "passive_http": AUTH_DISPATCH_PASSIVE_HTTP_LIMIT, + "passive_ai": AUTH_DISPATCH_PASSIVE_AI_LIMIT, + "passive_render": AUTH_DISPATCH_PASSIVE_RENDER_LIMIT, +} +_DISPATCH_LANE_SEMAPHORES = { + lane: asyncio.Semaphore(limit) + for lane, limit in _DISPATCH_LANE_LIMITS.items() + if limit > 0 +} +_DISPATCH_BUDGET_LANES = set(_DISPATCH_LANE_LIMITS) +_HANDLER_ACTIVATION_INDEX = HandlerActivationIndex() +_AUTH_PDP = PolicyDecisionPoint() _MATCHER_COMMAND_TYPE_CACHE: dict[type[Matcher], bool] = {} -_MATCHER_COMMAND_LITERAL_CACHE: dict[type[Matcher], tuple[str, ...] | None] = {} -_MATCHER_ALCONNA_SHORTCUT_CACHE: dict[type[Matcher], bool] = {} _CHECK_MATCHER_ROUTE_CACHE = CacheDict( "AUTH_MATCHER_ROUTE_CACHE", expire=MATCHER_ROUTE_PREFILTER_TTL ) + + _PREFILTER_STATS = { "checked": 0, "skipped": 0, @@ -121,90 +195,69 @@ _PREFILTER_STATS = { "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( r"^bot_manage\s+bot_switch\s+enable(?:\s+\S+)?$", re.IGNORECASE ) - - -class HookTraceRecorder: - def __init__(self, start_time: float) -> None: - self._start_time = start_time - self._enabled = False - self._data: dict[str, str] = {} - - def _ensure_enabled(self) -> bool: - if self._enabled: - return True - if time.time() - self._start_time <= WARNING_THRESHOLD: - return False - self._enabled = True - return True - - def set(self, key: str, value: str) -> None: - if self._ensure_enabled(): - self._data[key] = value - - def setdefault(self, key: str, value: str) -> None: - if self._ensure_enabled(): - self._data.setdefault(key, value) - - def contains(self, key: str) -> bool: - return key in self._data - - def snapshot(self) -> dict[str, str]: - return self._data if self._enabled else {} - - -def _debug_log(message: str, *args, **kwargs) -> None: - if is_overloaded(): - return - logger.debug(message, *args, **kwargs) +_URL_PATTERN = re.compile(r"(?:https?://|www\.|b23\.tv|t\.cn/)", re.IGNORECASE) def _normalize_command(command: str) -> str: text = command.strip() if not text: return "" - - # strip leading placeholders like "[引用消息] 撤回" text = re.sub(r"^(?:\s*(?:\[[^\]]*]|\<[^>]*>))+\s*", "", text) - - # keep command head: "点歌 [歌名]" -> "点歌", "foo " -> "foo" cut_points = [idx for idx in (text.find("["), text.find("<")) if idx >= 0] if cut_points: text = text[: min(cut_points)] - - # normalize spacing after trimming placeholders text = re.sub(r"\s+", " ", text).strip() - # remove trailing template markers left by forms like "foo ?[arg]" / "foo ?*[tags]" - text = re.sub(r"(?:\s+[?*]+|[?*]+)$", "", text).strip() - return text - - -def _is_bot_wake_command(module: str, text: str | None) -> bool: - if "bot_manage" not in (module or ""): - return False - if not text: - return False - normalized = re.sub(r"\s+", " ", text.strip()) - if not normalized: - return False - return ( - _BOT_WAKE_COMMAND_PATTERN.match(normalized) is not None - or _BOT_WAKE_CANONICAL_PATTERN.match(normalized) is not None - ) + return re.sub(r"(?:\s+[?*]+|[?*]+)$", "", text).strip() def _split_command_variants(command: str) -> tuple[str, ...]: text = command.strip() if not text: return () - # Keep slash-prefixed commands like "/info" as-is. if text.startswith("/"): return (text,) - # "今日运势/抽签/运势" => ("今日运势", "抽签", "运势") if "/" in text and " " not in text: parts = tuple(part.strip() for part in text.split("/") if part.strip()) if parts: @@ -216,12 +269,9 @@ def _is_ambiguous_route_command(command: str) -> bool: text = command.strip() if not text: return True - # Keep route-index strict only for literal, deterministic command heads. if any(token in text for token in ("?", "*", "|", "(", ")", "^", "$", "re:")): return True - if "xx" in text.lower(): - return True - return False + return "xx" in text.lower() def _extract_commands(extra: PluginExtraData | None) -> tuple[set[str], bool]: @@ -262,9 +312,7 @@ async def _ensure_route_index(): except Exception: continue command_set, has_ambiguous = _extract_commands(extra_data) - if not command_set: - continue - if has_ambiguous: + if not command_set or has_ambiguous: continue module = plugin.name _ROUTE_MODULES_WITH_COMMANDS.add(module) @@ -277,7 +325,7 @@ async def _ensure_route_index(): _ROUTE_INDEX_READY = True -def _command_matches(text: str, command: str) -> bool: +def _route_command_matches(text: str, command: str) -> bool: if not text or not command: return False if text == command: @@ -285,8 +333,7 @@ def _command_matches(text: str, command: str) -> bool: if text.startswith(command): if len(text) == len(command): return True - next_char = text[len(command)] - return next_char.isspace() + return text[len(command)].isspace() return False @@ -299,21 +346,52 @@ def _match_route_modules(text: str) -> set[str]: return set() matched_modules: set[str] = set() for command in commands: - if _command_matches(text, command): + if _route_command_matches(text, command): modules = _ROUTE_COMMAND_MAP.get(command) if modules: matched_modules.update(modules) return matched_modules -def _matcher_module_name(matcher_cls: type[Matcher]) -> str: - module = getattr(matcher_cls, "plugin_name", "") or "" - if module: - return module - plugin = getattr(matcher_cls, "plugin", None) - if not plugin: - return "" - return (getattr(plugin, "name", "") or "").strip() +def _is_bot_wake_command(module: str, text: str | None) -> bool: + if "bot_manage" not in (module or ""): + return False + if not text: + return False + normalized = re.sub(r"\s+", " ", text.strip()) + if not normalized: + return False + return ( + _BOT_WAKE_COMMAND_PATTERN.match(normalized) is not None + or _BOT_WAKE_CANONICAL_PATTERN.match(normalized) is not None + ) + + +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] + descriptor = _HANDLER_ACTIVATION_INDEX.descriptor_for(matcher_cls) + if descriptor is not None: + result = descriptor.command_like + else: + from .auth_activation import matcher_is_command_like + + result = matcher_is_command_like(matcher_cls) + _MATCHER_COMMAND_TYPE_CACHE[matcher_cls] = result + return result + + +def _matcher_has_alconna_shortcuts(matcher_cls: type[Matcher]) -> bool: + descriptor = _HANDLER_ACTIVATION_INDEX.descriptor_for(matcher_cls) + if descriptor is not None: + return bool(descriptor.shortcuts) + return bool(extract_matcher_alconna_shortcuts(matcher_cls)) def _collect_ai_route_modules(event: Event, state: dict | None = None) -> set[str]: @@ -366,62 +444,6 @@ def _collect_ai_route_heads(event: Event, state: dict | None = None) -> set[str] return result -def _matcher_matches_ai_route_heads( - matcher_cls: type[Matcher], - ai_route_heads: set[str], -) -> bool: - if not ai_route_heads: - return False - matcher_commands = _extract_matcher_command_literals(matcher_cls) - if not matcher_commands: - return False - for command in matcher_commands: - normalized_command = command.strip().casefold() - if not normalized_command: - continue - for head in ai_route_heads: - if not head: - continue - if _command_matches(head, normalized_command) or _command_matches( - normalized_command, head - ): - return True - return False - - -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] - if hasattr(matcher_cls, "command"): - _MATCHER_COMMAND_TYPE_CACHE[matcher_cls] = True - return True - 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 - call_type = call.__class__ - call_module = getattr(call_type, "__module__", "") - call_name = getattr(call_type, "__name__", "") - if call_module.startswith("nonebot.rule") and call_name in { - "CommandRule", - "ShellCommandRule", - "Command", - "ShellCommand", - }: - _MATCHER_COMMAND_TYPE_CACHE[matcher_cls] = True - return True - if ( - call_module.startswith("nonebot_plugin_alconna.rule") - and call_name == "AlconnaRule" - ): - _MATCHER_COMMAND_TYPE_CACHE[matcher_cls] = True - return True - _MATCHER_COMMAND_TYPE_CACHE[matcher_cls] = False - return False - - def _matcher_route_cache_key(event: Event) -> str: msg_id = getattr(event, "message_id", None) if msg_id is None: @@ -458,6 +480,19 @@ def _event_plain_text(event: Event) -> str: return "" +def _normalize_dispatch_text(text: str) -> str: + normalized = text.strip() + if not normalized: + return "" + # strip leading placeholders like "[reply:id=10004]撤回" + normalized = re.sub( + r"^(?:\s*(?:\[[^\]]*]|\<[^>]*>))+\s*", + "", + normalized, + ) + return normalized.strip() + + def _state_plain_text(state: dict | None) -> str: if state is None: return "" @@ -470,6 +505,402 @@ def _state_plain_text(state: dict | None) -> str: return "" +def _message_to_plain_text(message: object) -> str: + if message is None: + return "" + with contextlib.suppress(Exception): + extractor = getattr(message, "extract_plain_text", None) + if callable(extractor): + return _normalize_dispatch_text(str(extractor() or "")) + return _normalize_dispatch_text(str(message)) + + +def _trie_command_text_from_state(state: dict | None) -> str: + if state is None: + return "" + prefix = state.get(PREFIX_KEY) + if not isinstance(prefix, dict): + return "" + command = prefix.get(CMD_KEY) + if isinstance(command, tuple): + return _normalize_dispatch_text(" ".join(str(item) for item in command)) + if isinstance(command, str): + return _normalize_dispatch_text(command) + return "" + + +def _trie_raw_command_from_state(state: dict | None) -> str: + if state is None: + return "" + prefix = state.get(PREFIX_KEY) + if not isinstance(prefix, dict): + return "" + raw_command = prefix.get(RAW_CMD_KEY) + return _normalize_dispatch_text(raw_command) if isinstance(raw_command, str) else "" + + +def _trie_command_arg_text_from_state(state: dict | None) -> str: + if state is None: + return "" + prefix = state.get(PREFIX_KEY) + if not isinstance(prefix, dict): + return "" + return _message_to_plain_text(prefix.get(CMD_ARG_KEY)) + + +def _event_text_candidates( + event: Event, + state: dict | None, + plain_text: str = "", + raw_text: str = "", +) -> tuple[str, ...]: + candidates: list[str] = [] + + def add(text: object) -> None: + if not isinstance(text, str): + return + normalized = _normalize_dispatch_text(text) + if normalized and normalized not in candidates: + candidates.append(normalized) + + add(_trie_command_text_from_state(state)) + add(_trie_raw_command_from_state(state)) + trie_arg = _trie_command_arg_text_from_state(state) + if _trie_raw_command_from_state(state) and trie_arg: + add(f"{_trie_raw_command_from_state(state)} {trie_arg}") + add(plain_text) + if event is not None: + with contextlib.suppress(Exception): + getter = getattr(event, "get_plaintext", None) + if callable(getter): + add(getter()) + add(raw_text) + return tuple(candidates) + + +def _event_raw_message_text(event: Event) -> str: + with contextlib.suppress(Exception): + message = getattr(event, "message", None) + if message is not None: + return str(message) + return "" + + +def _event_has_image(event: Event) -> bool: + text = _event_raw_message_text(event) + lowered = text.casefold() + return "[cq:image" in lowered or "[image:" in lowered + + +def _event_has_url(text: str) -> bool: + return bool(_URL_PATTERN.search(text)) + + +def _event_to_me(event: Event) -> bool: + with contextlib.suppress(Exception): + getter = getattr(event, "is_tome", None) + if callable(getter): + return bool(getter()) + return bool(getattr(event, "to_me", False)) + + +def _context_from_state(state: dict | None) -> EventDispatchContext | None: + if state is None: + return None + context = state.get("_zx_dispatch_context") + return context if isinstance(context, EventDispatchContext) else None + + +def _build_dispatch_context_sync( + event: Event, state: dict | None = None +) -> EventDispatchContext: + context = _context_from_state(state) + if context is not None: + return context + + event_type = event.get_type() + plain_text = _state_plain_text(state) + if not plain_text: + plain_text = _event_plain_text(event) + if state is not None and plain_text: + state["_zx_plain_text"] = plain_text + + route_modules = ( + _get_route_modules_for_event(event, state) if _ROUTE_INDEX_READY else set() + ) + ai_route_modules = _collect_ai_route_modules(event, state) + ai_route_heads = _collect_ai_route_heads(event, state) + raw_text = _event_raw_message_text(event) + text_candidates = _event_text_candidates(event, state, plain_text, raw_text) + trie_command_text = _trie_command_text_from_state(state) + trie_raw_command = _trie_raw_command_from_state(state) + to_me = _event_to_me(event) + has_url = _event_has_url(raw_text) or _event_has_url(plain_text) + has_image = _event_has_image(event) + is_command_like = bool( + route_modules + or ai_route_modules + or trie_command_text + or trie_raw_command + or plain_text.startswith("/") + or plain_text.startswith("!") + or plain_text.startswith(".") + ) + context = EventDispatchContext( + event_type=event_type, + plain_text=plain_text, + raw_text=raw_text, + trie_command_text=trie_command_text, + trie_raw_command=trie_raw_command, + text_candidates=text_candidates, + to_me=to_me, + has_url=has_url, + has_image=has_image, + is_command_like=is_command_like, + route_modules=route_modules, + ai_route_modules=ai_route_modules, + ai_route_heads=ai_route_heads, + ) + if state is not None: + state["_zx_dispatch_context"] = context + return context + + +async def _build_dispatch_context( + event: Event, state: dict | None = None +) -> EventDispatchContext: + context = _build_dispatch_context_sync(event, state) + await _ensure_route_index() + if not context.route_modules: + route_modules = _get_route_modules_for_event(event, state) + context.route_modules = route_modules + context.is_command_like = bool( + route_modules + or context.ai_route_modules + or context.trie_command_text + or context.trie_raw_command + or context.plain_text.startswith("/") + or context.plain_text.startswith("!") + or context.plain_text.startswith(".") + ) + return context + + +def _dispatch_lane_for_matcher( + matcher_cls: type[Matcher], context: EventDispatchContext +) -> str: + descriptor = _HANDLER_ACTIVATION_INDEX.descriptor_for(matcher_cls) + if descriptor is not None: + return descriptor.lane + + event_type = context.event_type + if getattr(matcher_cls, "temp", False): + return "system" + matcher_type = getattr(matcher_cls, "type", "") or "" + if isinstance(matcher_type, str) and matcher_type and matcher_type != event_type: + return "system" + return classify_matcher_lane( + matcher_cls, + ai_route_modules=context.ai_route_modules, + ) + + +def _activation_context_from_dispatch( + context: EventDispatchContext, + event: Event, +) -> ActivationContext: + return ActivationContext( + event=event, + event_type=context.event_type, + plain_text=context.text_candidates[0] + if context.text_candidates + else context.plain_text, + raw_text="\n".join(context.text_candidates) + if context.text_candidates + else context.raw_text, + to_me=context.to_me, + has_url=context.has_url, + has_image=context.has_image, + is_command_like=context.is_command_like, + route_modules=set(context.route_modules), + ai_route_modules=set(context.ai_route_modules), + ai_route_heads=set(context.ai_route_heads), + ) + + +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) + + +def _merge_dispatch_budget( + target: dict[str, int], + source: dict[str, int], +) -> None: + for lane in _DISPATCH_BUDGET_LANES: + 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 "" + message_id = context.message_id if context.message_id is not None else "" + return ( + f"{context.platform}:{context.bot_id}:" + f"{context.user_id}:{group_id}:{channel_id}:{message_id}" + ) + + +def _auth_lane_context_from_state( + matcher_cls: type[Matcher], + auth_context: EventContext, + state: dict | None, +) -> AuthLaneContext: + dispatch_context = None + if state is not None: + value = state.get("_zx_dispatch_context") + if isinstance(value, EventDispatchContext): + dispatch_context = value + if dispatch_context is None: + dispatch_context = EventDispatchContext( + event_type=auth_context.event_type, + plain_text=auth_context.plain_text, + text_candidates=(auth_context.plain_text,) + if auth_context.plain_text + else (), + is_command_like=bool(auth_context.route_modules), + route_modules=set(auth_context.route_modules), + ) + lane = _dispatch_lane_for_matcher(matcher_cls, dispatch_context) + semaphore = _DISPATCH_LANE_SEMAPHORES.get(lane) + queue_size = 0 + if semaphore is not None: + limit = _DISPATCH_LANE_LIMITS.get(lane, 0) + value = getattr(semaphore, "_value", limit) + queue_size = max(limit - int(value), 0) + return AuthLaneContext( + lane=lane, + scope_key=_auth_scope_key(auth_context), + queue_size=queue_size, + ) + + +@contextlib.asynccontextmanager +async def _dispatch_lane_section(lane: str): + semaphore = _DISPATCH_LANE_SEMAPHORES.get(lane) + if semaphore is None: + yield + return + started = time.perf_counter() + await semaphore.acquire() + 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: + with contextlib.suppress(Exception): + semaphore.release() + + +def get_dispatch_snapshot() -> dict[str, object]: + lane_active = {} + for lane, semaphore in _DISPATCH_LANE_SEMAPHORES.items(): + limit = _DISPATCH_LANE_LIMITS.get(lane, 0) + 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), + } + + def _get_route_modules_for_event(event: Event, state: dict | None = None) -> set[str]: if state is not None: context = get_event_context(state) @@ -482,7 +913,11 @@ def _get_route_modules_for_event(event: Event, state: dict | None = None) -> set try: route_modules = _CHECK_MATCHER_ROUTE_CACHE[key] except KeyError: - route_modules = _match_route_modules(_event_plain_text(event)) + raw_text = _event_raw_message_text(event) + plain_text = _state_plain_text(state) or _event_plain_text(event) + route_modules = set() + for text in _event_text_candidates(event, state, plain_text, raw_text): + route_modules.update(_match_route_modules(text)) _CHECK_MATCHER_ROUTE_CACHE[key] = route_modules if state is not None: context = get_event_context(state) @@ -497,6 +932,15 @@ def _prepare_handle_event_state(event: Event, state: dict) -> None: get_permission_side_effect_cache(state=state) if event.get_type() != "message": return + raw_text = _event_raw_message_text(event) + text_candidates = _event_text_candidates( + event, + state, + _event_plain_text(event), + raw_text, + ) + if text_candidates: + state["_zx_text_candidates"] = text_candidates if _state_plain_text(state): return text = _event_plain_text(event) @@ -518,15 +962,17 @@ async def _run_selected_matcher( state: dict, stack, dependency_cache, + lane: str = "command_exact", ) -> None: - await nb_message.check_and_run_matcher( - matcher, - bot, - event, - state, - stack, - dependency_cache, - ) + async with _dispatch_lane_section(lane): + await nb_message.check_and_run_matcher( + matcher, + bot, + event, + state, + stack, + dependency_cache, + ) def _record_prefilter_stats( @@ -581,423 +1027,31 @@ def _record_prefilter_stats( ) -def _collect_command_literals(value, target: set[str], depth: int = 0) -> None: - if depth > 3 or value is None: - return - if isinstance(value, str): - text = value.strip() - if text: - target.add(text) - return - if isinstance(value, list | tuple | set | frozenset): - for item in value: - _collect_command_literals(item, target, depth + 1) - return - for attr in ("command", "commands", "cmd", "cmds"): - nested = getattr(value, attr, None) - if nested is not None and nested is not value: - _collect_command_literals(nested, target, depth + 1) - - -def _extract_matcher_command_literals( - matcher_cls: type[Matcher], -) -> tuple[str, ...] | None: - if matcher_cls in _MATCHER_COMMAND_LITERAL_CACHE: - return _MATCHER_COMMAND_LITERAL_CACHE[matcher_cls] - - commands: set[str] = set() - _collect_command_literals(getattr(matcher_cls, "command", None), commands) - - 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 - for attr in ("cmds", "command", "commands", "cmd"): - _collect_command_literals(getattr(call, attr, None), commands) - - if not commands: - _MATCHER_COMMAND_LITERAL_CACHE[matcher_cls] = None - return None - - sorted_commands = tuple(sorted(commands, key=len, reverse=True)) - _MATCHER_COMMAND_LITERAL_CACHE[matcher_cls] = sorted_commands - return sorted_commands - - -def _matcher_has_alconna_shortcuts(matcher_cls: type[Matcher]) -> bool: - cached = _MATCHER_ALCONNA_SHORTCUT_CACHE.get(matcher_cls) - if cached is not None: - return cached - - has_shortcuts = False - 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 - call_type = call.__class__ - call_module = getattr(call_type, "__module__", "") - call_name = getattr(call_type, "__name__", "") - if not ( - call_module.startswith("nonebot_plugin_alconna.rule") - and call_name == "AlconnaRule" - ): - continue - - # Alconna matcher supports shortcut-based parsing (regex/fuzzy expansion). - # Route prefilter only knows literal command heads, so shortcut matchers - # must bypass strict route miss to avoid false negative skips. - command_ref = getattr(call, "command", None) - command = None - if callable(command_ref): - with contextlib.suppress(Exception): - command = command_ref() - if command is not None: - get_shortcuts = getattr(command, "get_shortcuts", None) - if callable(get_shortcuts): - shortcuts = get_shortcuts() - if shortcuts: - has_shortcuts = True - break - formatter = getattr(command, "formatter", None) - if formatter is not None: - with contextlib.suppress(Exception): - data = getattr(formatter, "data", None) - if isinstance(data, dict): - for trace in data.values(): - if getattr(trace, "shortcuts", None): - has_shortcuts = True - break - if has_shortcuts: - break - - _MATCHER_ALCONNA_SHORTCUT_CACHE[matcher_cls] = has_shortcuts - return has_shortcuts - - -async def _check_matcher_prefilter( - matcher_cls: type[Matcher], event: Event, state: dict | None = None -) -> tuple[bool, str | None]: - event_type = event.get_type() - matcher_type = getattr(matcher_cls, "type", "") or "" - if isinstance(matcher_type, str) and matcher_type and matcher_type != event_type: - # Explicit matcher type mismatch cannot match this event. - return True, "type_miss" - - if event_type != "message": - return False, None - - # Session continuation matchers generated by pause/reject are temp=True. - # They must bypass command-route prefilter, otherwise follow-up messages - # (e.g. got_path waiting for plain text) will be dropped. - if getattr(matcher_cls, "temp", False): - return False, None - - is_command_matcher = _is_command_matcher_class(matcher_cls) - if not is_command_matcher: - return False, None - - text = _state_plain_text(state) - if is_command_matcher and not text: - text = _event_plain_text(event) - if state is not None and text: - state["_zx_plain_text"] = text - if is_command_matcher and not text: - return True, "empty_text" - - module = _matcher_module_name(matcher_cls) - if not module: - return False, None - - command_matched = False - matcher_commands = _extract_matcher_command_literals(matcher_cls) - if matcher_commands: - for command in matcher_commands: - if _command_matches(text, command): - command_matched = True - break - else: - if not _matcher_has_alconna_shortcuts(matcher_cls): - return True, "command_miss" - - ai_route_modules = _collect_ai_route_modules(event, state) - ai_route_heads = _collect_ai_route_heads(event, state) - if ai_route_modules and module not in ai_route_modules: - if not _matcher_matches_ai_route_heads(matcher_cls, ai_route_heads): - return True, "route_miss" - elif ai_route_modules: - return False, None - - if not _ROUTE_INDEX_READY: - await _ensure_route_index() - - if module not in _ROUTE_MODULES_WITH_COMMANDS: - return False, None - - route_modules = _get_route_modules_for_event(event, state) - if module not in route_modules: - if command_matched: - return False, None - if _matcher_has_alconna_shortcuts(matcher_cls): - return False, None - return True, "route_miss" - return False, None - - -def _check_matcher_prefilter_before_task( - matcher_cls: type[Matcher], event: Event, state: dict | None = None -) -> tuple[bool, str | None]: - """Conservative selector before creating matcher task. - - This mirrors the async matcher prefilter but never performs IO or route-index - rebuild. If anything is uncertain, let the existing check_and_run_matcher - patch handle it inside the task. - """ - event_type = event.get_type() - matcher_type = getattr(matcher_cls, "type", "") or "" - if isinstance(matcher_type, str) and matcher_type and matcher_type != event_type: - return True, "type_miss" - - if event_type != "message": - return False, None - - if getattr(matcher_cls, "temp", False): - return False, None - - if not _is_command_matcher_class(matcher_cls): - return False, None - - text = _state_plain_text(state) - if not text: - text = _event_plain_text(event) - if state is not None and text: - state["_zx_plain_text"] = text - if not text: - return True, "empty_text" - - module = _matcher_module_name(matcher_cls) - if not module: - return False, None - - command_matched = False - matcher_commands = _extract_matcher_command_literals(matcher_cls) - has_alconna_shortcuts = _matcher_has_alconna_shortcuts(matcher_cls) - if matcher_commands: - for command in matcher_commands: - if _command_matches(text, command): - command_matched = True - break - else: - if not has_alconna_shortcuts: - return True, "command_miss" - - ai_route_modules = _collect_ai_route_modules(event, state) - ai_route_heads = _collect_ai_route_heads(event, state) - if ai_route_modules and module not in ai_route_modules: - if not _matcher_matches_ai_route_heads(matcher_cls, ai_route_heads): - return True, "route_miss" - elif ai_route_modules: - return False, None - - if not _ROUTE_INDEX_READY: - return False, None - - if module not in _ROUTE_MODULES_WITH_COMMANDS: - return False, None - - route_modules = _get_route_modules_for_event(event, state) - if module not in route_modules: - if command_matched or has_alconna_shortcuts: - return False, None - return True, "route_miss" - return False, None - - _MAX_MATCHER_CACHE = 512 -async def _patched_check_and_run_matcher( - Matcher: type[Matcher], - bot: Bot, - event: Event, - state: dict, - stack=None, - dependency_cache=None, -) -> None: - skip, reason = await _check_matcher_prefilter( - Matcher, event, state if isinstance(state, dict) else None - ) - _record_prefilter_stats(skip, reason, "inside_task") - if skip: - return - - original = _ORIGINAL_CHECK_AND_RUN_MATCHER - if not original: - return - kwargs = { - "Matcher": Matcher, - "bot": bot, - "event": event, - "state": state, - "stack": stack, - "dependency_cache": dependency_cache, - } - await original(**kwargs) - - -def _install_matcher_prefilter() -> None: - global _CHECK_MATCHER_PATCHED, _ORIGINAL_CHECK_AND_RUN_MATCHER - if _CHECK_MATCHER_PATCHED: - return - _ORIGINAL_CHECK_AND_RUN_MATCHER = nb_message.check_and_run_matcher - nb_message.check_and_run_matcher = _patched_check_and_run_matcher # type: ignore[assignment] - _CHECK_MATCHER_PATCHED = True - - -def _uninstall_matcher_prefilter() -> None: - global _CHECK_MATCHER_PATCHED, _ORIGINAL_CHECK_AND_RUN_MATCHER - if not _CHECK_MATCHER_PATCHED: - return - if _ORIGINAL_CHECK_AND_RUN_MATCHER is not None: - nb_message.check_and_run_matcher = _ORIGINAL_CHECK_AND_RUN_MATCHER # type: ignore[assignment] - _CHECK_MATCHER_PATCHED = False - _ORIGINAL_CHECK_AND_RUN_MATCHER = None - - -async def _patched_handle_event(bot: Bot, event: Event) -> None: - show_log = True - escape_tag = getattr(nb_message, "escape_tag") - logger_ = getattr(nb_message, "logger") - no_log_exception = getattr(nb_message, "NoLogException") - - log_msg = f"{escape_tag(bot.type)} {escape_tag(bot.self_id)} | " - try: - log_msg += event.get_log_string() - except no_log_exception: - show_log = False - if show_log: - logger_.opt(colors=True).success(log_msg) - - state = {} - dependency_cache = {} - async_exit_stack = getattr(nb_message, "AsyncExitStack") - apply_event_preprocessors = getattr(nb_message, "_apply_event_preprocessors") - apply_event_postprocessors = getattr(nb_message, "_apply_event_postprocessors") - trie_rule = getattr(nb_message, "TrieRule") - matchers = getattr(nb_message, "matchers") - catch = getattr(nb_message, "catch") - stop_propagation = getattr(nb_message, "StopPropagation") - handle_exception = getattr(nb_message, "_handle_exception") - anyio_mod = getattr(nb_message, "anyio") - run_coro_with_shield = getattr(nb_message, "run_coro_with_shield") - - async with async_exit_stack() as stack: - if not await apply_event_preprocessors( - bot=bot, - event=event, - state=state, - stack=stack, - dependency_cache=dependency_cache, - ): - return - - try: - trie_rule.get_value(bot, event, state) - except Exception as e: - logger_.opt(colors=True, exception=e).warning( - "Error while parsing command for event" - ) - _prepare_handle_event_state(event, state) - - break_flag = False - - def _handle_stop_propagation(_exc_group) -> None: - nonlocal break_flag - break_flag = True - logger_.debug("Stop event propagation") - - for priority in sorted(matchers.keys()): - if break_flag: - break - - if show_log: - logger_.debug(f"Checking for matchers in priority {priority}...") - - if not (priority_matchers := matchers[priority]): - continue - - with catch( - { - stop_propagation: _handle_stop_propagation, - Exception: handle_exception( - "Error when checking Matcher." - ), - } - ): - async with anyio_mod.create_task_group() as tg: - for matcher in priority_matchers: - skip, reason = _check_matcher_prefilter_before_task( - matcher, - event, - state, - ) - _record_prefilter_stats(skip, reason, "before_task") - if skip: - continue - matcher_state = _build_matcher_state(state) - tg.start_soon( - run_coro_with_shield, - _run_selected_matcher( - matcher, - bot, - event, - matcher_state, - stack, - dependency_cache, - ), - ) - - if show_log: - logger_.debug("Checking for matchers completed") - - await apply_event_postprocessors(bot, event, state, stack, dependency_cache) +_SELECTOR_DEPS = HandleEventSelectorDependencies( + activation_index=_HANDLER_ACTIVATION_INDEX, + overload_selected_threshold=AUTH_OVERLOAD_SELECTED_THRESHOLD, + prepare_handle_event_state=_prepare_handle_event_state, + build_dispatch_context=_build_dispatch_context, + 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, +) def _install_handle_event_selector() -> None: - global _HANDLE_EVENT_PATCHED, _ORIGINAL_HANDLE_EVENT - if _HANDLE_EVENT_PATCHED: - return - _ORIGINAL_HANDLE_EVENT = nb_message.handle_event - nb_message.handle_event = _patched_handle_event # type: ignore[assignment] - for module_name in ( - "nonebot.adapters.onebot.v11.bot", - "nonebot.adapters.onebot.v12.bot", - "onebug.mixin.process", - ): - with contextlib.suppress(Exception): - module = importlib.import_module(module_name) - current = getattr(module, "handle_event", None) - if current is not None: - _ORIGINAL_ADAPTER_HANDLE_EVENTS[module] = current - setattr(module, "handle_event", _patched_handle_event) - _HANDLE_EVENT_PATCHED = True + install_handle_event_selector(_SELECTOR_DEPS) def _uninstall_handle_event_selector() -> None: - global _HANDLE_EVENT_PATCHED, _ORIGINAL_HANDLE_EVENT - if not _HANDLE_EVENT_PATCHED: - return - if _ORIGINAL_HANDLE_EVENT is not None: - nb_message.handle_event = _ORIGINAL_HANDLE_EVENT # type: ignore[assignment] - for module, original in list(_ORIGINAL_ADAPTER_HANDLE_EVENTS.items()): - with contextlib.suppress(Exception): - setattr(module, "handle_event", original) - _ORIGINAL_ADAPTER_HANDLE_EVENTS.clear() - _HANDLE_EVENT_PATCHED = False - _ORIGINAL_HANDLE_EVENT = None + uninstall_handle_event_selector() async def _get_route_context(text: str, event_cache: dict | None) -> set[str]: @@ -1006,12 +1060,23 @@ async def _get_route_context(text: str, event_cache: dict | None) -> set[str]: if event_cache is not None and "route_modules" in event_cache: return event_cache["route_modules"] await _ensure_route_index() - matched = _match_route_modules(text) + matched = set() + for candidate in text_match_candidates(text): + matched.update(_match_route_modules(candidate)) if event_cache is not None: event_cache["route_modules"] = matched 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) @@ -1019,11 +1084,7 @@ async def _cache_sweep_loop() -> None: if EVENT_CACHE is not None: _ = len(EVENT_CACHE) _ = len(_CHECK_MATCHER_ROUTE_CACHE) - for _mc in ( - _MATCHER_COMMAND_TYPE_CACHE, - _MATCHER_COMMAND_LITERAL_CACHE, - _MATCHER_ALCONNA_SHORTCUT_CACHE, - ): + for _mc in (_MATCHER_COMMAND_TYPE_CACHE,): if len(_mc) > _MAX_MATCHER_CACHE: _mc.clear() @@ -1031,7 +1092,6 @@ async def _cache_sweep_loop() -> None: async def start_auth_runtime_tasks() -> None: global _CACHE_SWEEP_TASK await _ensure_route_index() - _install_matcher_prefilter() _install_handle_event_selector() if _CACHE_SWEEP_TASK is None or _CACHE_SWEEP_TASK.done(): _CACHE_SWEEP_TASK = asyncio.create_task(_cache_sweep_loop()) @@ -1040,7 +1100,6 @@ async def start_auth_runtime_tasks() -> None: async def stop_auth_runtime_tasks() -> None: global _CACHE_SWEEP_TASK _uninstall_handle_event_selector() - _uninstall_matcher_prefilter() task = _CACHE_SWEEP_TASK _CACHE_SWEEP_TASK = None if task is not None: @@ -1049,20 +1108,63 @@ async def stop_auth_runtime_tasks() -> None: await task -async def _has_limits_cached(module: str, event_cache: dict | None) -> bool: +async def _has_limits_cached( + module: str, + event_cache: dict | None, + *, + known: bool | None = None, +) -> bool: module_limit_cache: dict[str, bool] = {} if event_cache is not None: module_limit_cache = event_cache.setdefault("module_limits", {}) if module in module_limit_cache: + ready_cache = ( + event_cache.setdefault("module_limits_ready", {}) + if event_cache is not None + else {} + ) + if ready_cache.get(module, True): + return module_limit_cache[module] + if known is True: + module_limit_cache[module] = True + ready_cache[module] = True + return True + elif known is not None: + module_limit_cache[module] = known + if event_cache is not None: + event_cache.setdefault("module_limits_ready", {})[module] = True return module_limit_cache[module] + limit_entries = None + if event_cache is not None: + entry_cache = event_cache.setdefault("module_limit_entries", {}) + if module in entry_cache: + limit_entries = entry_cache[module] + if limit_entries is not None: + has_limits = bool(limit_entries) + module_limit_cache[module] = has_limits + 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) + if limit_entries is not None: + has_limits = bool(limit_entries) + module_limit_cache[module] = has_limits + if event_cache is not None: + event_cache.setdefault("module_limit_entries", {})[module] = limit_entries + event_cache.setdefault("module_limits_ready", {})[module] = True + return has_limits limits = await LimitManager.get_module_limits(module) has_limits = bool(limits) module_limit_cache[module] = has_limits + if event_cache is not None: + event_cache.setdefault("module_limit_entries", {})[module] = limits + event_cache.setdefault("module_limits_ready", {})[module] = True return has_limits @contextlib.asynccontextmanager async def _db_section(): + """Legacy bounded DB section kept for explicit fallback callers.""" global DB_ACTIVE_COUNT if DB_SEMAPHORE.locked(): logger.warning( @@ -1080,87 +1182,25 @@ async def _db_section(): DB_ACTIVE_COUNT = max(DB_ACTIVE_COUNT - 1, 0) -async def _get_group_cached(entity, event_cache) -> GroupSnapshot | None: - if not entity.group_id: - return None - if event_cache is not None and "group" in event_cache: - return event_cache["group"] - group = GroupMemoryCache.get_if_ready(entity.group_id, entity.channel_id) - if event_cache is not None: - event_cache["group"] = group - return group - - -def _module_in_block_string(module: str, value: str | None) -> bool: - if not value: - return False - return f"<{module}," in value - - -def _group_has_plugin_block(group, module: str) -> bool: - if not group: - return False - block_set = getattr(group, "block_plugin_set", None) - super_block_set = getattr(group, "superuser_block_plugin_set", None) - if block_set is not None or super_block_set is not None: - if block_set and module in block_set: - return True - if super_block_set and module in super_block_set: - return True - return False - block_plugin = getattr(group, "block_plugin", "") or "" - super_block_plugin = getattr(group, "superuser_block_plugin", "") or "" - return _module_in_block_string(module, block_plugin) or _module_in_block_string( - module, super_block_plugin - ) - - -def _needs_auth_plugin(plugin: PluginInfo, context: PermissionContext) -> bool: - group = context.group - entity = context.entity - if plugin.block_type == BlockType.ALL and not plugin.status: - if group and getattr(group, "is_super", False): - return False - return True - if entity.group_id: - if plugin.block_type == BlockType.GROUP: - return True - return _group_has_plugin_block(group, plugin.module) - return plugin.block_type == BlockType.PRIVATE - - -def _needs_admin_check(plugin: PluginInfo) -> bool: - if plugin.admin_level and plugin.admin_level > 0: - return True - return plugin.plugin_type in { - PluginType.ADMIN, - PluginType.SUPERUSER, - PluginType.SUPER_AND_ADMIN, - } - - -async def _get_bot_data_cached( - bot_id: str, event_cache -) -> tuple[BotSnapshot | None, bool]: - if event_cache is not None and "bot_data" in event_cache: - return event_cache.get("bot_data"), event_cache.get("bot_timeout", False) - bot = await BotMemoryCache.get(bot_id) - if event_cache is not None: - event_cache["bot_data"] = bot - event_cache["bot_timeout"] = False - return bot, False - - -async def _get_admin_levels_cached( - entity, event_cache -) -> tuple[tuple[LevelUserSnapshot | None, LevelUserSnapshot | None] | None, bool]: - if event_cache is not None and "admin_levels" in event_cache: - return event_cache.get("admin_levels"), event_cache.get("admin_timeout", False) - levels = await LevelUserMemoryCache.get_levels(entity.user_id, entity.group_id) - if event_cache is not None: - event_cache["admin_levels"] = levels - event_cache["admin_timeout"] = False - return levels, False +def _policy_skip_message(reason: str) -> str: + return { + "user_or_group_banned": "user or group banned (cached)", + "superuser_required": "超级管理员权限不足...", + "admin_required": "管理员权限不足...", + "bot_not_found": "Bot不存在,阻断权限检测...", + "bot_sleeping": "Bot休眠中阻断权限检测...", + "bot_plugin_blocked": "Bot插件权限检查结果为关闭...", + "group_not_found": "群组信息不存在...", + "group_blacklisted": "群组黑名单, 目标群组群权限权限-1...", + "group_sleeping": "群组休眠状态...", + "group_level_low": "群等级限制...", + "admin_level_low": "管理员权限不足...", + "plugin_disabled_in_group": "该插件在群组中已被禁用...", + "plugin_superuser_blocked_in_group": "超级管理员禁用了该群此功能...", + "plugin_blocked_in_group": "该群未开启此功能...", + "plugin_disabled_in_private": "该插件在私聊中已被禁用...", + "plugin_global_disabled": "全局未开启此功能...", + }.get(reason, reason or "permission denied") # 超时装饰器 @@ -1232,61 +1272,29 @@ def _is_hidden_plugin(matcher: Matcher) -> bool: return extra.get("plugin_type") == PluginType.HIDDEN -async def _fetch_user_readonly( - user_dao: DataAccess, user_id: str -) -> UserConsole | None: - return await with_timeout( - user_dao.safe_get_or_none(user_id=user_id), name="get_user" - ) - - -async def get_plugin_and_user( +async def _get_plugin_cache_first( module: str, - user_id: str, - platform: str | None = None, - event_cache: dict | None = None, - need_user: bool = True, -) -> tuple[PluginInfo, UserConsole | None]: - """Fetch plugin info and read user only when cost is required.""" - user_dao = DataAccess(UserConsole) - + event_cache: dict | None, + *, + allow_cache_load: bool, +) -> tuple[PluginInfo | None, bool]: plugin = None if event_cache is not None: plugin_cache = event_cache.setdefault("plugin_cache", {}) if module in plugin_cache: - plugin = plugin_cache[module] - if plugin is None: + 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() + if plugin is None and allow_cache_load: plugin = await PluginInfoMemoryCache.get_by_module(module) - if event_cache is not None: - event_cache.setdefault("plugin_cache", {})[module] = plugin - plugin = cast(PluginInfo | None, plugin) - - if not plugin: - raise PermissionExemption(f"plugin:{module} not found, skip permission check") - if plugin.plugin_type == PluginType.HIDDEN: - raise PermissionExemption(f"plugin {plugin.name}:{plugin.module} hidden, skip") - - user = None - if need_user and plugin.cost_gold > 0: - if event_cache is not None: - user_cache = event_cache.setdefault("user_cache", {}) - if user_id in user_cache: - user = user_cache[user_id] - else: - try: - async with _db_section(): - user = await _fetch_user_readonly(user_dao, user_id) - except PermissionExemption: - user = None - user_cache[user_id] = user - else: - try: - async with _db_section(): - user = await _fetch_user_readonly(user_dao, user_id) - except PermissionExemption: - user = None - - return plugin, user + cache_miss = False + if event_cache is not None: + event_cache.setdefault("plugin_cache", {})[module] = plugin + event_cache.setdefault("auth_cache_misses", set()).discard("plugin") + if cache_miss: + event_cache.setdefault("auth_cache_misses", set()).add("plugin") + return plugin, cache_miss async def get_plugin_cost( @@ -1323,43 +1331,29 @@ async def get_plugin_cost( return cost_gold -async def reduce_gold(user_id: str, module: str, cost_gold: int, session: Uninfo): - """扣除用户金币 - - 参数: - user_id: 用户id - module: 插件模块名称 - cost_gold: 消耗金币 - session: Uninfo - """ - should_clear_cache = False +async def reserve_gold( + user_id: str, + module: str, + cost_gold: int, + session: Uninfo, +): + """预扣金币,matcher 未实际完成时由 SideEffectCommit 回滚。""" try: - await with_timeout( - UserConsole.reduce_gold( + reservation = await with_timeout( + UserConsole.reserve_gold( user_id, cost_gold, GoldHandle.PLUGIN, module, PlatformUtils.get_platform(session), ), - name="reduce_gold", + name="reserve_gold", ) except InsufficientGold: - if u := await UserConsole.get_user(user_id): - u.gold = 0 - await u.save(update_fields=["gold"]) - except asyncio.TimeoutError: - should_clear_cache = True - logger.error( - f"扣除金币超时,用户: {user_id}, 金币: {cost_gold}", - LOGGER_COMMAND, - session=session, - ) - - # 正常写入路径由 UserConsole.save() 统一失效缓存;超时状态不确定时兜底清理。 - if should_clear_cache: - await DataAccess(UserConsole).clear_cache(user_id=user_id) - logger.debug(f"调用功能花费金币: {cost_gold}", LOGGER_COMMAND, session=session) + raise + await DataAccess(UserConsole).clear_cache(user_id=user_id) + logger.debug(f"预扣功能花费金币: {cost_gold}", LOGGER_COMMAND, session=session) + return reservation # 辅助函数,用于记录每个 hook 的执行时间 @@ -1383,16 +1377,49 @@ async def time_hook(coro, name, recorder: HookTraceRecorder | None = None): recorder.set(name, f"{time.time() - start:.3f}s") -async def _enter_hooks_section(): - """尝试获取全局信号量并更新计数器,饱和时快速放行。""" +async def _record_backpressure( + *, + lane_context: AuthLaneContext, + reason: str, + 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, + ) + + +async def _enter_hooks_section(lane_context: AuthLaneContext): + """尝试获取全局信号量;过载时记录背压但不丢弃 matcher。""" global HOOKS_ACTIVE_COUNT if HOOKS_SEMAPHORE.locked(): + signal_overload(3.0) + await _record_backpressure( + lane_context=lane_context, + reason="hooks_semaphore_saturated", + action="wait", + ) logger.warning( - "hooks semaphore saturated, allowing pass", + "hooks semaphore saturated, matcher waiting", LOGGER_COMMAND, ) - raise PermissionExemption("hooks semaphore saturated, allow pass") + started = time.perf_counter() await HOOKS_SEMAPHORE.acquire() + wait_ms = (time.perf_counter() - started) * 1000 + if wait_ms >= AUTH_OVERLOAD_LANE_WAIT_MS: + signal_overload(2.0) + await _record_backpressure( + lane_context=lane_context, + reason="hooks_wait_slow", + action="execute", + duration_ms=wait_ms, + ) HOOKS_ACTIVE_COUNT += 1 @@ -1404,33 +1431,307 @@ async def _leave_hooks_section(): HOOKS_ACTIVE_COUNT = max(HOOKS_ACTIVE_COUNT - 1, 0) -async def route_precheck( - matcher: Matcher, +async def _prepare_auth_state( + *, + module: str, context: EventContext, -) -> bool: - 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, + bot: Bot, + event_cache: dict | None, + route_skip_checks: bool, + skip_ban: bool, + hook_recorder: HookTraceRecorder, + state: dict | None, + session: Uninfo, + allow_cache_load: bool = False, +) -> AuthPreparation | None: + plugin_user_start = time.time() + try: + plugin, plugin_cache_miss = await _get_plugin_cache_first( + module, + event_cache, + allow_cache_load=allow_cache_load, ) - set_route_modules(None, context, route_modules) + user = None + if plugin is None: + if not allow_cache_load and plugin_cache_miss: + return None + raise PermissionExemption( + f"plugin:{module} not found, skip permission check" + ) + if plugin.plugin_type == PluginType.HIDDEN: + raise PermissionExemption( + f"plugin {plugin.name}:{plugin.module} hidden, skip" + ) + hook_recorder.set("get_plugin_user", f"{time.time() - plugin_user_start:.3f}s") + except asyncio.TimeoutError: + logger.error( + f"获取插件和用户数据超时,模块: {module}", + LOGGER_COMMAND, + session=session, + ) + return None + except PermissionExemption: + raise - 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 + permission_context = PermissionContext( + event=context, + module=module, + plugin=plugin, + user=user, + ) + store_permission_context(state, permission_context) + + profile = await get_plugin_auth_profile( + plugin, + event_cache=event_cache, + allow_cache_load=allow_cache_load, + ) + snapshot = await get_or_build_auth_snapshot( + context=context, + plugin=plugin, + profile=profile, + bot=bot, + skip_ban=skip_ban, + allow_cache_load=allow_cache_load, + ) + permission_context.group = snapshot.group + permission_context.bot_data = snapshot.bot_data + if snapshot.admin_levels is not None: + permission_context.admin_levels = snapshot.admin_levels + store_permission_context(state, permission_context) + + 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), + ) + return AuthPreparation( + plugin=plugin, + user=user, + profile=profile, + snapshot=snapshot, + permission_context=permission_context, + policy_context=policy_context, + ) + + +async def _prepare_auth_state_with_fallback( + *, + module: str, + context: EventContext, + bot: Bot, + event_cache: dict | None, + route_skip_checks: bool, + skip_ban: bool, + hook_recorder: HookTraceRecorder, + state: dict | None, + session: Uninfo, +) -> AuthPreparation | None: + prep = await _prepare_auth_state( + module=module, + context=context, + bot=bot, + event_cache=event_cache, + route_skip_checks=route_skip_checks, + skip_ban=skip_ban, + hook_recorder=hook_recorder, + state=state, + session=session, + allow_cache_load=False, + ) + if prep is not None: + return prep + hook_recorder.set("auth_snapshot", "cache_miss_fallback") + return await _prepare_auth_state( + module=module, + context=context, + bot=bot, + event_cache=event_cache, + route_skip_checks=route_skip_checks, + skip_ban=skip_ban, + hook_recorder=hook_recorder, + state=state, + session=session, + allow_cache_load=True, + ) + + +async def _check_ban_from_snapshot( + *, + prep: AuthPreparation, + matcher: Matcher, + event_cache: dict | None, + skip_ban: bool, + hook_recorder: HookTraceRecorder, + session: Uninfo, +) -> None: + ban_cache_state = prep.snapshot.ban_state + if event_cache is not None: + ban_cache_state = event_cache.get("ban_state") + if ban_cache_state is True: + hook_recorder.set("auth_ban", "cached") + raise SkipPluginException("user or group banned (cached)") + if ban_cache_state is False: + hook_recorder.set("auth_ban", "cached") + return + if skip_ban: + hook_recorder.set("auth_ban", "skipped") + return + + ban_start = time.time() + try: + await auth_ban( + matcher, + session, + prep.plugin, + context=prep.permission_context, + ) + hook_recorder.set("auth_ban", f"{time.time() - ban_start:.3f}s") + if event_cache is not None: + event_cache["ban_state"] = False + except SkipPluginException: + hook_recorder.set("auth_ban", f"{time.time() - ban_start:.3f}s") + if event_cache is not None: + event_cache["ban_state"] = True + raise + + +async def _reserve_limit_side_effect( + *, + prep: AuthPreparation, + session: Uninfo, + side_effect_commit: SideEffectCommit, +) -> None: + reservation = await reserve_auth_limit( + prep.plugin, + session, + context=prep.permission_context, + ) + await side_effect_commit.reserve_limit(reservation) + + +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: + hook_recorder.set("cost_gold", "skipped") + return 0 + cost_start = time.time() + try: + cost_gold = await with_timeout( + get_plugin_cost( + prep.user, + plugin, + session, + context=prep.permission_context, + ), + name="get_plugin_cost", + ) + hook_recorder.set("cost_gold", f"{time.time() - cost_start:.3f}s") + return cost_gold + except asyncio.TimeoutError: + logger.error( + f"获取插件费用超时,模块: {prep.profile.module}", + LOGGER_COMMAND, + session=session, + ) + return 0 + + +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, +) -> float: + profile = prep.profile + hooks_start = time.time() + + 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, + ) + ) + else: + hook_recorder.set("auth_limit", "skipped") + else: + hook_recorder.set("auth_limit", "skipped") + + if not hook_tasks: + return time.time() - hooks_start + + try: + await with_timeout( + asyncio.gather(*hook_tasks), + timeout=TIMEOUT_SECONDS * 2, + name="auth_hooks_gather", + ) + except asyncio.TimeoutError: + logger.error( + f"权限检查 hooks 总体执行超时,模块: {profile.module}", + LOGGER_COMMAND, + session=session, + ) + finally: + await _leave_hooks_section() + 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, + is_hidden_plugin=_is_hidden_plugin, + is_command_matcher_class=_is_command_matcher_class, + matcher_has_alconna_shortcuts=_matcher_has_alconna_shortcuts, + prepare_auth_state_with_fallback=_prepare_auth_state_with_fallback, + prepare_auth_state=_prepare_auth_state, + policy_decision_point=_AUTH_PDP, + policy_skip_message=_policy_skip_message, + legacy_pure_auth_fallback=legacy_pure_auth_fallback, + check_ban_from_snapshot=_check_ban_from_snapshot, + resolve_cost_gold=_resolve_cost_gold, + 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, +) +_AUTH_PIPELINE = build_auth_pipeline(_AUTH_PIPELINE_DEPS) async def auth( @@ -1453,426 +1754,92 @@ async def auth( context: EventContext """ start_time = time.time() - cost_gold = 0 - ignore_flag = False entity = context.entity event_cache = context.event_cache text = context.plain_text - is_superuser = context.is_superuser route_modules = context.route_modules if context.route_modules_loaded else None module = matcher.plugin_name or "" is_command_matcher = _is_command_matcher_class(type(matcher)) - auth_allowed = None - auth_result_cache = None - admin_checked_pre = False - permission_context: PermissionContext | None = None + lane_context = _auth_lane_context_from_state(type(matcher), context, state) side_effect_cache = get_permission_side_effect_cache( state=state, event_cache=event_cache, ) - side_effect_lock = None - entered_side_effect_lock = False + side_effect_commit = SideEffectCommit( + session=session, + module=module, + owner_matcher_id=id(matcher), + limit_entity=entity, + ) # 仅在慢请求时记录 hook 明细,避免热路径高频构造字符串 hook_recorder = HookTraceRecorder(start_time) - hooks_time = 0 # 初始化 hooks_time 变量 - - # 记录是否已进入 hooks 区域(用于 finally 中释放) - entered_hooks = False + pipeline_context = AuthPipelineContext( + matcher=matcher, + event=event, + bot=bot, + session=session, + event_context=context, + skip_ban=skip_ban, + state=state, + start_time=start_time, + module=module, + entity=entity, + event_cache=event_cache, + text=text, + route_modules=route_modules, + is_command_matcher=is_command_matcher, + lane_context=lane_context, + side_effect_cache=side_effect_cache, + side_effect_commit=side_effect_commit, + hook_recorder=hook_recorder, + ) try: - if not module: - auth_allowed = True - return - - side_effect_lock = side_effect_cache.lock_for(module) - await side_effect_lock.acquire() - entered_side_effect_lock = True - - auth_result_cache = side_effect_cache.auth_results - cached_result = auth_result_cache.get(module) - if cached_result is not None: - allowed, reason = cached_result - if not allowed: - raise SkipPluginException(reason or "auth cached skip") - return - - if _is_hidden_plugin(matcher): - auth_allowed = True - return - if event_cache is not None and event_cache.get("ban_state") is True: - raise SkipPluginException("user or group banned (cached)") - - if route_modules is None: - route_modules = await _get_route_context(text, event_cache) - set_route_modules(state, context, route_modules) - route_skip_checks = ( - is_command_matcher - and module in _ROUTE_MODULES_WITH_COMMANDS - and module not in route_modules - and not _matcher_has_alconna_shortcuts(type(matcher)) - ) - if route_skip_checks: - if event_cache is not None: - event_cache["route_skip"] = True - hook_recorder.set("route", "miss") - auth_allowed = True - return - - platform = context.platform - # 获取插件和用户数据 - plugin_user_start = time.time() - try: - plugin, user = await with_timeout( - get_plugin_and_user( - module, - entity.user_id, - platform, - event_cache=event_cache, - need_user=not route_skip_checks, - ), - name="get_plugin_and_user", - ) - hook_recorder.set( - "get_plugin_user", f"{time.time() - plugin_user_start:.3f}s" - ) - except asyncio.TimeoutError: - logger.error( - f"获取插件和用户数据超时,模块: {module}", - LOGGER_COMMAND, - session=session, - ) - auth_allowed = True - return - - permission_context = PermissionContext( - event=context, - module=module, - plugin=plugin, - user=user, - ) - store_permission_context(state, permission_context) - - if not route_skip_checks and _needs_admin_check(plugin): - if plugin.plugin_type in { - PluginType.SUPERUSER, - PluginType.SUPER_AND_ADMIN, - }: - if is_superuser: - hook_recorder.set("auth_admin", "superuser") - admin_checked_pre = True - elif plugin.plugin_type == PluginType.SUPERUSER: - raise SkipPluginException("超级管理员权限不足...") - if not admin_checked_pre: - if event_cache is not None and event_cache.get("admin_precheck_done"): - hook_recorder.set("auth_admin", "precheck") - admin_checked_pre = True - else: - await LevelUserMemoryCache.ensure_fresh() - admin_levels = None - admin_timeout = False - if event_cache is not None: - admin_levels, admin_timeout = await _get_admin_levels_cached( - entity, event_cache - ) - permission_context.admin_levels = admin_levels - if admin_timeout: - hook_recorder.set("auth_admin", "timeout") - else: - admin_start = time.time() - await auth_admin( - plugin, - session, - context=permission_context, - ) - hook_recorder.set( - "auth_admin", f"{time.time() - admin_start:.3f}s(pre)" - ) - admin_checked_pre = True - - ban_cache_state = None - if event_cache is not None: - ban_cache_state = event_cache.get("ban_state") - if ban_cache_state is True: - hook_recorder.set("auth_ban", "cached") - raise SkipPluginException("user or group banned (cached)") - if ban_cache_state is False: - hook_recorder.set("auth_ban", "cached") - elif ban_cache_state is None: - if skip_ban: - hook_recorder.set("auth_ban", "skipped") - else: - ban_start = time.time() - try: - await auth_ban( - matcher, - session, - plugin, - context=permission_context, - ) - hook_recorder.set("auth_ban", f"{time.time() - ban_start:.3f}s") - if event_cache is not None: - event_cache["ban_state"] = False - except SkipPluginException: - hook_recorder.set("auth_ban", f"{time.time() - ban_start:.3f}s") - if event_cache is not None: - event_cache["ban_state"] = True - raise - - # 获取插件费用 - if not route_skip_checks and plugin.cost_gold > 0: - cost_start = time.time() - try: - cost_gold = await with_timeout( - get_plugin_cost( - user, - plugin, - session, - context=permission_context, - ), - name="get_plugin_cost", - ) - hook_recorder.set("cost_gold", f"{time.time() - cost_start:.3f}s") - except asyncio.TimeoutError: - logger.error( - f"获取插件费用超时,模块: {module}", LOGGER_COMMAND, session=session - ) - # 继续执行,不阻止权限检查 - else: - hook_recorder.set("cost_gold", "skipped") - - # 执行 bot_filter - bot_filter(session, context=permission_context) - - group = await _get_group_cached(entity, event_cache) - - bot_data = None - bot_timeout = False - if event_cache is not None: - bot_data, bot_timeout = await _get_bot_data_cached(bot.self_id, event_cache) - - admin_levels = None - admin_timeout = False - if ( - not admin_checked_pre - and plugin.admin_level - and event_cache is not None - and not route_skip_checks - ): - admin_levels, admin_timeout = await _get_admin_levels_cached( - entity, event_cache - ) - - permission_context.group = group - permission_context.bot_data = bot_data - if admin_levels is not None: - permission_context.admin_levels = admin_levels - store_permission_context(state, permission_context) - - # 并行执行所有 hook 检查,并记录执行时间 - hooks_start = time.time() - allow_sleep_bypass = _is_bot_wake_command(module, text) - - # 先进入 hooks 并行检查区域;饱和时快速放行,避免创建并积压协程。 - await _enter_hooks_section() - entered_hooks = True - - # 创建所有 hook 任务 - hook_tasks = [] - if event_cache is None: - hook_tasks.append( - time_hook( - auth_bot( - plugin, - bot.self_id, - allow_sleep_bypass=allow_sleep_bypass, - context=permission_context, - ), - "auth_bot", - hook_recorder, - ) - ) - else: - if bot_timeout: - hook_recorder.set("auth_bot", "timeout") - else: - hook_tasks.append( - time_hook( - auth_bot( - plugin, - bot.self_id, - bot_data=bot_data, - skip_fetch=True, - allow_sleep_bypass=allow_sleep_bypass, - context=permission_context, - ), - "auth_bot", - hook_recorder, - ) - ) - - if is_superuser: - hook_recorder.set("auth_group", "superuser") - else: - hook_tasks.append( - time_hook( - auth_group( - plugin, - group, - text, - entity.group_id, - context=permission_context, - ), - "auth_group", - hook_recorder, - ) - ) - - if not route_skip_checks and plugin.admin_level and not admin_checked_pre: - if event_cache is None: - hook_tasks.append( - time_hook( - auth_admin(plugin, session, context=permission_context), - "auth_admin", - hook_recorder, - ) - ) - else: - if admin_timeout: - hook_recorder.set("auth_admin", "timeout") - else: - hook_tasks.append( - time_hook( - auth_admin( - plugin, - session, - context=permission_context, - ), - "auth_admin", - hook_recorder, - ) - ) - else: - hook_recorder.setdefault("auth_admin", "skipped") - - if is_superuser: - hook_recorder.set("auth_plugin", "superuser") - elif not route_skip_checks and _needs_auth_plugin(plugin, permission_context): - hook_tasks.append( - time_hook( - auth_plugin( - plugin, - group, - session, - event, - context=permission_context, - skip_group_block=is_superuser, - ), - "auth_plugin", - hook_recorder, - ) - ) - else: - hook_recorder.set("auth_plugin", "skipped") - - if not route_skip_checks: - has_limits = await _has_limits_cached(module, event_cache) - if has_limits: - hook_tasks.append( - time_hook( - auth_limit(plugin, session, context=permission_context), - "auth_limit", - hook_recorder, - ) - ) - else: - hook_recorder.set("auth_limit", "skipped") - else: - hook_recorder.set("auth_limit", "skipped") - - # 使用 gather 并行执行所有 hook,但添加总体超时控制 - try: - await with_timeout( - asyncio.gather(*hook_tasks), - timeout=TIMEOUT_SECONDS * 2, # 给总体执行更多时间 - name="auth_hooks_gather", - ) - except asyncio.TimeoutError: - logger.error( - f"权限检查 hooks 总体执行超时,模块: {module}", - LOGGER_COMMAND, - session=session, - ) - # 不抛出异常,允许继续执行 - - hooks_time = time.time() - hooks_start - auth_allowed = True + await _AUTH_PIPELINE.run(pipeline_context) except SkipPluginException as e: LimitManager.unblock(module, entity.user_id, entity.group_id, entity.channel_id) + await side_effect_commit.rollback_all("auth_skip") if e.tip_message: - try: - tip_coro = send_message( - session, - e.tip_message, - e.tip_check_tag, - background=e.tip_background, - ) - if e.tip_timeout and not e.tip_background: - await asyncio.wait_for(tip_coro, timeout=e.tip_timeout) - else: - await tip_coro - except asyncio.TimeoutError: - logger.error("发送权限提示超时", LOGGER_COMMAND, session=session) + await side_effect_commit.send_permission_tip( + e.tip_message, + e.tip_check_tag, + background=e.tip_background, + timeout=e.tip_timeout, + ) logger.info(str(e), LOGGER_COMMAND, session=session) - ignore_flag = True - auth_allowed = False + pipeline_context.ignore_flag = True + pipeline_context.auth_allowed = False + pipeline_context.decision_effect = "defer" if "deferred" in str(e) else "skip" + pipeline_context.decision_reason = str(e) or "skip_plugin" except IsSuperuserException: logger.debug("超级用户跳过权限检测...", LOGGER_COMMAND, session=session) - auth_allowed = True + pipeline_context.auth_allowed = True + pipeline_context.decision_effect = "allow" + pipeline_context.decision_reason = "superuser" except PermissionExemption as e: + await side_effect_commit.rollback_all("permission_exemption") logger.info(str(e), LOGGER_COMMAND, session=session) - auth_allowed = True + pipeline_context.auth_allowed = True + pipeline_context.decision_effect = "allow" + pipeline_context.decision_reason = str(e) or "permission_exemption" + except Exception: + await side_effect_commit.rollback_all("auth_exception") + raise finally: - # 如果进入过 hooks 区域,确保释放信号量(即使上层处理抛出了异常) - if entered_hooks: - try: - await _leave_hooks_section() - except Exception: - logger.error( - "释放 hooks 信号量时出错", - LOGGER_COMMAND, - session=session, - ) - if auth_result_cache is not None and auth_allowed is not None: - auth_result_cache[module] = (auth_allowed, None) - if entered_side_effect_lock and side_effect_lock is not None: - with contextlib.suppress(Exception): - side_effect_lock.release() - # 扣除金币 - if not ignore_flag and cost_gold > 0: - gold_start = time.time() - try: - await with_timeout( - reduce_gold(entity.user_id, module, cost_gold, session), - name="reduce_gold", - ) - hook_recorder.set("reduce_gold", f"{time.time() - gold_start:.3f}s") - except asyncio.TimeoutError: - logger.error( - f"扣除金币超时,模块: {module}", LOGGER_COMMAND, session=session - ) + await decision_log_stage(pipeline_context, _AUTH_PIPELINE_DEPS) # 记录总执行时间 total_time = time.time() - start_time if total_time > WARNING_THRESHOLD: # 如果总时间超过500ms,记录详细信息 logger.warning( f"权限检查耗时过长: {total_time:.3f}s, 模块: {module}, " - f"hooks时间: {hooks_time:.3f}s, " + f"hooks时间: {pipeline_context.hooks_time:.3f}s, " f"详情: {hook_recorder.snapshot()}", LOGGER_COMMAND, session=session, ) - if ignore_flag: + if pipeline_context.ignore_flag: raise IgnoredException("权限检测 ignore") diff --git a/zhenxun/builtin_plugins/hooks/auth_event_selector.py b/zhenxun/builtin_plugins/hooks/auth_event_selector.py new file mode 100644 index 00000000..811c7ebd --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth_event_selector.py @@ -0,0 +1,271 @@ +from __future__ import annotations + +from collections.abc import Awaitable, Callable +import contextlib +from dataclasses import dataclass +import importlib +from typing import Any + +from nonebot.adapters import Bot, Event +from nonebot.matcher import Matcher +import nonebot.message as nb_message + +from zhenxun.services.log import logger +from zhenxun.services.message_load import signal_overload + +from .auth.config import LOGGER_COMMAND +from .auth_activation import HandlerActivationIndex +from .auth_patch_guard import validate_handle_event_patch +from .auth_types import EventDispatchContext + + +@dataclass(slots=True) +class HandleEventSelectorDependencies: + activation_index: HandlerActivationIndex + overload_selected_threshold: int + prepare_handle_event_state: Callable[[Event, dict], None] + build_dispatch_context: Callable[ + [Event, dict | None], + Awaitable[EventDispatchContext], + ] + 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]] + + +_HANDLE_EVENT_PATCHED = False +_ORIGINAL_HANDLE_EVENT: Callable[..., Awaitable[None]] | None = None +_ORIGINAL_ADAPTER_HANDLE_EVENTS: dict[object, object] = {} + + +async def patched_handle_event( + bot: Bot, + event: Event, + deps: HandleEventSelectorDependencies, +) -> None: + show_log = True + escape_tag = getattr(nb_message, "escape_tag") + logger_ = getattr(nb_message, "logger") + no_log_exception = getattr(nb_message, "NoLogException") + + log_msg = f"{escape_tag(bot.type)} {escape_tag(bot.self_id)} | " + try: + log_msg += event.get_log_string() + except no_log_exception: + show_log = False + if show_log: + logger_.opt(colors=True).success(log_msg) + + state = {} + dependency_cache = {} + async_exit_stack = getattr(nb_message, "AsyncExitStack") + apply_event_preprocessors = getattr(nb_message, "_apply_event_preprocessors") + apply_event_postprocessors = getattr(nb_message, "_apply_event_postprocessors") + trie_rule = getattr(nb_message, "TrieRule") + matchers = getattr(nb_message, "matchers") + catch = getattr(nb_message, "catch") + stop_propagation = getattr(nb_message, "StopPropagation") + handle_exception = getattr(nb_message, "_handle_exception") + anyio_mod = getattr(nb_message, "anyio") + run_coro_with_shield = getattr(nb_message, "run_coro_with_shield") + + async with async_exit_stack() as stack: + if not await apply_event_preprocessors( + bot=bot, + event=event, + state=state, + stack=stack, + dependency_cache=dependency_cache, + ): + return + + try: + trie_rule.get_value(bot, event, state) + except Exception as e: + logger_.opt(colors=True, exception=e).warning( + "Error while parsing command for event" + ) + deps.prepare_handle_event_state(event, state) + dispatch_context = await deps.build_dispatch_context(event, state) + activation_context = deps.activation_context_from_dispatch( + dispatch_context, + event, + ) + activation_available = True + try: + deps.activation_index.ensure_fresh(matchers) + except Exception as exc: + activation_available = False + logger.warning( + "HandlerActivationIndex 构建失败,回退到旧 matcher 选择逻辑", + LOGGER_COMMAND, + e=exc, + ) + + break_flag = False + + def _handle_stop_propagation(_exc_group) -> None: + nonlocal break_flag + break_flag = True + logger_.debug("Stop event propagation") + + for priority in sorted(matchers.keys()): + if break_flag: + break + + if show_log: + logger_.debug(f"Checking for matchers in priority {priority}...") + + if not (priority_matchers := matchers[priority]): + continue + + with catch( + { + stop_propagation: _handle_stop_propagation, + Exception: handle_exception( + "Error when checking Matcher." + ), + } + ): + priority_budget = deps.new_dispatch_budget() + if activation_available: + try: + activation_result = deps.activation_index.select_priority( + priority, + priority_matchers, + activation_context, + priority_budget, + ) + except Exception as exc: + logger.warning( + "HandlerActivationIndex 选择失败,当前 priority 回退", + LOGGER_COMMAND, + e=exc, + ) + activation_result = None + else: + activation_result = None + + 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 + ): + signal_overload(3.0) + else: + selected_matchers = priority_matchers + + async with anyio_mod.create_task_group() as tg: + for matcher in selected_matchers: + lane = deps.dispatch_lane_for_matcher(matcher, dispatch_context) + if activation_result is None: + descriptor = deps.activation_index.descriptor_for(matcher) + if descriptor is not None: + single_budget = dict(priority_budget) + try: + single_result = ( + deps.activation_index.select_priority( + priority, + [matcher], + activation_context, + single_budget, + ) + ) + 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, + ) + if not single_result.selected: + continue + matcher_state = deps.build_matcher_state(state) + tg.start_soon( + run_coro_with_shield, + deps.run_selected_matcher( + matcher, + bot, + event, + matcher_state, + stack, + dependency_cache, + lane, + ), + ) + + if show_log: + logger_.debug("Checking for matchers completed") + + await apply_event_postprocessors(bot, event, state, stack, dependency_cache) + + +def install_handle_event_selector(deps: HandleEventSelectorDependencies) -> None: + global _HANDLE_EVENT_PATCHED, _ORIGINAL_HANDLE_EVENT + if _HANDLE_EVENT_PATCHED: + return + guard = validate_handle_event_patch() + if not guard.ok: + logger.warning( + f"权限事件分发选择器 patch 未安装,回退 NoneBot 原生分发: {guard.reason}", + LOGGER_COMMAND, + ) + return + _ORIGINAL_HANDLE_EVENT = nb_message.handle_event + + async def _patched(bot: Bot, event: Event) -> None: + await patched_handle_event(bot, event, deps) + + nb_message.handle_event = _patched # type: ignore[assignment] + for module_name in ( + "nonebot.adapters.onebot.v11.bot", + "nonebot.adapters.onebot.v12.bot", + "onebug.mixin.process", + ): + with contextlib.suppress(Exception): + module = importlib.import_module(module_name) + current = getattr(module, "handle_event", None) + if current is not None: + _ORIGINAL_ADAPTER_HANDLE_EVENTS[module] = current + setattr(module, "handle_event", _patched) + _HANDLE_EVENT_PATCHED = True + + +def uninstall_handle_event_selector() -> None: + global _HANDLE_EVENT_PATCHED, _ORIGINAL_HANDLE_EVENT + if not _HANDLE_EVENT_PATCHED: + return + if _ORIGINAL_HANDLE_EVENT is not None: + nb_message.handle_event = _ORIGINAL_HANDLE_EVENT # type: ignore[assignment] + for module, original in list(_ORIGINAL_ADAPTER_HANDLE_EVENTS.items()): + with contextlib.suppress(Exception): + setattr(module, "handle_event", original) + _ORIGINAL_ADAPTER_HANDLE_EVENTS.clear() + _HANDLE_EVENT_PATCHED = False + _ORIGINAL_HANDLE_EVENT = None + + +__all__ = [ + "HandleEventSelectorDependencies", + "install_handle_event_selector", + "patched_handle_event", + "uninstall_handle_event_selector", +] diff --git a/zhenxun/builtin_plugins/hooks/auth_hook.py b/zhenxun/builtin_plugins/hooks/auth_hook.py index 94c08da4..c6fe0566 100644 --- a/zhenxun/builtin_plugins/hooks/auth_hook.py +++ b/zhenxun/builtin_plugins/hooks/auth_hook.py @@ -18,6 +18,7 @@ from .auth.config import LOGGER_COMMAND from .auth.context import ( get_event_context, get_or_create_event_context, + get_permission_side_effect_cache, resolve_actor_user_id, resolve_event_channel_id, resolve_event_group_id, @@ -25,12 +26,13 @@ from .auth.context import ( ) from .auth_checker import ( LimitManager, + _get_auth_route_precheck_deps, _get_route_context, auth, - route_precheck, 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 @@ -111,7 +113,7 @@ async def _auth_preprocessor( ) set_route_modules(state, event_context, route_modules) - if await route_precheck(matcher, event_context): + if await route_precheck(matcher, event_context, **_get_auth_route_precheck_deps()): return try: @@ -146,6 +148,7 @@ async def _unblock_after_matcher( session: Uninfo, event: Event, state: T_State, + exception: Exception | None = None, ): context = get_event_context(state) if context is not None: @@ -164,4 +167,37 @@ async def _unblock_after_matcher( group_id = session.group.id if user_id and matcher.plugin: module = matcher.plugin.name - LimitManager.unblock(module, user_id, group_id, channel_id) + side_effects = get_permission_side_effect_cache( + state=state, + event_cache=context.event_cache if context is not None else None, + ) + commit = side_effects.commits.get(module) + if ( + commit is not None + and not commit.committed + and commit.owner_matcher_id == id(matcher) + ): + side_effects.commits.pop(module, None) + if exception is None: + try: + await commit.commit_all() + side_effects.auth_results[module] = (True, None) + except Exception as exc: + await commit.rollback_all("commit_failed") + logger.error( + "auth side effect commit failed", + LOGGER_COMMAND, + e=exc, + ) + else: + await commit.rollback_all("matcher_exception") + if commit.limit_should_auto_unblock: + limit_entity = commit.limit_entity + LimitManager.unblock( + module, + limit_entity.user_id if limit_entity else user_id, + limit_entity.group_id if limit_entity else group_id, + limit_entity.channel_id if limit_entity else channel_id, + ) + else: + LimitManager.unblock(module, user_id, group_id, channel_id) diff --git a/zhenxun/builtin_plugins/hooks/auth_legacy_fallback.py b/zhenxun/builtin_plugins/hooks/auth_legacy_fallback.py new file mode 100644 index 00000000..aee36ebd --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth_legacy_fallback.py @@ -0,0 +1,54 @@ +from __future__ import annotations + +from nonebot.adapters import Event +from nonebot_plugin_uninfo import Uninfo + +from .auth.auth_admin import auth_admin +from .auth.auth_bot import auth_bot +from .auth.auth_group import auth_group +from .auth.auth_plugin import auth_plugin +from .auth_types import AuthPreparation + + +async def legacy_pure_auth_fallback( + *, + prep: AuthPreparation, + event: Event, + session: Uninfo, + text: str, +) -> None: + """Compatibility fallback for cache-deferred pure permission checks.""" + + await auth_bot( + prep.plugin, + prep.snapshot.context.bot_id, + prep.snapshot.bot_data, + skip_fetch=prep.snapshot.bot_data is not None, + allow_sleep_bypass=prep.policy_context.allow_sleep_bypass, + context=prep.permission_context, + ) + await auth_group( + prep.plugin, + prep.snapshot.group, + text, + prep.snapshot.group_id, + context=prep.permission_context, + ) + await auth_plugin( + prep.plugin, + prep.snapshot.group, + session, + event, + context=prep.permission_context, + user_id=prep.snapshot.user_id, + ) + await auth_admin( + prep.plugin, + session, + cached_levels=prep.snapshot.admin_levels, + context=prep.permission_context, + entity=prep.snapshot.context.entity, + ) + + +__all__ = ["legacy_pure_auth_fallback"] diff --git a/zhenxun/builtin_plugins/hooks/auth_patch_guard.py b/zhenxun/builtin_plugins/hooks/auth_patch_guard.py new file mode 100644 index 00000000..518e711d --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth_patch_guard.py @@ -0,0 +1,61 @@ +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass +import inspect +from typing import Any + +import nonebot.message as nb_message + + +@dataclass(frozen=True, slots=True) +class AuthPatchGuardResult: + ok: bool + reason: str = "" + + +_HANDLE_EVENT_PARAMS = {"bot", "event"} +_HANDLE_EVENT_REQUIRED_ATTRS = ( + "escape_tag", + "logger", + "NoLogException", + "AsyncExitStack", + "_apply_event_preprocessors", + "_apply_event_postprocessors", + "TrieRule", + "matchers", + "catch", + "StopPropagation", + "_handle_exception", + "anyio", + "run_coro_with_shield", +) + + +def _signature_param_names(func: Callable[..., Any]) -> set[str]: + return set(inspect.signature(func).parameters) + + +def validate_handle_event_patch() -> AuthPatchGuardResult: + target = getattr(nb_message, "handle_event", None) + if target is None: + return AuthPatchGuardResult(False, "missing nonebot.message.handle_event") + try: + params = _signature_param_names(target) + except Exception as exc: + return AuthPatchGuardResult(False, f"inspect signature failed: {exc}") + missing_params = sorted(_HANDLE_EVENT_PARAMS - params) + if missing_params: + return AuthPatchGuardResult( + False, + "handle_event signature missing params: " + ", ".join(missing_params), + ) + missing_attrs = [ + attr for attr in _HANDLE_EVENT_REQUIRED_ATTRS if not hasattr(nb_message, attr) + ] + if missing_attrs: + return AuthPatchGuardResult( + False, + "nonebot.message missing attrs: " + ", ".join(missing_attrs), + ) + return AuthPatchGuardResult(True) diff --git a/zhenxun/builtin_plugins/hooks/auth_pipeline.py b/zhenxun/builtin_plugins/hooks/auth_pipeline.py new file mode 100644 index 00000000..a5513089 --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth_pipeline.py @@ -0,0 +1,480 @@ +from __future__ import annotations + +import asyncio +from collections.abc import Awaitable, Callable +from dataclasses import dataclass, field +import time +from typing import TYPE_CHECKING, Any + +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 ( + EventContext, + PermissionSideEffectCache, + set_route_modules, +) +from .auth.exception import PermissionExemption, SkipPluginException +from .auth_policy import ( + action_from_snapshot, + principal_from_snapshot, + raise_for_policy, + resource_from_snapshot, +) +from .auth_types import AuthLaneContext, AuthPolicyFlags, AuthPreparation + +if TYPE_CHECKING: + from .auth_side_effect import SideEffectCommit + from .auth_trace import HookTraceRecorder + + +def _require(value: Any, name: str): + if value is None: + raise RuntimeError(f"AuthPipelineContext.{name} is required") + return value + + +def _prep(ctx: AuthPipelineContext) -> AuthPreparation: + return _require(ctx.prep, "prep") + + +def _recorder(ctx: AuthPipelineContext) -> "HookTraceRecorder": + return _require(ctx.hook_recorder, "hook_recorder") + + +def _side_effect_commit(ctx: AuthPipelineContext) -> "SideEffectCommit": + return _require(ctx.side_effect_commit, "side_effect_commit") + + +def _side_effect_cache(ctx: AuthPipelineContext) -> PermissionSideEffectCache: + return _require(ctx.side_effect_cache, "side_effect_cache") + + +def _entity(ctx: AuthPipelineContext) -> EntityIDs: + return _require(ctx.entity, "entity") + + +def _lane_context(ctx: AuthPipelineContext) -> AuthLaneContext: + return _require(ctx.lane_context, "lane_context") + + +PipelineHandler = Callable[["AuthPipelineContext"], Awaitable[None]] + + +@dataclass(slots=True) +class AuthPipelineStage: + name: str + handler: PipelineHandler + + +@dataclass(slots=True) +class AuthPipelineContext: + matcher: Matcher + event: Event + bot: Bot + session: Uninfo + event_context: EventContext + skip_ban: bool = False + state: dict | None = None + start_time: float = field(default_factory=time.time) + module: str = "" + entity: EntityIDs | None = None + 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 + side_effect_commit: "SideEffectCommit | None" = None + side_effect_lock: asyncio.Lock | None = None + entered_side_effect_lock: bool = False + auth_result_cache: dict | None = None + hook_recorder: "HookTraceRecorder | None" = None + prep: AuthPreparation | None = None + flags: AuthPolicyFlags | None = None + cost_gold: int = 0 + hooks_time: float = 0.0 + ignore_flag: bool = False + auth_allowed: bool | None = None + decision_effect: str | None = None + decision_reason: str | None = None + stopped: bool = False + stage_timings: dict[str, float] = field(default_factory=dict) + + def stop( + self, + *, + allowed: bool, + effect: str, + reason: str, + ) -> None: + self.auth_allowed = allowed + self.decision_effect = effect + self.decision_reason = reason + self.stopped = True + + +class AuthPipeline: + def __init__(self, stages: list[AuthPipelineStage]) -> None: + self._stages = tuple(stages) + + async def run(self, context: AuthPipelineContext) -> None: + for stage in self._stages: + started = time.perf_counter() + await stage.handler(context) + context.stage_timings[stage.name] = (time.perf_counter() - started) * 1000 + if context.stopped: + break + + +@dataclass(slots=True) +class AuthPipelineDependencies: + route_modules_with_commands: set[str] + get_route_context: Callable[[str, dict | None], Awaitable[set[str]]] + is_hidden_plugin: Callable[[Matcher], bool] + is_command_matcher_class: Callable[[type[Matcher]], bool] + matcher_has_alconna_shortcuts: Callable[[type[Matcher]], bool] + prepare_auth_state_with_fallback: Callable[..., Awaitable[Any]] + prepare_auth_state: Callable[..., Awaitable[Any]] + policy_decision_point: Any + policy_skip_message: Callable[[str], str] + legacy_pure_auth_fallback: Callable[..., Awaitable[None]] + check_ban_from_snapshot: Callable[..., Awaitable[None]] + resolve_cost_gold: Callable[..., Awaitable[int]] + 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 + + +def apply_policy_precheck( + ctx: AuthPipelineContext, + deps: AuthPipelineDependencies, +) -> AuthPolicyFlags: + prep = _prep(ctx) + hook_recorder = _recorder(ctx) + flags = AuthPolicyFlags() + snapshot = prep.snapshot + decision = deps.policy_decision_point.decide( + principal_from_snapshot(snapshot), + action_from_snapshot(snapshot), + resource_from_snapshot(snapshot), + prep.policy_context, + ) + if decision.deferred: + 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", + }: + flags.should_return_allowed = True + return flags + + bot_decision = deps.policy_decision_point.decide_bot(prep.policy_context) + if bot_decision.allowed: + hook_recorder.set("auth_bot", "policy") + elif bot_decision.denied: + raise_for_policy(bot_decision, deps.policy_skip_message(bot_decision.reason)) + elif bot_decision.deferred: + raise PermissionExemption(f"auth_bot deferred: {bot_decision.reason}") + + group_decision = deps.policy_decision_point.decide_group(prep.policy_context) + if group_decision.allowed or group_decision.skipped: + hook_recorder.set("auth_group", f"policy:{group_decision.reason}") + elif group_decision.denied: + raise_for_policy( + group_decision, + deps.policy_skip_message(group_decision.reason), + ) + elif group_decision.deferred: + raise PermissionExemption(f"auth_group deferred: {group_decision.reason}") + + plugin_decision = deps.policy_decision_point.decide_plugin(prep.policy_context) + if plugin_decision.allowed or plugin_decision.skipped: + hook_recorder.set("auth_plugin", f"policy:{plugin_decision.reason}") + elif plugin_decision.denied: + raise_for_policy( + plugin_decision, + deps.policy_skip_message(plugin_decision.reason), + ) + else: + raise PermissionExemption(f"auth_plugin deferred: {plugin_decision.reason}") + + admin_decision = deps.policy_decision_point.decide_admin(prep.policy_context) + if admin_decision.allowed or admin_decision.skipped: + hook_recorder.set("auth_admin", f"policy:{admin_decision.reason}") + elif admin_decision.denied: + raise_for_policy( + admin_decision, + deps.policy_skip_message(admin_decision.reason), + ) + else: + raise PermissionExemption(f"auth_admin deferred: {admin_decision.reason}") + + return flags + + +async def route_gate_stage( + ctx: AuthPipelineContext, + deps: AuthPipelineDependencies, +) -> None: + if not ctx.module: + ctx.stop(allowed=True, effect="allow", reason="empty_module") + return + + side_effect_cache = _side_effect_cache(ctx) + ctx.side_effect_lock = side_effect_cache.lock_for(ctx.module) + await ctx.side_effect_lock.acquire() + ctx.entered_side_effect_lock = True + + auth_result_cache = side_effect_cache.auth_results + ctx.auth_result_cache = auth_result_cache + cached_result = auth_result_cache.get(ctx.module) + if cached_result is not None: + allowed, reason = cached_result + if not allowed: + ctx.decision_effect = "skip" + ctx.decision_reason = reason or "auth_cached_skip" + raise SkipPluginException(reason or "auth cached skip") + ctx.stop(allowed=True, effect="allow", reason="auth_cached_allow") + return + + if deps.is_hidden_plugin(ctx.matcher): + ctx.stop(allowed=True, effect="allow", reason="hidden_plugin") + return + if ctx.event_cache is not None and ctx.event_cache.get("ban_state") is True: + ctx.decision_effect = "skip" + ctx.decision_reason = "ban_cached" + raise SkipPluginException("user or group banned (cached)") + + 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 = ( + 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 ctx.event_cache is not None: + ctx.event_cache["route_skip"] = True + _recorder(ctx).set("route", "miss") + ctx.stop(allowed=True, effect="allow", reason="route_miss_skip_checks") + + +async def prepare_snapshot_stage( + ctx: AuthPipelineContext, + deps: AuthPipelineDependencies, +) -> None: + ctx.prep = await deps.prepare_auth_state_with_fallback( + module=ctx.module, + 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, + session=ctx.session, + ) + if ctx.prep is None: + ctx.stop(allowed=True, effect="allow", reason="prepare_timeout_allow") + + +async def policy_precheck_stage( + ctx: AuthPipelineContext, + deps: AuthPipelineDependencies, +) -> None: + try: + ctx.flags = apply_policy_precheck(ctx, deps) + except PermissionExemption as exc: + _recorder(ctx).set("policy_fallback", str(exc)) + ctx.prep = await deps.prepare_auth_state( + module=ctx.module, + 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, + session=ctx.session, + allow_cache_load=True, + ) + if ctx.prep is None: + ctx.stop(allowed=True, effect="allow", reason="policy_fallback_timeout") + return + try: + ctx.flags = apply_policy_precheck(ctx, deps) + except PermissionExemption as fallback_exc: + _recorder(ctx).set("legacy_pure_auth", str(fallback_exc)) + await deps.legacy_pure_auth_fallback( + prep=ctx.prep, + event=ctx.event, + session=ctx.session, + text=ctx.text, + ) + ctx.flags = AuthPolicyFlags() + flags = _require(ctx.flags, "flags") + if flags.should_return_allowed: + ctx.stop(allowed=True, effect="allow", reason="policy_precheck_allow") + return + await deps.check_ban_from_snapshot( + prep=ctx.prep, + matcher=ctx.matcher, + event_cache=ctx.event_cache, + skip_ban=ctx.skip_ban, + hook_recorder=ctx.hook_recorder, + session=ctx.session, + ) + 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, + ) + + +async def legacy_hook_adapter_stage( + ctx: AuthPipelineContext, + deps: AuthPipelineDependencies, +) -> None: + prep = _prep(ctx) + deps.bot_filter(ctx.session, context=prep.permission_context) + ctx.hooks_time = await deps.run_auth_hooks( + 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), + ) + ctx.auth_allowed = True + ctx.decision_effect = "allow" + ctx.decision_reason = "auth_passed" + + +async def side_effect_commit_stage( + ctx: AuthPipelineContext, + deps: AuthPipelineDependencies, +) -> None: + commit = _side_effect_commit(ctx) + side_effect_cache = _side_effect_cache(ctx) + if ctx.ignore_flag: + await commit.rollback_all("auth_ignored") + return + if ctx.cost_gold <= 0: + if commit.has_pending: + side_effect_cache.commits[ctx.module] = commit + return + gold_start = time.time() + try: + reservation = await deps.reserve_gold( + _entity(ctx).user_id, + ctx.module, + ctx.cost_gold, + ctx.session, + ) + await commit.reserve_gold( + reservation, + amount=ctx.cost_gold, + metadata={"module": ctx.module}, + ) + _recorder(ctx).set("reserve_gold", f"{time.time() - gold_start:.3f}s") + except deps.insufficient_gold_error: + deps.logger.debug( + f"预扣金币失败,金币不足: {ctx.module}", + deps.log_command, + session=ctx.session, + ) + raise SkipPluginException(f"{ctx.module} 金币不足,已取消执行...") from None + except TimeoutError: + deps.logger.error( + f"预扣金币超时,模块: {ctx.module}", + deps.log_command, + session=ctx.session, + ) + raise + side_effect_cache.commits[ctx.module] = commit + + +async def decision_log_stage( + ctx: AuthPipelineContext, + deps: AuthPipelineDependencies, +) -> None: + commit = ctx.side_effect_commit + has_deferred_commit = commit is not None and commit.has_pending + if ( + ctx.auth_result_cache is not None + and ctx.auth_allowed is not None + and not has_deferred_commit + ): + ctx.auth_result_cache[ctx.module] = ( + 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: + return AuthPipeline( + [ + AuthPipelineStage("route_gate", lambda ctx: route_gate_stage(ctx, deps)), + AuthPipelineStage( + "prepare_snapshot", + lambda ctx: prepare_snapshot_stage(ctx, deps), + ), + AuthPipelineStage( + "policy_precheck", + lambda ctx: policy_precheck_stage(ctx, deps), + ), + AuthPipelineStage( + "legacy_hook_adapter", + lambda ctx: legacy_hook_adapter_stage(ctx, deps), + ), + AuthPipelineStage( + "side_effect_commit", + lambda ctx: side_effect_commit_stage(ctx, deps), + ), + ] + ) diff --git a/zhenxun/builtin_plugins/hooks/auth_policy.py b/zhenxun/builtin_plugins/hooks/auth_policy.py new file mode 100644 index 00000000..541d8d45 --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth_policy.py @@ -0,0 +1,258 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any, Literal + +from zhenxun.services.cache.runtime_cache import _parse_block_modules +from zhenxun.utils.common_utils import CommonUtils +from zhenxun.utils.enum import BlockType, PluginType + +from .auth.exception import IsSuperuserException, SkipPluginException +from .auth_profile import PluginAuthProfile +from .auth_snapshot import AuthSnapshot + +PolicyEffect = Literal["allow", "deny", "skip", "defer"] + + +@dataclass(frozen=True, slots=True) +class PolicyDecision: + effect: PolicyEffect + reason: str = "" + metadata: dict[str, Any] = field(default_factory=dict) + + @property + def allowed(self) -> bool: + return self.effect == "allow" + + @property + def denied(self) -> bool: + return self.effect == "deny" + + @property + def skipped(self) -> bool: + return self.effect == "skip" + + @property + def deferred(self) -> bool: + return self.effect == "defer" + + +@dataclass(frozen=True, slots=True) +class PolicyPrincipal: + user_id: str + group_id: str | None = None + channel_id: str | None = None + is_superuser: bool = False + + +@dataclass(frozen=True, slots=True) +class PolicyAction: + name: str + module: str + + +@dataclass(frozen=True, slots=True) +class PolicyResource: + plugin: object + profile: PluginAuthProfile + + +@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 + + +class PolicyDecisionPoint: + """Structured permission decision helpers. + + This layer mirrors existing auth semantics and deliberately does not add a + new policy table. Side-effecting checks such as limit counters remain + deferred to the old hooks. + """ + + @staticmethod + def _missing(snapshot: AuthSnapshot, name: str) -> bool: + return name in snapshot.cache_misses + + @staticmethod + def _private_disabled(profile: PluginAuthProfile) -> bool: + return profile.block_type == BlockType.PRIVATE + + @staticmethod + def _group_disabled(profile: PluginAuthProfile) -> bool: + return profile.block_type == BlockType.GROUP + + @staticmethod + def _globally_disabled(profile: PluginAuthProfile) -> bool: + return profile.block_type == BlockType.ALL and not profile.status + + def decide( + self, + principal: PolicyPrincipal, + action: PolicyAction, + resource: PolicyResource, + context: PolicyContext, + ) -> PolicyDecision: + del action + snapshot = context.snapshot + 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: + return PolicyDecision("deny", "superuser_required") + return PolicyDecision("defer", "needs_legacy_hooks") + + def decide_bot(self, context: PolicyContext) -> PolicyDecision: + snapshot = context.snapshot + bot_data = snapshot.bot_data + if bot_data is None: + if self._missing(snapshot, "bot"): + return PolicyDecision("defer", "bot_cache_unavailable") + return PolicyDecision("deny", "bot_not_found") + if not bot_data.status and not context.allow_sleep_bypass: + return PolicyDecision("deny", "bot_sleeping") + module = snapshot.profile.module + if module and self._module_in_block_string(module, bot_data.block_plugins): + return PolicyDecision("deny", "bot_plugin_blocked") + return PolicyDecision("allow", "bot_allowed") + + def decide_group(self, context: PolicyContext) -> PolicyDecision: + snapshot = context.snapshot + if not snapshot.group_id: + return PolicyDecision("skip", "not_group_event") + group = snapshot.group + profile = snapshot.profile + if group is None: + if self._missing(snapshot, "group"): + return PolicyDecision("defer", "group_cache_unavailable") + return PolicyDecision("deny", "group_not_found") + if group.level < 0: + return PolicyDecision("deny", "group_blacklisted") + if ( + not group.status + and not context.allow_group_sleep_bypass + and not snapshot.is_superuser + ): + return PolicyDecision("deny", "group_sleeping") + if profile.level > group.level: + return PolicyDecision("deny", "group_level_low") + return PolicyDecision("allow", "group_allowed") + + def decide_admin(self, context: PolicyContext) -> PolicyDecision: + snapshot = context.snapshot + profile = snapshot.profile + if not profile.need_admin: + return PolicyDecision("skip", "admin_not_required") + if profile.plugin_type in {PluginType.SUPERUSER, PluginType.SUPER_AND_ADMIN}: + if snapshot.is_superuser: + return PolicyDecision("allow", "superuser") + if profile.plugin_type == PluginType.SUPERUSER: + return PolicyDecision("deny", "superuser_required") + if not profile.admin_level: + return PolicyDecision("skip", "admin_level_empty") + if snapshot.admin_levels is None: + return PolicyDecision("defer", "admin_levels_unavailable") + global_user, group_user = snapshot.admin_levels + user_level = global_user.user_level if global_user else 0 + if snapshot.group_id and group_user: + user_level = max(user_level, group_user.user_level) + if user_level < profile.admin_level: + return PolicyDecision("deny", "admin_level_low") + return PolicyDecision("allow", "admin_allowed") + + def decide_plugin(self, context: PolicyContext) -> PolicyDecision: + snapshot = context.snapshot + profile = snapshot.profile + group = snapshot.group + if snapshot.is_superuser: + return PolicyDecision("allow", "superuser") + if snapshot.group_id: + if group is None: + if self._missing(snapshot, "group"): + return PolicyDecision("defer", "group_cache_unavailable") + return PolicyDecision("deny", "group_not_found") + if profile.status and not self._group_disabled(profile): + block_set, super_block_set = self._group_block_sets(group) + if not block_set and not super_block_set: + return PolicyDecision("allow", "plugin_group_fast_allow") + block_set, super_block_set = self._group_block_sets(group) + if profile.module in super_block_set: + return PolicyDecision("deny", "plugin_superuser_blocked_in_group") + if profile.module in block_set: + return PolicyDecision("deny", "plugin_blocked_in_group") + if self._group_disabled(profile): + return PolicyDecision("deny", "plugin_disabled_in_group") + elif self._private_disabled(profile): + return PolicyDecision("deny", "plugin_disabled_in_private") + if self._globally_disabled(profile): + if group is not None and getattr(group, "is_super", False): + return PolicyDecision("allow", "super_group_bypass") + return PolicyDecision("deny", "plugin_global_disabled") + return PolicyDecision("allow", "plugin_allowed") + + @staticmethod + def _group_block_sets(group: object) -> tuple[frozenset[str], frozenset[str]]: + block_set = getattr(group, "block_plugin_set", None) + super_block_set = getattr(group, "superuser_block_plugin_set", None) + if block_set is None: + block_set = _parse_block_modules(getattr(group, "block_plugin", "") or "") + setattr(group, "block_plugin_set", block_set) + if super_block_set is None: + super_block_set = _parse_block_modules( + getattr(group, "superuser_block_plugin", "") or "" + ) + setattr(group, "superuser_block_plugin_set", super_block_set) + return block_set, super_block_set + + @staticmethod + def _module_in_block_string(module: str, value: str | None) -> bool: + if not value: + return False + return CommonUtils.format(module) in value or module in _parse_block_modules( + value + ) + + +def principal_from_snapshot(snapshot: AuthSnapshot) -> PolicyPrincipal: + return PolicyPrincipal( + user_id=snapshot.user_id, + group_id=snapshot.group_id, + channel_id=snapshot.channel_id, + is_superuser=snapshot.is_superuser, + ) + + +def action_from_snapshot(snapshot: AuthSnapshot) -> PolicyAction: + return PolicyAction(name="invoke_plugin", module=snapshot.module) + + +def resource_from_snapshot(snapshot: AuthSnapshot) -> PolicyResource: + return PolicyResource(plugin=snapshot.plugin, profile=snapshot.profile) + + +def raise_for_policy(decision: PolicyDecision, message: str | None = None) -> None: + if decision.denied: + raise SkipPluginException(message or decision.reason) + if decision.allowed and decision.reason == "super_group_bypass": + raise IsSuperuserException() + + +__all__ = [ + "PolicyAction", + "PolicyContext", + "PolicyDecision", + "PolicyDecisionPoint", + "PolicyPrincipal", + "PolicyResource", + "action_from_snapshot", + "principal_from_snapshot", + "raise_for_policy", + "resource_from_snapshot", +] diff --git a/zhenxun/builtin_plugins/hooks/auth_profile.py b/zhenxun/builtin_plugins/hooks/auth_profile.py new file mode 100644 index 00000000..d1c47c73 --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth_profile.py @@ -0,0 +1,122 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from zhenxun.services.cache.runtime_cache import ( + PluginLimitMemoryCache, + PluginLimitSnapshot, +) +from zhenxun.utils.enum import BlockType, PluginType + + +@dataclass(frozen=True, slots=True) +class PluginAuthProfile: + module: str + name: str + hidden: bool = False + status: bool = True + block_type: BlockType | None = None + plugin_type: PluginType | None = None + need_admin: bool = False + need_group_check: bool = False + has_limit: bool = False + cost_gold: int = 0 + admin_level: int = 0 + limit_superuser: bool = False + level: int = 0 + + @property + def superuser_only(self) -> bool: + return self.plugin_type == PluginType.SUPERUSER + + @property + def superuser_or_admin(self) -> bool: + return self.plugin_type == PluginType.SUPER_AND_ADMIN + + +def _plugin_admin_level(plugin) -> int: + try: + return int(getattr(plugin, "admin_level", 0) or 0) + except (TypeError, ValueError): + return 0 + + +def _plugin_cost_gold(plugin) -> int: + try: + return int(getattr(plugin, "cost_gold", 0) or 0) + except (TypeError, ValueError): + return 0 + + +def build_plugin_auth_profile(plugin, *, has_limit: bool = False) -> PluginAuthProfile: + plugin_type = getattr(plugin, "plugin_type", None) + admin_level = _plugin_admin_level(plugin) + block_type = getattr(plugin, "block_type", None) + module = str(getattr(plugin, "module", "") or "") + need_admin = bool(admin_level > 0) or plugin_type in { + PluginType.ADMIN, + PluginType.SUPERUSER, + PluginType.SUPER_AND_ADMIN, + } + return PluginAuthProfile( + module=module, + name=str(getattr(plugin, "name", "") or module), + hidden=plugin_type == PluginType.HIDDEN, + status=bool(getattr(plugin, "status", True)), + block_type=block_type, + plugin_type=plugin_type, + need_admin=need_admin, + need_group_check=block_type + in {BlockType.ALL, BlockType.GROUP, BlockType.PRIVATE}, + has_limit=bool(has_limit), + cost_gold=_plugin_cost_gold(plugin), + admin_level=admin_level, + limit_superuser=bool(getattr(plugin, "limit_superuser", False)), + level=int(getattr(plugin, "level", 0) or 0), + ) + + +async def get_plugin_auth_profile( + plugin, + *, + event_cache: dict | None = None, + allow_cache_load: bool = True, +) -> PluginAuthProfile: + module = str(getattr(plugin, "module", "") or "") + profile_cache: dict[str, PluginAuthProfile] = {} + if event_cache is not None: + profile_cache = event_cache.setdefault("plugin_auth_profiles", {}) + cached = profile_cache.get(module) + if cached is not None: + return cached + + limits: list[PluginLimitSnapshot] | None = None + limits_ready = False + if event_cache is not None: + limit_cache = event_cache.setdefault("module_limit_entries", {}) + if module in limit_cache: + limits = limit_cache[module] + limits_ready = True + if limits is None: + limits = PluginLimitMemoryCache.get_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_ready = True + if limits is None: + limits = [] + profile = build_plugin_auth_profile(plugin, has_limit=bool(limits)) + if event_cache is not None: + profile_cache[module] = profile + event_cache.setdefault("module_limits", {})[module] = profile.has_limit + event_cache.setdefault("module_limits_ready", {})[module] = limits_ready + if limits_ready: + event_cache.setdefault("module_limit_entries", {})[module] = limits + return profile + + +__all__ = [ + "PluginAuthProfile", + "build_plugin_auth_profile", + "get_plugin_auth_profile", +] diff --git a/zhenxun/builtin_plugins/hooks/auth_route.py b/zhenxun/builtin_plugins/hooks/auth_route.py new file mode 100644 index 00000000..88b9b43d --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth_route.py @@ -0,0 +1,60 @@ +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"] diff --git a/zhenxun/builtin_plugins/hooks/auth_runtime_config.py b/zhenxun/builtin_plugins/hooks/auth_runtime_config.py new file mode 100644 index 00000000..6cbb210e --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth_runtime_config.py @@ -0,0 +1,135 @@ +from __future__ import annotations + +from dataclasses import dataclass, fields +import os + + +@dataclass(frozen=True, slots=True) +class AuthDispatchRuntimeConfig: + hooks_concurrency_limit: int = 5 + db_concurrency_limit: int = 6 + command_exact_limit: int = 96 + command_shortcut_limit: int = 32 + command_regex_limit: int = 8 + system_limit: int = 64 + passive_light_limit: int = 12 + passive_db_limit: int = 4 + passive_http_limit: int = 4 + passive_ai_limit: int = 2 + passive_render_limit: int = 2 + overload_selected_threshold: int = 48 + overload_lane_wait_ms: float = 200.0 + timeout_seconds: float = 5.0 + circuit_reset_time: int = 300 + matcher_route_prefilter_ttl: int = 2 + prefilter_stats_log_interval: float = 10.0 + cache_sweep_interval: float = 1.0 + dispatch_stats_log_interval: float = 10.0 + + +@dataclass(frozen=True, slots=True) +class AuthObservabilityRuntimeConfig: + buffer_max_retain: int = 20_000 + flush_trigger_size: int = 256 + flush_batch_size: int = 500 + flush_interval_seconds: float = 30.0 + drop_log_interval_seconds: float = 10.0 + allow_sample_rate: float = 0.005 + overloaded_allow_sample_rate: float = 0.02 + non_allow_sample_rate: float = 1.0 + backpressure_sample_rate: float = 0.2 + backpressure_severe_active_threshold: int = 5 + + +_WARNED_ENV_KEYS: set[str] = set() +_ENV_ALIASES: dict[str, tuple[str, ...]] = { + "hooks_concurrency_limit": ("ZX_AUTH_HOOKS_CONCURRENCY_LIMIT",), + "db_concurrency_limit": ("ZX_AUTH_DB_CONCURRENCY_LIMIT",), + "command_exact_limit": ("ZX_AUTH_DISPATCH_COMMAND_EXACT_LIMIT",), + "command_shortcut_limit": ("ZX_AUTH_DISPATCH_COMMAND_SHORTCUT_LIMIT",), + "command_regex_limit": ("ZX_AUTH_DISPATCH_COMMAND_REGEX_LIMIT",), + "system_limit": ("ZX_AUTH_DISPATCH_SYSTEM_LIMIT",), + "passive_light_limit": ("ZX_AUTH_DISPATCH_PASSIVE_LIGHT_LIMIT",), + "passive_db_limit": ("ZX_AUTH_DISPATCH_PASSIVE_DB_LIMIT",), + "passive_http_limit": ("ZX_AUTH_DISPATCH_PASSIVE_HTTP_LIMIT",), + "passive_ai_limit": ("ZX_AUTH_DISPATCH_PASSIVE_AI_LIMIT",), + "passive_render_limit": ("ZX_AUTH_DISPATCH_PASSIVE_RENDER_LIMIT",), + "overload_selected_threshold": ("ZX_AUTH_OVERLOAD_SELECTED_THRESHOLD",), + "overload_lane_wait_ms": ("ZX_AUTH_OVERLOAD_LANE_WAIT_MS",), + "timeout_seconds": ("ZX_AUTH_TIMEOUT_SECONDS",), + "circuit_reset_time": ("ZX_AUTH_CIRCUIT_RESET_TIME",), + "matcher_route_prefilter_ttl": ("ZX_AUTH_MATCHER_ROUTE_PREFILTER_TTL",), + "prefilter_stats_log_interval": ("ZX_AUTH_PREFILTER_STATS_LOG_INTERVAL",), + "cache_sweep_interval": ("ZX_AUTH_CACHE_SWEEP_INTERVAL",), + "dispatch_stats_log_interval": ("ZX_AUTH_DISPATCH_STATS_LOG_INTERVAL",), +} + + +def _env_name(prefix: str, field_name: str) -> str: + return f"{prefix}_{field_name.upper()}" + + +def _env_names(prefix: str, field_name: str) -> tuple[str, ...]: + generated = _env_name(prefix, field_name) + aliases = _ENV_ALIASES.get(field_name, ()) + return (*aliases, generated) + + +def _coerce_env_value(raw: str, default: object) -> object: + if isinstance(default, bool): + return raw.strip().lower() in {"1", "true", "yes", "on"} + if isinstance(default, int) and not isinstance(default, bool): + return int(raw) + if isinstance(default, float): + return float(raw) + return raw + + +def _warn_invalid_env(env_name: str, raw: str, exc: Exception) -> None: + if env_name in _WARNED_ENV_KEYS: + return + _WARNED_ENV_KEYS.add(env_name) + try: + from zhenxun.services.log import logger + + logger.warning( + f"{env_name}={raw!r} 解析失败,使用默认值: {exc}", + "AuthRuntimeConfig", + ) + except Exception: + # Config is imported early on the auth hot path; logging must be optional. + return + + +def _load_config(cls: type, prefix: str): + values = {} + default_obj = cls() + for item in fields(default_obj): + default = getattr(default_obj, item.name) + env_name = "" + raw = None + for candidate in _env_names(prefix, item.name): + candidate_value = os.getenv(candidate) + if candidate_value is not None and candidate_value.strip(): + env_name = candidate + raw = candidate_value + break + if raw is None or not raw.strip(): + values[item.name] = default + continue + try: + values[item.name] = _coerce_env_value(raw, default) + except Exception as exc: + _warn_invalid_env(env_name, raw, exc) + values[item.name] = default + return cls(**values) + + +AUTH_DISPATCH_RUNTIME_CONFIG = _load_config( + AuthDispatchRuntimeConfig, + "ZX_AUTH", +) +AUTH_OBSERVABILITY_RUNTIME_CONFIG = _load_config( + AuthObservabilityRuntimeConfig, + "ZX_AUTH_OBSERVABILITY", +) diff --git a/zhenxun/builtin_plugins/hooks/auth_side_effect.py b/zhenxun/builtin_plugins/hooks/auth_side_effect.py new file mode 100644 index 00000000..bd051bb7 --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth_side_effect.py @@ -0,0 +1,236 @@ +from __future__ import annotations + +import asyncio +from collections.abc import Awaitable, Callable, Sequence +from dataclasses import dataclass, field +import time +from typing import Any, Protocol + +from nonebot_plugin_uninfo import Uninfo + +from zhenxun.services.log import logger +from zhenxun.utils.utils import EntityIDs + +from .auth.config import LOGGER_COMMAND +from .auth.utils import send_message + +AsyncAction = Callable[[], Awaitable[None]] + + +class SyncReservation(Protocol): + def commit(self) -> None: ... + + def release(self) -> None: ... + + +class AsyncReservation(Protocol): + async def commit(self) -> None: ... + + async def release(self) -> None: ... + + +ReservationLike = AsyncAction | SyncReservation | AsyncReservation +SideEffectKind = str +SideEffectState = str + + +@dataclass(slots=True) +class SideEffectReservation: + kind: SideEffectKind + reservation: ReservationLike + amount: int = 0 + metadata: dict[str, Any] = field(default_factory=dict) + state: SideEffectState = "reserved" + reserved_at: float = field(default_factory=time.monotonic) + committed_at: float | None = None + released_at: float | None = None + reason: str | None = None + + @property + def should_auto_unblock(self) -> bool: + return bool(getattr(self.reservation, "should_auto_unblock", False)) + + +async def _maybe_await(value: Any) -> None: + if hasattr(value, "__await__"): + await value + + +async def _commit_reservation(reservation: ReservationLike) -> None: + commit = getattr(reservation, "commit", None) + if callable(commit): + await _maybe_await(commit()) + return + if callable(reservation): + await reservation() + + +async def _release_reservation(reservation: ReservationLike) -> None: + release = getattr(reservation, "release", None) + if callable(release): + await _maybe_await(release()) + + +@dataclass(slots=True) +class SideEffectCommit: + """权限链副作用提交器。 + + 第一阶段只封装既有调用点,不改变扣金币、限流提交、权限提示发送时机。 + """ + + session: Uninfo + module: str + owner_matcher_id: int | None = None + limit_entity: EntityIDs | None = None + _reservations: dict[SideEffectKind, SideEffectReservation] = field( + default_factory=dict + ) + committed: bool = False + + @property + def limit_should_auto_unblock(self) -> bool: + record = self._reservations.get("limit") + return bool(record and record.should_auto_unblock) + + @property + def has_pending(self) -> bool: + return any(record.state == "reserved" for record in self._reservations.values()) + + @property + def pending_kinds(self) -> tuple[str, ...]: + return tuple( + kind + for kind, record in self._reservations.items() + if record.state == "reserved" + ) + + def snapshot(self) -> dict[str, Any]: + return { + "module": self.module, + "committed": self.committed, + "pending": list(self.pending_kinds), + "reservations": { + kind: { + "state": record.state, + "amount": record.amount, + "metadata": record.metadata, + "reason": record.reason, + } + for kind, record in self._reservations.items() + }, + } + + async def send_permission_tip( + self, + message: list | str, + check_tag: str | None = None, + *, + background: bool = False, + timeout: float | None = None, + ) -> None: + try: + tip_coro = send_message( + self.session, + message, + check_tag, + background=background, + ) + if timeout and not background: + await asyncio.wait_for(tip_coro, timeout=timeout) + else: + await tip_coro + except asyncio.TimeoutError: + logger.error("发送权限提示超时", LOGGER_COMMAND, session=self.session) + + async def reduce_gold( + self, + func: ReservationLike, + ) -> None: + await self.reserve_gold(func) + await self.commit_gold() + + async def reserve( + self, + kind: SideEffectKind, + reservation: ReservationLike, + *, + amount: int = 0, + metadata: dict[str, Any] | None = None, + ) -> None: + await self.release(kind, f"replace_{kind}_reservation") + self._reservations[kind] = SideEffectReservation( + kind=kind, + reservation=reservation, + amount=amount, + metadata=metadata or {}, + ) + + async def commit(self, kind: SideEffectKind) -> None: + record = self._reservations.get(kind) + if record is None or record.state != "reserved": + return + try: + await _commit_reservation(record.reservation) + except Exception: + record.reason = "commit_failed" + raise + record.state = "committed" + record.committed_at = time.monotonic() + + async def release( + self, + kind: SideEffectKind, + reason: str | None = None, + ) -> None: + record = self._reservations.get(kind) + if record is None or record.state != "reserved": + return + try: + await _release_reservation(record.reservation) + finally: + record.state = "released" + record.released_at = time.monotonic() + record.reason = reason + + async def reserve_limit(self, reservation: ReservationLike) -> None: + await self.reserve("limit", reservation) + + async def commit_limit( + self, + reservation: ReservationLike | None = None, + ) -> None: + if reservation is not None: + await self.reserve_limit(reservation) + await self.commit("limit") + + async def release_limit(self, reason: str | None = None) -> None: + await self.release("limit", reason) + + async def reserve_gold( + self, + reservation: ReservationLike, + *, + amount: int = 0, + metadata: dict[str, Any] | None = None, + ) -> None: + await self.reserve( + "gold", + reservation, + amount=amount, + metadata=metadata, + ) + + async def commit_gold(self) -> None: + await self.commit("gold") + + async def rollback_gold(self, reason: str | None = None) -> None: + await self.release("gold", reason) + + async def rollback_all(self, reason: str | None = None) -> None: + for kind in list(self._reservations): + await self.release(kind, reason) + + async def commit_all(self, *, order: Sequence[str] = ("gold", "limit")) -> None: + for name in order: + await self.commit(name) + self.committed = True diff --git a/zhenxun/builtin_plugins/hooks/auth_snapshot.py b/zhenxun/builtin_plugins/hooks/auth_snapshot.py new file mode 100644 index 00000000..4c79486a --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth_snapshot.py @@ -0,0 +1,199 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import TYPE_CHECKING + +from zhenxun.services.cache.runtime_cache import ( + BanMemoryCache, + BotMemoryCache, + BotSnapshot, + GroupMemoryCache, + GroupSnapshot, + LevelUserMemoryCache, + LevelUserSnapshot, +) + +from .auth.context import EventContext +from .auth_profile import PluginAuthProfile + +if TYPE_CHECKING: + from nonebot.adapters import Bot + + +@dataclass(slots=True) +class AuthSnapshot: + context: EventContext + plugin: object + profile: PluginAuthProfile + bot_data: BotSnapshot | None = None + group: GroupSnapshot | None = None + admin_levels: tuple[LevelUserSnapshot | None, LevelUserSnapshot | None] | None = ( + None + ) + ban_state: bool | None = None + user_balance_loaded: bool = False + user_balance: int | None = None + cache_misses: frozenset[str] = field(default_factory=frozenset) + + @property + def module(self) -> str: + return self.profile.module + + @property + def is_superuser(self) -> bool: + return self.context.is_superuser + + @property + def user_id(self) -> str: + return self.context.user_id + + @property + def group_id(self) -> str | None: + return self.context.group_id + + @property + def channel_id(self) -> str | None: + return self.context.channel_id + + @property + def has_ban_cache(self) -> bool: + return self.ban_state is not None + + @property + def cache_ready(self) -> bool: + return not self.cache_misses + + +async def build_auth_snapshot( + *, + context: EventContext, + plugin: object, + profile: PluginAuthProfile, + bot: "Bot", + skip_ban: bool = False, + allow_cache_load: bool = False, +) -> AuthSnapshot: + event_cache = context.event_cache + entity = context.entity + cache_misses: set[str] = set() + + bot_data: BotSnapshot | None = None + if ( + event_cache is not None + and "bot_data" in event_cache + and (event_cache.get("bot_cache_ready") or not allow_cache_load) + ): + bot_data = event_cache.get("bot_data") + else: + bot_data = BotMemoryCache.get_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(): + cache_misses.add("bot") + if event_cache is not None: + event_cache["bot_data"] = bot_data + event_cache["bot_cache_ready"] = BotMemoryCache.is_loaded() + + group = None + if entity.group_id: + if ( + event_cache is not None + and "group" in event_cache + and (event_cache.get("group_cache_ready") or not allow_cache_load) + ): + 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(): + cache_misses.add("group") + elif group is None and allow_cache_load: + group = await GroupMemoryCache.get(entity.group_id, entity.channel_id) + if event_cache is not None: + event_cache["group"] = group + event_cache["group_cache_ready"] = GroupMemoryCache.is_loaded() + + admin_levels = None + if profile.need_admin: + if ( + event_cache is not None + and "admin_levels" in event_cache + and (event_cache.get("admin_cache_ready") or not allow_cache_load) + ): + admin_levels = event_cache.get("admin_levels") + else: + admin_levels = LevelUserMemoryCache.get_levels_if_ready( + entity.user_id, + entity.group_id, + ) + if admin_levels is None: + if allow_cache_load: + admin_levels = await LevelUserMemoryCache.get_levels( + entity.user_id, + entity.group_id, + ) + else: + 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() + + 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) + 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) + if event_cache is not None: + event_cache["ban_state"] = ban_state + else: + cache_misses.add("ban") + + return AuthSnapshot( + context=context, + plugin=plugin, + profile=profile, + bot_data=bot_data, + group=group, + admin_levels=admin_levels, + ban_state=ban_state, + cache_misses=frozenset(cache_misses), + ) + + +async def get_or_build_auth_snapshot( + *, + context: EventContext, + plugin: object, + profile: PluginAuthProfile, + bot: "Bot", + skip_ban: bool = False, + allow_cache_load: bool = False, +) -> AuthSnapshot: + event_cache = context.event_cache + module = profile.module + if event_cache is not None: + snapshot_cache = event_cache.setdefault("auth_snapshots", {}) + cached = snapshot_cache.get(module) + if isinstance(cached, AuthSnapshot): + if not (allow_cache_load and cached.cache_misses): + return cached + snapshot = await build_auth_snapshot( + context=context, + plugin=plugin, + profile=profile, + bot=bot, + skip_ban=skip_ban, + allow_cache_load=allow_cache_load, + ) + if event_cache is not None: + event_cache.setdefault("auth_snapshots", {})[module] = snapshot + return snapshot + + +__all__ = ["AuthSnapshot", "build_auth_snapshot", "get_or_build_auth_snapshot"] diff --git a/zhenxun/builtin_plugins/hooks/auth_trace.py b/zhenxun/builtin_plugins/hooks/auth_trace.py new file mode 100644 index 00000000..0d7cf36b --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth_trace.py @@ -0,0 +1,37 @@ +from __future__ import annotations + +import time + +from .auth.config import WARNING_THRESHOLD + + +class HookTraceRecorder: + def __init__(self, start_time: float) -> None: + self._start_time = start_time + self._enabled = False + self._data: dict[str, str] = {} + + def _ensure_enabled(self) -> bool: + if self._enabled: + return True + if time.time() - self._start_time <= WARNING_THRESHOLD: + return False + self._enabled = True + return True + + def set(self, key: str, value: str) -> None: + if self._ensure_enabled(): + self._data[key] = value + + def setdefault(self, key: str, value: str) -> None: + if self._ensure_enabled(): + self._data.setdefault(key, value) + + def contains(self, key: str) -> bool: + return key in self._data + + def snapshot(self) -> dict[str, str]: + return self._data if self._enabled else {} + + +__all__ = ["HookTraceRecorder"] diff --git a/zhenxun/builtin_plugins/hooks/auth_types.py b/zhenxun/builtin_plugins/hooks/auth_types.py new file mode 100644 index 00000000..1e21a8bc --- /dev/null +++ b/zhenxun/builtin_plugins/hooks/auth_types.py @@ -0,0 +1,62 @@ +from __future__ import annotations + +from dataclasses import dataclass, field + +from zhenxun.models.plugin_info import PluginInfo +from zhenxun.models.user_console import UserConsole + +from .auth.context import PermissionContext +from .auth_policy import PolicyContext +from .auth_profile import PluginAuthProfile +from .auth_snapshot import AuthSnapshot + + +@dataclass(slots=True) +class AuthPreparation: + plugin: PluginInfo + user: UserConsole | None + profile: PluginAuthProfile + snapshot: AuthSnapshot + permission_context: PermissionContext + policy_context: PolicyContext + + +@dataclass(slots=True) +class AuthPolicyFlags: + should_return_allowed: bool = False + + +@dataclass(slots=True) +class AuthLaneContext: + lane: str = "passive_light" + scope_key: str = "" + queue_size: int = 0 + + @property + def is_guaranteed(self) -> bool: + return self.lane.startswith("command_") or self.lane == "system" + + +@dataclass(slots=True) +class EventDispatchContext: + event_type: str + plain_text: str = "" + raw_text: str = "" + trie_command_text: str = "" + trie_raw_command: str = "" + text_candidates: tuple[str, ...] = () + to_me: bool = False + has_url: bool = False + has_image: bool = False + is_command_like: bool = False + route_modules: set[str] = field(default_factory=set) + ai_route_modules: set[str] = field(default_factory=set) + ai_route_heads: set[str] = field(default_factory=set) + + +__all__ = [ + "AuthLaneContext", + "AuthPolicyFlags", + "AuthPreparation", + "EventDispatchContext", +] diff --git a/zhenxun/builtin_plugins/hooks/call_hook.py b/zhenxun/builtin_plugins/hooks/call_hook.py index 923480be..ac40fa37 100644 --- a/zhenxun/builtin_plugins/hooks/call_hook.py +++ b/zhenxun/builtin_plugins/hooks/call_hook.py @@ -65,7 +65,7 @@ async def handle_api_result( if not Config.get_config("hook", "RECORD_BOT_SENT_MESSAGES"): return try: - await BotMessageStore.create( + await BotMessageStore.append_buffered( bot_id=bot.self_id, user_id=user_id, group_id=group_id, diff --git a/zhenxun/builtin_plugins/platform/qq/group_handle/data_source.py b/zhenxun/builtin_plugins/platform/qq/group_handle/data_source.py index 33742258..751dfb33 100644 --- a/zhenxun/builtin_plugins/platform/qq/group_handle/data_source.py +++ b/zhenxun/builtin_plugins/platform/qq/group_handle/data_source.py @@ -18,6 +18,10 @@ from zhenxun.models.group_console import GroupConsole from zhenxun.models.group_member_info import GroupInfoUser from zhenxun.models.level_user import LevelUser from zhenxun.models.plugin_info import PluginInfo +from zhenxun.services.hot_query_cache import ( + invalidate_group_members, + invalidate_member_names, +) from zhenxun.services.log import logger from zhenxun.utils.common_utils import CommonUtils from zhenxun.utils.enum import RequestHandleType @@ -82,6 +86,8 @@ async def _refresh_member_info_async( "platform": platform, }, ) + await invalidate_group_members(group_id, [user_id]) + await invalidate_member_names([user_id]) class GroupManager: @@ -333,6 +339,8 @@ class GroupManager: "platform": platform, }, ) + await invalidate_group_members(group_id, [user_id]) + await invalidate_member_names([user_id]) task = asyncio.create_task( _refresh_member_info_async( bot, str(group_id), str(user_id), _normalize_platform(platform) @@ -400,6 +408,8 @@ class GroupManager: user_name = f"{user_id}" if user: await user.delete() + await invalidate_group_members(group_id, [user_id]) + await invalidate_member_names([user_id]) logger.info( f"名称: {user_name} 退出群聊", "group_decrease_handle", diff --git a/zhenxun/builtin_plugins/platform/qq_api/ug_watch.py b/zhenxun/builtin_plugins/platform/qq_api/ug_watch.py index 4435e880..a27e9fb5 100644 --- a/zhenxun/builtin_plugins/platform/qq_api/ug_watch.py +++ b/zhenxun/builtin_plugins/platform/qq_api/ug_watch.py @@ -4,6 +4,10 @@ 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 @@ -27,6 +31,8 @@ async def _(session: Uninfo): 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) diff --git a/zhenxun/builtin_plugins/shop/_data_source.py b/zhenxun/builtin_plugins/shop/_data_source.py index f8076c3f..c4bdc7a3 100644 --- a/zhenxun/builtin_plugins/shop/_data_source.py +++ b/zhenxun/builtin_plugins/shop/_data_source.py @@ -17,11 +17,11 @@ from zhenxun import ui from zhenxun.configs.config import BotConfig from zhenxun.models.friend_user import FriendUser from zhenxun.models.goods_info import GoodsInfo -from zhenxun.models.group_member_info import GroupInfoUser from zhenxun.models.user_console import UserConsole from zhenxun.models.user_props_log import UserPropsLog from zhenxun.services import avatar_service from zhenxun.services.buffered_writers import append_user_gold_log +from zhenxun.services.hot_query_cache import get_group_user_ids, get_member_names from zhenxun.services.log import logger from zhenxun.ui.models import ImageCell, TextCell from zhenxun.utils.enum import GoldHandle, PropHandle @@ -96,9 +96,7 @@ class ShopParam(BaseModel): async def gold_rank(session: Uninfo, group_id: str | None, num: int) -> bytes | str: query = UserConsole if group_id: - uid_list = await GroupInfoUser.filter(group_id=group_id).values_list( - "user_id", flat=True - ) + uid_list = await get_group_user_ids(group_id) if uid_list: query = query.filter(user_id__in=uid_list) user_list = await query.annotate().order_by("-gold").values_list("user_id", "gold") @@ -115,11 +113,7 @@ async def gold_rank(session: Uninfo, group_id: str | None, num: int) -> bytes | ) uid2name = {user[0]: user[1] for user in friend_user} if diff_id := set(user_id_list).difference(set(uid2name.keys())): - group_user = await GroupInfoUser.filter(user_id__in=diff_id).values_list( - "user_id", "user_name" - ) - for g in group_user: - uid2name[g[0]] = g[1] + uid2name.update(await get_member_names(diff_id)) column_name = ["排名", "-", "名称", "金币", "平台"] data_list = [] platform = PlatformUtils.get_platform(session) diff --git a/zhenxun/builtin_plugins/sign_in/_data_source.py b/zhenxun/builtin_plugins/sign_in/_data_source.py index 4dbe3ff0..5593383e 100644 --- a/zhenxun/builtin_plugins/sign_in/_data_source.py +++ b/zhenxun/builtin_plugins/sign_in/_data_source.py @@ -10,12 +10,12 @@ from zhenxun import ui from zhenxun.configs.path_config import IMAGE_PATH from zhenxun.models.friend_user import FriendUser from zhenxun.models.goods_info import GoodsInfo -from zhenxun.models.group_member_info import GroupInfoUser from zhenxun.models.sign_log import SignLog from zhenxun.models.sign_user import SignUser from zhenxun.models.user_console import UserConsole from zhenxun.services.avatar_service import avatar_service from zhenxun.services.buffered_writers import append_user_gold_log +from zhenxun.services.hot_query_cache import get_group_user_ids, get_member_names from zhenxun.services.log import logger from zhenxun.ui.models import ImageCell, TextCell from zhenxun.utils.enum import GoldHandle @@ -52,9 +52,7 @@ class SignManage: """ query = SignUser if group_id: - user_list = await GroupInfoUser.filter(group_id=group_id).values_list( - "user_id", flat=True - ) + user_list = await get_group_user_ids(group_id) if user_list: query = query.filter(user_id__in=user_list) user_list = ( @@ -76,11 +74,7 @@ class SignManage: ) uid2name = {f[0]: f[1] for f in friend_list} if diff_id := set(user_id_list).difference(set(uid2name.keys())): - group_user = await GroupInfoUser.filter(user_id__in=diff_id).values_list( - "user_id", "user_name" - ) - for g in group_user: - uid2name[g[0]] = g[1] + uid2name.update(await get_member_names(diff_id)) data_list = [] platform = PlatformUtils.get_platform(session) for i, user in enumerate(user_list): diff --git a/zhenxun/builtin_plugins/sign_in/utils.py b/zhenxun/builtin_plugins/sign_in/utils.py index 4b611b7e..f01589f0 100644 --- a/zhenxun/builtin_plugins/sign_in/utils.py +++ b/zhenxun/builtin_plugins/sign_in/utils.py @@ -327,7 +327,7 @@ async def _generate_html_card( "pages/builtin/sign", data=card_data, clip_selector=".wrapper", - clip_padding=8, + clip_padding=0, disable_animations=True, screenshot_scale="css", ) diff --git a/zhenxun/builtin_plugins/statistics/_data_source.py b/zhenxun/builtin_plugins/statistics/_data_source.py index 9e5dc34f..069bf335 100644 --- a/zhenxun/builtin_plugins/statistics/_data_source.py +++ b/zhenxun/builtin_plugins/statistics/_data_source.py @@ -1,15 +1,61 @@ -from tortoise.functions import Count +from dataclasses import dataclass +from datetime import datetime, timedelta from zhenxun.models.group_console import GroupConsole -from zhenxun.models.group_member_info import GroupInfoUser from zhenxun.models.plugin_info import PluginInfo -from zhenxun.models.statistics import Statistics +from zhenxun.services.hot_query_cache import ( + get_member_name, + get_statistics_plugin_counts_cached, +) from zhenxun.utils.echart_utils import ChartUtils from zhenxun.utils.echart_utils.models import Barh from zhenxun.utils.enum import PluginType from zhenxun.utils.time_utils import TimeUtils +@dataclass(frozen=True) +class _StatisticsPeriod: + title: str + start_time: datetime | None + + +def _get_statistics_period(search_type: str | None) -> _StatisticsPeriod: + if search_type == "day": + return _StatisticsPeriod("日(1天)", TimeUtils.get_day_start()) + if search_type == "week": + return _StatisticsPeriod( + "周(7天)", + TimeUtils.get_day_start( + datetime.now(TimeUtils.DEFAULT_TIMEZONE) - timedelta(days=6) + ), + ) + if search_type == "month": + return _StatisticsPeriod( + "月(30天)", + TimeUtils.get_day_start( + datetime.now(TimeUtils.DEFAULT_TIMEZONE) - timedelta(days=29) + ), + ) + return _StatisticsPeriod("", None) + + +def _build_statistics_title( + *, + target_name: str | None, + is_global: bool, + period_title: str, +) -> str: + title = f"{period_title}功能调用统计" if period_title else "功能调用统计" + prefixes: list[str] = [] + if target_name: + prefixes.append(target_name) + if is_global: + prefixes.append("全局") + if prefixes: + return f"{' '.join(prefixes)} {title}" + return title + + class StatisticsManage: @classmethod async def get_statistics( @@ -20,55 +66,49 @@ class StatisticsManage: user_id: str | None = None, group_id: str | None = None, ): - day = None - day_type = "" - if search_type == "day": - day = 1 - day_type = "日" - elif search_type == "month": - day = 30 - day_type = "月" - elif search_type == "week": - day = 7 - day_type = "周" - if day_type: - day_type += f"({day}天)" - title = "" + period = _get_statistics_period(search_type) if user_id: """查用户""" - query = GroupInfoUser.filter(user_id=user_id) - if group_id: - query = query.filter(group_id=group_id) - user = await query.first() - title = f"{user.user_name if user else user_id} {day_type}功能调用统计" + user_name = await get_member_name(user_id, group_id) + title = _build_statistics_title( + target_name=user_name or user_id, + is_global=is_global and not group_id, + period_title=period.title, + ) elif group_id: """查群组""" group = await GroupConsole.get_group(group_id=group_id) - title = f"{group.group_name if group else group_id} {day_type}功能调用统计" + title = _build_statistics_title( + target_name=group.group_name if group else group_id, + is_global=False, + period_title=period.title, + ) else: - title = "功能调用统计" + title = _build_statistics_title( + target_name=None, + is_global=is_global, + period_title=period.title, + ) if is_global and not user_id: - title = f"全局 {title}" - return await cls.get_global_statistics(plugin_name, day, title) + return await cls.get_global_statistics( + plugin_name, period.start_time, title + ) if user_id: - return await cls.get_my_statistics(user_id, group_id, day, title) + return await cls.get_my_statistics( + user_id, group_id, period.start_time, title + ) if group_id: - return await cls.get_group_statistics(group_id, day, title) + return await cls.get_group_statistics(group_id, period.start_time, title) return None @classmethod async def get_global_statistics( - cls, plugin_name: str | None, day: int | None, title: str + cls, plugin_name: str | None, start_time: datetime | None, title: str ) -> bytes | str: - query = Statistics - if plugin_name: - query = query.filter(plugin_name=plugin_name) - if day: - query = query.filter(create_time__gte=TimeUtils.get_day_start()) - data_list = ( - await query.annotate(count=Count("id")) - .group_by("plugin_name") - .values_list("plugin_name", "count") + data_list = await get_statistics_plugin_counts_cached( + "global", + plugin_name=plugin_name, + start_time=start_time, ) return ( await cls.__build_image(data_list, title) @@ -78,17 +118,18 @@ class StatisticsManage: @classmethod async def get_my_statistics( - cls, user_id: str, group_id: str | None, day: int | None, title: str + cls, + user_id: str, + group_id: str | None, + start_time: datetime | None, + title: str, ): - query = Statistics.filter(user_id=user_id) - if group_id: - query = query.filter(group_id=group_id) - if day: - query = query.filter(create_time__gte=TimeUtils.get_day_start()) - data_list = ( - await query.annotate(count=Count("id")) - .group_by("plugin_name") - .values_list("plugin_name", "count") + data_list = await get_statistics_plugin_counts_cached( + "user", + plugin_name=None, + start_time=start_time, + user_id=user_id, + group_id=group_id, ) return ( await cls.__build_image(data_list, title) @@ -97,14 +138,14 @@ class StatisticsManage: ) @classmethod - async def get_group_statistics(cls, group_id: str, day: int | None, title: str): - query = Statistics.filter(group_id=group_id) - if day: - query = query.filter(create_time__gte=TimeUtils.get_day_start()) - data_list = ( - await query.annotate(count=Count("id")) - .group_by("plugin_name") - .values_list("plugin_name", "count") + async def get_group_statistics( + cls, group_id: str, start_time: datetime | None, title: str + ): + data_list = await get_statistics_plugin_counts_cached( + "group", + plugin_name=None, + start_time=start_time, + group_id=group_id, ) return ( await cls.__build_image(data_list, title) diff --git a/zhenxun/cli.py b/zhenxun/cli.py index 5a908d89..5526d426 100644 --- a/zhenxun/cli.py +++ b/zhenxun/cli.py @@ -23,6 +23,121 @@ RESTART_POLL_INTERVAL = 0.5 WORKER_SOFT_EXIT_TIMEOUT = 15.0 WORKER_TERMINATE_TIMEOUT = 5.0 WORKER_KILL_TIMEOUT = 5.0 +ENV_EXAMPLE_FILE = ".env.example" +ENV_DEV_FILE = ".env.dev" + + +def _env_assignment_key(line: str, *, include_commented: bool = False) -> str | None: + stripped = line.strip() + if include_commented and stripped.startswith("#"): + stripped = stripped[1:].lstrip() + if not stripped or stripped.startswith("#") or "=" not in stripped: + return None + key = stripped.split("=", 1)[0].strip() + return key if key.replace("_", "").isalnum() else None + + +def _env_key(line: str) -> str | None: + return _env_assignment_key(line) + + +def _env_block_key(block: list[str]) -> str | None: + for line in block: + if key := _env_key(line): + return key + return None + + +def _env_block_anchor_key(block: list[str]) -> str | None: + for line in block: + if key := _env_assignment_key(line, include_commented=True): + return key + return None + + +def _split_env_blocks(lines: list[str]) -> list[tuple[int, list[str]]]: + blocks: list[tuple[int, list[str]]] = [] + current: list[str] = [] + start_index = 0 + for index, line in enumerate(lines): + if line.strip(): + if not current: + start_index = index + current.append(line) + elif current: + blocks.append((start_index, current)) + current = [] + + if current: + blocks.append((start_index, current)) + return blocks + + +def _find_env_block_start(lines: list[str], key: str) -> int | None: + for start_index, block in _split_env_blocks(lines): + if _env_block_anchor_key(block) == key: + return start_index + return None + + +def _insert_env_block_before( + lines: list[str], + index: int, + block: list[str], +) -> list[str]: + insert_block = block.copy() + if index > 0 and lines[index - 1].strip(): + insert_block.insert(0, "\n") + if index < len(lines) and insert_block and insert_block[-1].strip(): + insert_block.append("\n") + return lines[:index] + insert_block + lines[index:] + + +def _sync_env_missing_items(project_root: Path) -> None: + """Copy missing .env keys from .env.example without touching existing values.""" + example_path = project_root / ENV_EXAMPLE_FILE + env_path = project_root / ENV_DEV_FILE + if not example_path.exists(): + return + if not env_path.exists(): + env_path.write_text(example_path.read_text(encoding="utf-8"), encoding="utf-8") + _launcher_log("已根据 .env.example 生成 .env.dev") + return + + example_lines = example_path.read_text(encoding="utf-8").splitlines(keepends=True) + env_lines = env_path.read_text(encoding="utf-8").splitlines(keepends=True) + example_blocks = _split_env_blocks(example_lines) + existing_keys = {key for line in env_lines if (key := _env_key(line))} + missing_blocks: list[tuple[int, list[str]]] = [] + + for block_index, (_, block) in enumerate(example_blocks): + key = _env_block_key(block) + if key and key not in existing_keys: + missing_blocks.append((block_index, block)) + + if not missing_blocks: + return + + updated_lines = env_lines + added_keys: list[str] = [] + for block_index, block in missing_blocks: + key = _env_block_key(block) + if not key: + continue + anchor_index = len(updated_lines) + for _, next_block in example_blocks[block_index + 1 :]: + next_key = _env_block_anchor_key(next_block) + if not next_key: + continue + if (found := _find_env_block_start(updated_lines, next_key)) is not None: + anchor_index = found + break + updated_lines = _insert_env_block_before(updated_lines, anchor_index, block) + existing_keys.add(key) + added_keys.append(key) + + env_path.write_text("".join(updated_lines), encoding="utf-8") + _launcher_log(f"已补齐 .env.dev 缺失配置: {', '.join(added_keys)}") def _launcher_log(message: str) -> None: @@ -53,7 +168,8 @@ def _ensure_project_root() -> Path: def _run_worker() -> None: """启动 Bot worker(必须在项目目录下执行)""" - _ensure_project_root() + project_root = _ensure_project_root() + _sync_env_missing_items(project_root) import contextlib import platform @@ -91,18 +207,32 @@ def _run_worker() -> None: f"使用 {htmlrender_browser_channel} 作为 htmlrender 驱动启动..." ) + nonebot.init(htmlrender_browser_channel=htmlrender_browser_channel) + from nonebot.adapters.onebot.v11 import Adapter as OneBotV11Adapter - nonebot.init(htmlrender_browser_channel=htmlrender_browser_channel) + from zhenxun.configs.config import BotConfig driver = nonebot.get_driver() driver.register_adapter(OneBotV11Adapter) + enabled_adapters = ["OneBot V11"] + + if BotConfig.qq_adapter_load: + try: + from nonebot.adapters.qq import Adapter as QQAdapter # type: ignore + except ImportError as e: + raise RuntimeError( + "QQ_ADAPTER_LOAD=True 但未安装 nonebot-adapter-qq," + "请安装后再开启 QQ 官方适配器。" + ) from e + driver.register_adapter(QQAdapter) + enabled_adapters.append("QQ") + + nonebot.logger.info(f"已启用适配器: {', '.join(enabled_adapters)}") nonebot.load_plugins("zhenxun/builtin_plugins") nonebot.load_plugins("zhenxun/plugins") - from zhenxun.configs.config import BotConfig - for ext in BotConfig.ext_path: ext = ext.strip() if ext: diff --git a/zhenxun/configs/config.py b/zhenxun/configs/config.py index 0ac30f44..b1e0d473 100644 --- a/zhenxun/configs/config.py +++ b/zhenxun/configs/config.py @@ -21,6 +21,8 @@ class BotSetting(BaseModel): """官bot id:账号id""" ext_path: list[str] = Field(default_factory=list) """第三方插件路径""" + qq_adapter_load: bool = False + """是否加载 QQ 官方适配器""" def get_qbot_uid(self, qbot_id: str) -> str | None: """获取官bot账号id diff --git a/zhenxun/models/_bot_message_buffer.py b/zhenxun/models/_bot_message_buffer.py new file mode 100644 index 00000000..d2e30129 --- /dev/null +++ b/zhenxun/models/_bot_message_buffer.py @@ -0,0 +1,113 @@ +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("关闭") diff --git a/zhenxun/models/auth_decision_log.py b/zhenxun/models/auth_decision_log.py new file mode 100644 index 00000000..33578a9e --- /dev/null +++ b/zhenxun/models/auth_decision_log.py @@ -0,0 +1,49 @@ +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 [] diff --git a/zhenxun/models/ban_console.py b/zhenxun/models/ban_console.py index f20d1bc2..2080486e 100644 --- a/zhenxun/models/ban_console.py +++ b/zhenxun/models/ban_console.py @@ -229,8 +229,4 @@ class BanConsole(Model): @classmethod async def _run_script(cls): - return [ - "CREATE INDEX idx_ban_console_user_id ON ban_console(user_id);", - "CREATE INDEX idx_ban_console_group_id ON ban_console(group_id);", - "ALTER TABLE ban_console ADD COLUMN ban_reason TEXT DEFAULT NULL;", - ] + return [] diff --git a/zhenxun/models/bot_console.py b/zhenxun/models/bot_console.py index 8d21a763..176e8e29 100644 --- a/zhenxun/models/bot_console.py +++ b/zhenxun/models/bot_console.py @@ -2,9 +2,9 @@ from typing import Literal, overload from tortoise import fields -from zhenxun.configs.config import BotConfig from zhenxun.services.cache.runtime_cache import BotMemoryCache from zhenxun.services.db_context import Model +from zhenxun.services.db_context.schema_ops import AlterColumnType, RenameColumn from zhenxun.utils.enum import CacheType @@ -470,32 +470,9 @@ class BotConsole(Model): @classmethod async def _run_script(cls): - db_type = (BotConfig.get_sql_type() or "").lower() - - scripts = [ - "ALTER TABLE bot_console RENAME COLUMN block_plugin TO block_plugins;", - "ALTER TABLE bot_console RENAME COLUMN block_task TO block_tasks;", - "ALTER TABLE bot_console ADD available_plugins text default '';", - "ALTER TABLE bot_console ADD available_tasks text default '';", + return [ + RenameColumn("bot_console", "block_plugin", "block_plugins"), + RenameColumn("bot_console", "block_task", "block_tasks"), + AlterColumnType("bot_console", "block_plugins", "TEXT"), + AlterColumnType("bot_console", "block_tasks", "TEXT"), ] - - if "postgres" in db_type: - scripts.extend( - [ - "ALTER TABLE bot_console ALTER COLUMN block_plugins TYPE TEXT;", - "ALTER TABLE bot_console ALTER COLUMN block_tasks TYPE TEXT;", - "ALTER TABLE bot_console ALTER COLUMN available_plugins TYPE TEXT;", - "ALTER TABLE bot_console ALTER COLUMN available_tasks TYPE TEXT;", - ] - ) - elif "mysql" in db_type: - scripts.extend( - [ - "ALTER TABLE bot_console MODIFY COLUMN block_plugins TEXT;", - "ALTER TABLE bot_console MODIFY COLUMN block_tasks TEXT;", - "ALTER TABLE bot_console MODIFY COLUMN available_plugins TEXT;", - "ALTER TABLE bot_console MODIFY COLUMN available_tasks TEXT;", - ] - ) - - return scripts diff --git a/zhenxun/models/bot_message_store.py b/zhenxun/models/bot_message_store.py index fa1244f9..6159e094 100644 --- a/zhenxun/models/bot_message_store.py +++ b/zhenxun/models/bot_message_store.py @@ -3,6 +3,8 @@ 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) @@ -27,3 +29,27 @@ class BotMessageStore(Model): 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, + ) + ) diff --git a/zhenxun/models/chat_history.py b/zhenxun/models/chat_history.py index c9f0a91e..12ed6ad4 100644 --- a/zhenxun/models/chat_history.py +++ b/zhenxun/models/chat_history.py @@ -1,11 +1,12 @@ from datetime import datetime, timedelta -from typing import ClassVar, Literal +from typing import Any, ClassVar, Literal from typing_extensions import Self from tortoise import fields -from tortoise.functions import Count +from tortoise.expressions import Q from zhenxun.services.db_context import Model +from zhenxun.services.db_context.schema_ops import AlterColumnType, RenameColumn class ChatHistory(Model): @@ -35,6 +36,41 @@ class ChatHistory(Model): ("user_id", "group_id"), ] + @classmethod + def _platform_from_scope(cls, platform_scope: str | None) -> str | None: + """Map the new fine-grained scope back to the legacy platform column.""" + if not platform_scope: + return None + scope = str(platform_scope).lower() + if scope in {"qq", "qq_client", "qq_api"} or scope.startswith("qq_"): + return "qq" + if "onebot" in scope: + return "qq" + return scope + + @classmethod + def scoped_query(cls, platform_scope: str | None = None, **filters: Any): + """Return a chat-history query compatible with platform_scope callers. + + chat_history currently stores the coarse legacy ``platform`` column rather + than a dedicated ``platform_scope`` column, so this method intentionally + stays as a thin compatibility shim. + """ + query = cls.filter(**filters) + if not platform_scope: + return query + if "platform" in filters or any(k.startswith("platform__") for k in filters): + return query + + platform = cls._platform_from_scope(platform_scope) + if not platform: + return query + if platform == "qq" and str(platform_scope).lower() in {"qq", "qq_client"}: + return query.filter( + Q(platform=platform) | Q(platform__isnull=True) | Q(platform="") + ) + return query.filter(platform=platform) + @classmethod async def get_group_msg_rank( cls, @@ -42,7 +78,7 @@ class ChatHistory(Model): limit: int = 10, order: str = "DESC", date_scope: tuple[datetime, datetime] | None = None, - ) -> list[Self]: + ) -> list[tuple[str, int]]: """获取排行数据 参数: @@ -51,18 +87,9 @@ class ChatHistory(Model): order: 排序类型,desc,des date_scope: 日期范围 """ - o = "-" if order == "DESC" else "" - query = cls.filter(group_id=gid) if gid else cls - if date_scope: - filter_scope = (date_scope[0].isoformat(" "), date_scope[1].isoformat(" ")) - query = query.filter(create_time__range=filter_scope) - return list( - await query.annotate(count=Count("user_id")) - .order_by(f"{o}count") - .group_by("user_id") - .limit(limit) - .values_list("user_id", "count") - ) # type: ignore + from zhenxun.services.hot_query_cache import get_chat_history_rank_cached + + return await get_chat_history_rank_cached(cls, gid, limit, order, date_scope) @classmethod async def get_group_first_msg_datetime( @@ -73,22 +100,21 @@ class ChatHistory(Model): 参数: group_id: 群组id """ - if group_id: - message = ( - await cls.filter(group_id=group_id).order_by("create_time").first() - ) - else: - message = await cls.all().order_by("create_time").first() - return message.create_time if message else None + from zhenxun.services.hot_query_cache import ( + get_chat_history_first_msg_datetime_cached, + ) + + return await get_chat_history_first_msg_datetime_cached(cls, group_id) @classmethod async def get_message( cls, - uid: str, - gid: str, + uid: str | None, + gid: str | None, type_: Literal["user", "group"], msg_type: Literal["private", "group"] | None = None, days: int | tuple[datetime, datetime] | None = None, + platform_scope: str | None = None, ) -> list[Self]: """获取消息查询query @@ -98,15 +124,16 @@ class ChatHistory(Model): type_: 类型,私聊或群聊 msg_type: 消息类型,用户或群聊 days: 限制日期 + platform_scope: 兼容细粒度平台作用域 """ if type_ == "user": - query = cls.filter(user_id=uid) + query = cls.scoped_query(platform_scope=platform_scope, user_id=uid) if msg_type == "private": query = query.filter(group_id__isnull=True) elif msg_type == "group": query = query.filter(group_id__not_isnull=True) else: - query = cls.filter(group_id=gid) + query = cls.scoped_query(platform_scope=platform_scope, group_id=gid) if uid: query = query.filter(user_id=uid) if days: @@ -128,13 +155,15 @@ class ChatHistory(Model): # 允许 plain_text 为空 "alter table chat_history alter plain_text drop not null;", # 将user_id改为user_id - "ALTER TABLE chat_history RENAME COLUMN user_qq TO user_id;", - "ALTER TABLE chat_history " - "ALTER COLUMN user_id TYPE character varying(255);", - "ALTER TABLE chat_history " - "ALTER COLUMN group_id TYPE character varying(255);", - # 添加bot_id字段 - "ALTER TABLE chat_history ADD bot_id VARCHAR(255);", - "ALTER TABLE chat_history ALTER COLUMN bot_id TYPE character varying(255);", - "ALTER TABLE chat_history ADD COLUMN platform character varying(255);", + RenameColumn("chat_history", "user_qq", "user_id"), + AlterColumnType( + "chat_history", + "user_id", + {"postgres": "character varying(255)", "mysql": "VARCHAR(255)"}, + ), + AlterColumnType( + "chat_history", + "group_id", + {"postgres": "character varying(255)", "mysql": "VARCHAR(255)"}, + ), ] diff --git a/zhenxun/models/fg_request.py b/zhenxun/models/fg_request.py index 95677f2a..2342abbb 100644 --- a/zhenxun/models/fg_request.py +++ b/zhenxun/models/fg_request.py @@ -8,7 +8,6 @@ from zhenxun.configs.config import BotConfig from zhenxun.models.group_console import GroupConsole from zhenxun.services.db_context import Model from zhenxun.services.log import logger -from zhenxun.utils.common_utils import SqlUtils from zhenxun.utils.enum import RequestHandleType, RequestType from zhenxun.utils.exception import NotFoundError from zhenxun.utils.manager.bot_profile_manager import BotProfileManager @@ -173,6 +172,4 @@ class FgRequest(Model): @classmethod async def _run_script(cls): - return [ - SqlUtils.add_column("fg_request", "message_ids", "character varying(255)") - ] + return [] diff --git a/zhenxun/models/friend_user.py b/zhenxun/models/friend_user.py index 7b235466..7283e3ca 100644 --- a/zhenxun/models/friend_user.py +++ b/zhenxun/models/friend_user.py @@ -1,6 +1,7 @@ from tortoise import fields from zhenxun.services.db_context import Model +from zhenxun.services.db_context.schema_ops import DropColumn class FriendUser(Model): @@ -30,4 +31,4 @@ class FriendUser(Model): @classmethod def _run_script(cls): - return ["ALTER TABLE friend_users DROP COLUMN nickname;"] + return [DropColumn("friend_users", "nickname")] diff --git a/zhenxun/models/goods_info.py b/zhenxun/models/goods_info.py index 07efa6f4..0c764a33 100644 --- a/zhenxun/models/goods_info.py +++ b/zhenxun/models/goods_info.py @@ -4,6 +4,7 @@ import uuid from tortoise import fields from zhenxun.services.db_context import Model +from zhenxun.services.db_context.schema_ops import DropColumn class GoodsInfo(Model): @@ -158,11 +159,6 @@ class GoodsInfo(Model): @classmethod async def _run_script(cls): return [ - "ALTER TABLE goods_info ADD uuid VARCHAR(255);", - "ALTER TABLE goods_info ADD daily_limit Integer DEFAULT 0;", - "ALTER TABLE goods_info ADD is_passive boolean DEFAULT False;", - "ALTER TABLE goods_info ADD icon VARCHAR(255);", # 删除 daily_purchase_limit 字段 - "ALTER TABLE goods_info DROP daily_purchase_limit;", - "ALTER TABLE goods_info ADD partition VARCHAR(255);", + DropColumn("goods_info", "daily_purchase_limit"), ] diff --git a/zhenxun/models/group_console.py b/zhenxun/models/group_console.py index c5c5644d..c52650fb 100644 --- a/zhenxun/models/group_console.py +++ b/zhenxun/models/group_console.py @@ -4,13 +4,13 @@ from typing_extensions import Self from tortoise import fields from tortoise.backends.base.client import BaseDBAsyncClient -from zhenxun.configs.config import BotConfig from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.task_info import TaskInfo from zhenxun.services.cache import CacheRoot from zhenxun.services.cache.runtime_cache import GroupMemoryCache from zhenxun.services.data_access import DataAccess from zhenxun.services.db_context import Model +from zhenxun.services.db_context.schema_ops import AlterColumnType, CreateIndex from zhenxun.utils.enum import CacheType, DbLockType, PluginType if TYPE_CHECKING: @@ -516,49 +516,13 @@ class GroupConsole(Model): @classmethod def _run_script(cls): - db_type = (BotConfig.get_sql_type() or "").lower() - - scripts = [ - "ALTER TABLE group_console ADD superuser_block_plugin" - " Text NOT NULL DEFAULT '';", - "ALTER TABLE group_console ADD superuser_block_task" - " Text NOT NULL DEFAULT '';", - "CREATE INDEX idx_group_console_group_id ON group_console(group_id);", - ( - "CREATE INDEX idx_group_console_group_null_channel ON " - "group_console(group_id) WHERE channel_id IS NULL;" + return [ + CreateIndex( + "group_console", + ("group_id",), + name="idx_group_console_group_null_channel", + where="channel_id IS NULL", ), + AlterColumnType("group_console", "block_plugin", "TEXT"), + AlterColumnType("group_console", "block_task", "TEXT"), ] - - if "postgres" in db_type: - scripts.extend( - [ - ("ALTER TABLE group_console ALTER COLUMN block_plugin TYPE TEXT;"), - ( - "ALTER TABLE group_console ALTER COLUMN " - "superuser_block_plugin TYPE TEXT;" - ), - ("ALTER TABLE group_console ALTER COLUMN block_task TYPE TEXT;"), - ( - "ALTER TABLE group_console ALTER COLUMN " - "superuser_block_task TYPE TEXT;" - ), - ] - ) - elif "mysql" in db_type: - scripts.extend( - [ - ("ALTER TABLE group_console MODIFY COLUMN block_plugin TEXT;"), - ( - "ALTER TABLE group_console MODIFY COLUMN " - "superuser_block_plugin TEXT;" - ), - ("ALTER TABLE group_console MODIFY COLUMN block_task TEXT;"), - ( - "ALTER TABLE group_console MODIFY COLUMN " - "superuser_block_task TEXT;" - ), - ] - ) - - return scripts diff --git a/zhenxun/models/group_member_info.py b/zhenxun/models/group_member_info.py index ee5f3584..87d441ac 100644 --- a/zhenxun/models/group_member_info.py +++ b/zhenxun/models/group_member_info.py @@ -2,8 +2,8 @@ from typing import ClassVar from tortoise import fields -from zhenxun.configs.config import BotConfig from zhenxun.services.db_context import Model +from zhenxun.services.db_context.schema_ops import AlterColumnType, DropColumn class GroupInfoUser(Model): @@ -35,9 +35,9 @@ class GroupInfoUser(Model): 参数: group_id: 群号 """ - return set( - await cls.filter(group_id=group_id).values_list("user_id", flat=True) - ) # type: ignore + from zhenxun.services.hot_query_cache import get_group_user_ids + + return await get_group_user_ids(group_id) @classmethod async def get_user_all_group(cls, user_id: str) -> list[str]: @@ -46,60 +46,24 @@ class GroupInfoUser(Model): 参数: user_id: 用户id """ - return list( - await cls.filter(user_id=user_id).values_list("group_id", flat=True) - ) # type: ignore + from zhenxun.services.hot_query_cache import get_user_group_ids + + return await get_user_group_ids(user_id) @classmethod async def _run_script(cls): - db_type = (BotConfig.get_sql_type() or "").lower() - scripts = ["ALTER TABLE group_info_users DROP COLUMN nickname;"] - - if "postgres" in db_type: - scripts.extend( - [ - ( - "ALTER TABLE group_info_users ADD COLUMN IF NOT EXISTS " - "platform character varying(255);" - ), - ( - "ALTER TABLE group_info_users ALTER COLUMN user_id " - "TYPE character varying(255) USING user_id::character varying;" - ), - ( - "ALTER TABLE group_info_users ALTER COLUMN group_id " - "TYPE character varying(255) USING group_id::character varying;" - ), - ( - "ALTER TABLE group_info_users ALTER COLUMN platform " - "TYPE character varying(255) USING platform::character varying;" - ), - ] - ) - elif "mysql" in db_type: - scripts.extend( - [ - ( - "ALTER TABLE group_info_users ADD COLUMN " - "platform VARCHAR(255) NULL;" - ), - ( - "ALTER TABLE group_info_users MODIFY COLUMN " - "user_id VARCHAR(255) NOT NULL;" - ), - ( - "ALTER TABLE group_info_users MODIFY COLUMN " - "group_id VARCHAR(255) NOT NULL;" - ), - ( - "ALTER TABLE group_info_users MODIFY COLUMN " - "platform VARCHAR(255) NULL;" - ), - ] - ) - elif "sqlite" in db_type: - scripts.append( - "ALTER TABLE group_info_users ADD COLUMN platform VARCHAR(255);" - ) - - return scripts + return [ + DropColumn("group_info_users", "nickname"), + AlterColumnType( + "group_info_users", + "user_id", + {"postgres": "character varying(255)", "mysql": "VARCHAR(255)"}, + nullable=False, + ), + AlterColumnType( + "group_info_users", + "group_id", + {"postgres": "character varying(255)", "mysql": "VARCHAR(255)"}, + nullable=False, + ), + ] diff --git a/zhenxun/models/level_user.py b/zhenxun/models/level_user.py index 3d0ddbb2..13afbe93 100644 --- a/zhenxun/models/level_user.py +++ b/zhenxun/models/level_user.py @@ -2,6 +2,7 @@ from tortoise import fields from zhenxun.services.cache.runtime_cache import LevelUserMemoryCache from zhenxun.services.db_context import Model +from zhenxun.services.db_context.schema_ops import AlterColumnType, RenameColumn from zhenxun.utils.enum import CacheType @@ -151,9 +152,16 @@ class LevelUser(Model): async def _run_script(cls): return [ # 将user_id改为user_id - "ALTER TABLE level_users RENAME COLUMN user_qq TO user_id;", - "ALTER TABLE level_users ALTER COLUMN user_id TYPE character varying(255);", + RenameColumn("level_users", "user_qq", "user_id"), + AlterColumnType( + "level_users", + "user_id", + {"postgres": "character varying(255)", "mysql": "VARCHAR(255)"}, + ), # 将user_id字段类型改为character varying(255) - "ALTER TABLE level_users " - "ALTER COLUMN group_id TYPE character varying(255);", + AlterColumnType( + "level_users", + "group_id", + {"postgres": "character varying(255)", "mysql": "VARCHAR(255)"}, + ), ] diff --git a/zhenxun/models/plugin_info.py b/zhenxun/models/plugin_info.py index 53a7136d..79f2be0f 100644 --- a/zhenxun/models/plugin_info.py +++ b/zhenxun/models/plugin_info.py @@ -206,12 +206,4 @@ class PluginInfo(Model): @classmethod async def _run_script(cls): - return [ - "ALTER TABLE plugin_info ADD COLUMN parent character varying(255);", - "ALTER TABLE plugin_info ADD COLUMN is_show boolean DEFAULT true;", - "ALTER TABLE plugin_info ADD COLUMN ignore_prompt boolean DEFAULT false;", - "ALTER TABLE plugin_info ADD COLUMN impression float DEFAULT 0;", - "CREATE INDEX idx_plugin_info_module ON plugin_info(module);", - "ALTER TABLE plugin_info ADD COLUMN ignore_statistics" - " boolean DEFAULT false;", - ] + return [] diff --git a/zhenxun/models/runtime_backpressure_log.py b/zhenxun/models/runtime_backpressure_log.py new file mode 100644 index 00000000..3b8f050f --- /dev/null +++ b/zhenxun/models/runtime_backpressure_log.py @@ -0,0 +1,39 @@ +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 [] diff --git a/zhenxun/models/statistics.py b/zhenxun/models/statistics.py index a3172654..a4a9fc4e 100644 --- a/zhenxun/models/statistics.py +++ b/zhenxun/models/statistics.py @@ -3,6 +3,7 @@ from typing import ClassVar from tortoise import fields from zhenxun.services.db_context import Model +from zhenxun.services.db_context.schema_ops import AlterColumnType, RenameColumn class Statistics(Model): @@ -32,9 +33,16 @@ class Statistics(Model): @classmethod async def _run_script(cls): return [ - "ALTER TABLE statistics RENAME COLUMN user_qq TO user_id;", + RenameColumn("statistics", "user_qq", "user_id"), # 将user_qq改为user_id - "ALTER TABLE statistics ALTER COLUMN user_id TYPE character varying(255);", - "ALTER TABLE statistics ALTER COLUMN group_id TYPE character varying(255);", - "ALTER TABLE statistics ADD bot_id Text DEFAULT '';", + AlterColumnType( + "statistics", + "user_id", + {"postgres": "character varying(255)", "mysql": "VARCHAR(255)"}, + ), + AlterColumnType( + "statistics", + "group_id", + {"postgres": "character varying(255)", "mysql": "VARCHAR(255)"}, + ), ] diff --git a/zhenxun/models/task_info.py b/zhenxun/models/task_info.py index 94e8ead2..ccbdba51 100644 --- a/zhenxun/models/task_info.py +++ b/zhenxun/models/task_info.py @@ -101,8 +101,4 @@ class TaskInfo(Model): @classmethod async def _run_script(cls): - return [ - "ALTER TABLE task_info ADD default_status boolean DEFAULT true;", - "ALTER TABLE task_info ADD load_status boolean DEFAULT false;", - # 默认状态 - ] + return [] diff --git a/zhenxun/models/user_console.py b/zhenxun/models/user_console.py index 50f49fe2..bcc8a6a1 100644 --- a/zhenxun/models/user_console.py +++ b/zhenxun/models/user_console.py @@ -1,8 +1,11 @@ import asyncio +from dataclasses import dataclass from typing import ClassVar from tortoise import fields from tortoise.exceptions import IntegrityError +from tortoise.expressions import F +from tortoise.transactions import in_transaction from zhenxun.models.goods_info import GoodsInfo from zhenxun.services.buffered_writers import append_user_gold_log @@ -11,6 +14,41 @@ from zhenxun.utils.enum import CacheType, GoldHandle from zhenxun.utils.exception import GoodsNotFound, InsufficientGold +@dataclass(slots=True) +class GoldReservation: + user_id: str + gold: int + handle: GoldHandle + plugin_module: str + platform: str | None = None + committed: bool = False + released: bool = False + + async def commit(self) -> None: + if self.committed or self.released: + return + await append_user_gold_log( + user_id=self.user_id, + gold=self.gold, + handle=self.handle, + source=self.plugin_module, + ) + self.committed = True + + async def release(self) -> None: + if self.released or self.committed: + return + self.released = True + async with in_transaction() as connection: + updated = ( + await UserConsole.filter(user_id=self.user_id) + .using_db(connection) + .update(gold=F("gold") + self.gold) + ) + if updated: + await UserConsole.invalidate_user_cache(self.user_id) + + class UserConsole(Model): id = fields.IntField(pk=True, generated=True, auto_increment=True) """自增id""" @@ -150,6 +188,51 @@ class UserConsole(Model): user_id=user_id, gold=gold, handle=handle, source=plugin_module ) + @classmethod + async def reserve_gold( + cls, + user_id: str, + gold: int, + handle: GoldHandle, + plugin_module: str, + platform: str | None = None, + ) -> GoldReservation: + """预扣金币;插件最终未执行时可 release 补偿。""" + async with in_transaction() as connection: + user = await cls.filter(user_id=user_id).using_db(connection).get_or_none() + if user is None: + try: + user = await cls.create( + using_db=connection, + user_id=user_id, + platform=platform, + uid=await cls.get_new_uid(), + ) + except IntegrityError: + user = await cls.filter(user_id=user_id).using_db(connection).get() + if user.gold < gold: + raise InsufficientGold() + updated = ( + await cls.filter(user_id=user_id, gold__gte=gold) + .using_db(connection) + .update(gold=F("gold") - gold) + ) + if not updated: + raise InsufficientGold() + return GoldReservation( + user_id=user_id, + gold=gold, + handle=handle, + plugin_module=plugin_module, + platform=platform, + ) + + @classmethod + async def invalidate_user_cache(cls, user_id: str) -> None: + from zhenxun.services.cache import CacheRoot + + await CacheRoot.invalidate_cache(CacheType.USERS, user_id) + @classmethod async def add_props( cls, user_id: str, goods_uuid: str, num: int = 1, platform: str | None = None @@ -223,7 +306,4 @@ class UserConsole(Model): @classmethod async def _run_script(cls): - return [ - "CREATE INDEX idx_user_console_user_id ON user_console(user_id);", - "CREATE INDEX idx_user_console_uid ON user_console(uid);", - ] + return [] diff --git a/zhenxun/services/auth_observability.py b/zhenxun/services/auth_observability.py new file mode 100644 index 00000000..e4b76fff --- /dev/null +++ b/zhenxun/services/auth_observability.py @@ -0,0 +1,547 @@ +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 "") + 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 "") + buckets.setdefault(lane, []).append(row) + advice: dict[str, dict[str, Any]] = {} + for lane, items in buckets.items(): + if lane == "": + 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 ""), []).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 ''}:{row.get('reason') or ''}" + 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() diff --git a/zhenxun/services/cache/runtime_cache.py b/zhenxun/services/cache/runtime_cache.py index e779a47c..4eb0cfd7 100644 --- a/zhenxun/services/cache/runtime_cache.py +++ b/zhenxun/services/cache/runtime_cache.py @@ -715,12 +715,22 @@ class PluginInfoMemoryCache: return await cls.refresh() + @classmethod + def is_loaded(cls) -> bool: + return cls._loaded + @classmethod async def get_by_module(cls, module: str) -> "PluginInfo | None": if not cls._loaded: await cls.ensure_loaded() return cls._to_model(cls._by_module.get(module)) + @classmethod + def get_by_module_if_ready(cls, module: str) -> "PluginInfo | None": + if not cls._loaded: + return None + return cls._to_model(cls._by_module.get(module)) + @classmethod async def get_all(cls) -> dict[str, "PluginInfo"]: if not cls._loaded: @@ -851,6 +861,10 @@ class BotMemoryCache: return await cls.refresh() + @classmethod + def is_loaded(cls) -> bool: + return cls._loaded + @classmethod async def get(cls, bot_id: str | None) -> BotSnapshot | None: bot_id = cls._normalize(bot_id) @@ -866,6 +880,19 @@ class BotMemoryCache: cls._mark_negative(bot_id) return None + @classmethod + def get_if_ready(cls, bot_id: str | None) -> BotSnapshot | None: + bot_id = cls._normalize(bot_id) + if not bot_id or not cls._loaded: + return None + entry = cls._by_id.get(bot_id) + if entry: + return entry + if cls._is_negative(bot_id): + return None + cls._mark_negative(bot_id) + return None + @classmethod async def get_all(cls) -> dict[str, BotSnapshot]: if not cls._loaded: @@ -1199,6 +1226,10 @@ class LevelUserMemoryCache: return await cls.refresh() + @classmethod + def is_loaded(cls) -> bool: + return cls._loaded + @classmethod async def ensure_fresh(cls) -> None: interval = LEVEL_MEM_REFRESH_INTERVAL @@ -1244,6 +1275,23 @@ class LevelUserMemoryCache: group_user = cls._by_key.get(group_key) return global_user, group_user + @classmethod + def get_levels_if_ready( + cls, user_id: str | None, group_id: str | None + ) -> tuple[LevelUserSnapshot | None, LevelUserSnapshot | None] | None: + if not cls._loaded: + return None + global_user = None + group_user = None + global_key = cls._key(user_id, "") + if global_key: + global_user = cls._by_key.get(global_key) + if group_id: + group_key = cls._key(user_id, group_id) + if group_key: + group_user = cls._by_key.get(group_key) + return global_user, group_user + @classmethod async def get_max_level(cls, user_id: str | None) -> int: user_id = cls._normalize(user_id) @@ -1576,6 +1624,10 @@ class PluginLimitMemoryCache: return await cls.refresh() + @classmethod + def is_loaded(cls) -> bool: + return cls._loaded + @classmethod async def get_limits(cls, module: str) -> list[PluginLimitSnapshot]: normalized = cls._normalize(module) @@ -1592,6 +1644,22 @@ class PluginLimitMemoryCache: cls._mark_negative(module) return [] + @classmethod + def get_limits_if_ready(cls, module: str) -> list[PluginLimitSnapshot] | None: + normalized = cls._normalize(module) + if not normalized: + return [] + module = normalized + if not cls._loaded: + return None + limits = cls._by_module.get(module) + if limits is not None: + return limits + if cls._is_negative(module): + return [] + cls._mark_negative(module) + return [] + @classmethod def get_all_limits(cls) -> list[PluginLimitSnapshot]: return list(cls._by_id.values()) diff --git a/zhenxun/services/db_context/__init__.py b/zhenxun/services/db_context/__init__.py index e2f7f86c..b144dae0 100644 --- a/zhenxun/services/db_context/__init__.py +++ b/zhenxun/services/db_context/__init__.py @@ -1,8 +1,10 @@ import asyncio import hashlib import json +import os from pathlib import Path import re +from typing import Literal from urllib.parse import urlparse import aiofiles @@ -27,8 +29,12 @@ from .config import ( prompt, ) from .exceptions import DbConnectError, DbUrlIsNode +from .schema_guard import repair_safe_schema_drift +from .schema_ops import SchemaOpRisk, normalize_schema_ops from .utils import with_db_timeout +Dialect = Literal["sqlite", "postgres", "mysql", "unknown"] + MODELS = db_model.models SCRIPT_METHOD = db_model.script_method @@ -47,7 +53,66 @@ __all__ = [ driver = nonebot.get_driver() -_SCRIPT_HASH_FILE = Path() / "data" / ".db_script_hash" +_SCRIPT_HASH_DIR = Path() / "data" / ".db_script_hashes" +_TRUE_VALUES = {"1", "true", "yes", "on"} + + +def _connection_dialect() -> Dialect: + try: + connection = Tortoise.get_connection("default") + capabilities = getattr(connection, "capabilities", None) + raw = str(getattr(capabilities, "dialect", "") or "").lower() + if raw.startswith("sqlite"): + return "sqlite" + if raw.startswith("postgres"): + return "postgres" + if raw.startswith("mysql"): + return "mysql" + except Exception: + pass + return "unknown" + + +def _allow_guarded_schema_ops() -> bool: + """Whether startup may run guarded SchemaOp migrations. + + Safe SchemaOps are limited to non-destructive changes such as adding nullable + columns and non-unique indexes. Guarded operations may rename, drop, or alter + columns, so keep them opt-in to avoid damaging existing databases during + normal startup. + """ + return os.getenv("DB_SCHEMA_RUN_GUARDED_OPS", "").strip().lower() in _TRUE_VALUES + + +def _extract_alter_table_name(sql: str) -> str | None: + match = re.match(r"ALTER\s+TABLE\s+[`\"]?(\w+)[`\"]?", sql, re.IGNORECASE) + return match.group(1) if match else None + + +def _extract_create_index_table_name(sql: str) -> str | None: + match = re.search(r"\bON\s+[`\"]?(\w+)[`\"]?\s*\(", sql, re.IGNORECASE) + return match.group(1) if match else None + + +def _db_script_hash_file(script_fingerprint: str) -> Path: + parsed = urlparse(BotConfig.db_url or "") + dialect = parsed.scheme or "unknown" + if dialect == "sqlite": + db_identity = str(Path(parsed.path).resolve()) + else: + db_identity = f"{parsed.hostname or ''}:{parsed.port or ''}{parsed.path}" + db_hash = hashlib.md5( + json.dumps( + { + "dialect": dialect, + "db": db_identity, + "script": script_fingerprint, + }, + ensure_ascii=False, + sort_keys=True, + ).encode() + ).hexdigest() + return _SCRIPT_HASH_DIR / f"{db_hash}.json" def get_config() -> dict: @@ -131,20 +196,39 @@ async def init(): f"{len(db_model.script_method)} 个..." ) sql_list = [] + allow_guarded_ops = _allow_guarded_schema_ops() for module, func in db_model.script_method: try: - sql = await func() if is_coroutine_callable(func) else func() - if sql: - sql_list += sql + items = await func() if is_coroutine_callable(func) else func() + if not items: + continue + for item in items: + if not isinstance(item, str): + if item.risk == SchemaOpRisk.MANUAL: + logger.debug(f"{module} 跳过手动迁移动作: {item}") + continue + if ( + item.risk == SchemaOpRisk.GUARDED + and not allow_guarded_ops + ): + logger.debug(f"{module} 跳过受保护迁移动作: {item}") + continue + if item.risk != SchemaOpRisk.SAFE and not allow_guarded_ops: + logger.debug(f"{module} 跳过未知风险迁移动作: {item}") + continue + sql_list += normalize_schema_ops([item], _connection_dialect()) except Exception as e: logger.debug(f"{module} 执行SCRIPT_METHOD方法出错...", e=e) if sql_list: fingerprint = hashlib.md5( json.dumps(sorted(sql_list), ensure_ascii=False).encode() ).hexdigest() + script_hash_file = _db_script_hash_file(fingerprint) need_run = not ( - _SCRIPT_HASH_FILE.exists() - and _SCRIPT_HASH_FILE.read_text(encoding="utf-8").strip() + script_hash_file.exists() + and json.loads(script_hash_file.read_text(encoding="utf-8")).get( + "script_fingerprint" + ) == fingerprint ) if need_run: @@ -186,12 +270,16 @@ async def init(): for sql in sql_list: # 对于 ALTER TABLE 操作,先检查表是否存在 - if sql.strip().upper().startswith("ALTER TABLE"): - match = re.match( - r"ALTER\s+TABLE\s+(\w+)", sql, re.IGNORECASE - ) - if match: - table_name = match.group(1) + sql_upper = sql.strip().upper() + if sql_upper.startswith("ALTER TABLE"): + table_name = _extract_alter_table_name(sql) + if table_name: + if not await table_exists(table_name): + logger.debug(f"跳过SQL(表不存在): {sql}") + continue + elif sql_upper.startswith("CREATE INDEX"): + table_name = _extract_create_index_table_name(sql) + if table_name: if not await table_exists(table_name): logger.debug(f"跳过SQL(表不存在): {sql}") continue @@ -236,13 +324,31 @@ async def init(): except Exception as e: logger.debug(f"执行SQL: {sql} 错误...", e=e) logger.debug("SCRIPT_METHOD方法执行完毕!") - _SCRIPT_HASH_FILE.parent.mkdir(parents=True, exist_ok=True) - _SCRIPT_HASH_FILE.write_text(fingerprint, encoding="utf-8") + script_hash_file.parent.mkdir(parents=True, exist_ok=True) + script_hash_file.write_text( + json.dumps( + { + "dialect": urlparse(BotConfig.db_url or "").scheme, + "db_url_hash": hashlib.md5( + (BotConfig.db_url or "").encode() + ).hexdigest(), + "script_fingerprint": fingerprint, + }, + ensure_ascii=False, + indent=2, + ), + encoding="utf-8", + ) else: logger.debug("迁移脚本无变化,跳过执行") + # Tortoise may emit column comments/index SQL during generate_schemas(). + # On existing databases with newly added nullable fields, PostgreSQL can + # fail before the post-generate SchemaGuard gets a chance to repair drift. + await repair_safe_schema_drift() logger.debug("开始生成数据库表结构...") await Tortoise.generate_schemas() logger.debug("数据库表结构生成完毕!") + await repair_safe_schema_drift() logger.info("Database loaded successfully!") except Exception as e: raise DbConnectError(f"数据库连接错误... e:{e}") from e diff --git a/zhenxun/services/db_context/config.py b/zhenxun/services/db_context/config.py index ae6d6b8c..c4889ffe 100644 --- a/zhenxun/services/db_context/config.py +++ b/zhenxun/services/db_context/config.py @@ -24,7 +24,8 @@ MYSQL_CONFIG = { SQLITE_CONFIG = { "journal_mode": "WAL", # 提高并发写入性能 - "timeout": 30, # 锁等待超时(可选) + "busy_timeout": 30000, # SQLite 锁等待超时,单位毫秒 + "foreign_keys": "ON", } diff --git a/zhenxun/services/db_context/schema_guard.py b/zhenxun/services/db_context/schema_guard.py new file mode 100644 index 00000000..fd00ae96 --- /dev/null +++ b/zhenxun/services/db_context/schema_guard.py @@ -0,0 +1,593 @@ +from __future__ import annotations + +import asyncio +import contextlib +from dataclasses import dataclass +from datetime import datetime +import json +from pathlib import Path +from typing import Any, Literal + +from tortoise import Tortoise +from tortoise.exceptions import OperationalError + +from zhenxun.services.log import logger + +from .config import DB_TIMEOUT_SECONDS, LOG_COMMAND + +Dialect = Literal["sqlite", "postgres", "mysql", "unknown"] + + +@dataclass(slots=True) +class SchemaGuardResult: + checked_tables: int = 0 + repaired_columns: int = 0 + repaired_indexes: int = 0 + skipped_columns: int = 0 + skipped_indexes: int = 0 + type_mismatches: int = 0 + warnings: int = 0 + drift: list[dict[str, Any]] | None = None + + +@dataclass(slots=True) +class ColumnInfo: + name: str + data_type: str + nullable: bool | None = None + default: str | None = None + + +def _quote_identifier(identifier: str, dialect: Dialect) -> str: + escaped = identifier.replace('"', '""') + if dialect == "mysql": + return f"`{identifier.replace('`', '``')}`" + return f'"{escaped}"' + + +def _connection_dialect(connection: Any) -> Dialect: + capabilities = getattr(connection, "capabilities", None) + raw = str(getattr(capabilities, "dialect", "") or "").lower() + if raw.startswith("sqlite"): + return "sqlite" + if raw.startswith("postgres"): + return "postgres" + if raw.startswith("mysql"): + return "mysql" + return "unknown" + + +def _is_safe_missing_field(field: Any) -> bool: + if getattr(field, "pk", False): + return False + if getattr(field, "generated", False): + return False + if bool(getattr(field, "null", False)): + return True + default = getattr(field, "default", None) + return default is not None + + +def _field_sql_type(field: Any, dialect: Dialect) -> str | None: + if dialect != "unknown" and hasattr(field, "get_for_dialect"): + with contextlib.suppress(Exception): + value = field.get_for_dialect(dialect, "SQL_TYPE") + if value: + return str(value) + value = getattr(field, "SQL_TYPE", None) + return str(value) if value else None + + +def _is_db_field(field: Any, dialect: Dialect) -> bool: + if getattr(field, "virtual", False): + return False + return _field_sql_type(field, dialect) is not None + + +def _field_default_sql(field: Any, dialect: Dialect) -> str: + default = getattr(field, "default", None) + if default is None or callable(default): + return "" + if isinstance(default, bool): + if dialect == "postgres": + return " DEFAULT TRUE" if default else " DEFAULT FALSE" + return " DEFAULT 1" if default else " DEFAULT 0" + if isinstance(default, int | float): + return f" DEFAULT {default}" + if isinstance(default, str): + escaped = default.replace("'", "''") + return f" DEFAULT '{escaped}'" + return "" + + +def _missing_column_sql( + table: str, + column: str, + field: Any, + dialect: Dialect, +) -> str | None: + sql_type = _field_sql_type(field, dialect) + if not sql_type: + return None + table_sql = _quote_identifier(table, dialect) + column_sql = _quote_identifier(column, dialect) + null_sql = "" if bool(getattr(field, "null", False)) else " NOT NULL" + default_sql = _field_default_sql(field, dialect) + if not bool(getattr(field, "null", False)) and not default_sql: + return None + if dialect == "postgres": + return ( + f"ALTER TABLE {table_sql} ADD COLUMN IF NOT EXISTS " + f"{column_sql} {sql_type}{default_sql}{null_sql}" + ) + return ( + f"ALTER TABLE {table_sql} ADD COLUMN " + f"{column_sql} {sql_type}{default_sql}{null_sql}" + ) + + +async def _table_columns( + connection: Any, table: str, dialect: Dialect +) -> dict[str, ColumnInfo] | None: + if dialect == "sqlite": + rows = await connection.execute_query_dict( + f"PRAGMA table_xinfo({_quote_identifier(table, dialect)})" + ) + if not rows: + return None + return { + str(row.get("name")): ColumnInfo( + name=str(row.get("name")), + data_type=str(row.get("type") or ""), + nullable=not bool(row.get("notnull")), + default=str(row.get("dflt_value")) + if row.get("dflt_value") is not None + else None, + ) + for row in rows + if row.get("name") + } + if dialect == "postgres": + rows = await connection.execute_query_dict( + "SELECT column_name, data_type, is_nullable, column_default " + "FROM information_schema.columns " + "WHERE table_schema = current_schema() AND table_name = $1", + [table], + ) + if not rows: + return None + return { + str(row.get("column_name")): ColumnInfo( + name=str(row.get("column_name")), + data_type=str(row.get("data_type") or ""), + nullable=str(row.get("is_nullable") or "").upper() == "YES", + default=str(row.get("column_default")) + if row.get("column_default") is not None + else None, + ) + for row in rows + if row.get("column_name") + } + if dialect == "mysql": + rows = await connection.execute_query_dict( + "SELECT COLUMN_NAME AS column_name, COLUMN_TYPE AS column_type, " + "IS_NULLABLE AS is_nullable, COLUMN_DEFAULT AS column_default " + "FROM INFORMATION_SCHEMA.COLUMNS " + "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = %s", + [table], + ) + if not rows: + return None + return { + str(row.get("column_name")): ColumnInfo( + name=str(row.get("column_name")), + data_type=str(row.get("column_type") or ""), + nullable=str(row.get("is_nullable") or "").upper() == "YES", + default=str(row.get("column_default")) + if row.get("column_default") is not None + else None, + ) + for row in rows + if row.get("column_name") + } + return None + + +def _normalize_type(value: str, dialect: Dialect) -> str: + text = value.lower().strip() + if not text: + return "" + if dialect == "sqlite": + if "int" in text: + return "integer" + if any(part in text for part in ("char", "clob", "text", "varchar")): + return "text" + if any(part in text for part in ("real", "floa", "doub")): + return "real" + if "blob" in text: + return "blob" + if "bool" in text: + return "integer" + if "timestamp" in text or "datetime" in text: + return "text" + return text + if dialect == "postgres": + text = text.replace("character varying", "varchar") + if text.startswith("timestamp"): + return "timestamp" + if text in {"boolean", "bool"}: + return "boolean" + return text + if dialect == "mysql": + if text.startswith("tinyint(1)") or text == "bool" or text == "boolean": + return "boolean" + if text.startswith("datetime") or text.startswith("timestamp"): + return "datetime" + return text + return text + + +def _field_type_compatible(field: Any, column: ColumnInfo, dialect: Dialect) -> bool: + expected = _field_sql_type(field, dialect) + if not expected: + return True + expected_norm = _normalize_type(expected, dialect) + actual_norm = _normalize_type(column.data_type, dialect) + if not expected_norm or not actual_norm: + return True + if expected_norm == actual_norm: + return True + if dialect == "sqlite": + # SQLite affinity is intentionally loose; varchar/text and bool/integer are + # compatible enough for startup validation. + compatible = { + ("varchar", "text"), + ("text", "varchar"), + ("boolean", "integer"), + ("integer", "boolean"), + } + return (expected_norm, actual_norm) in compatible + return False + + +def _index_name(table: str, columns: tuple[str, ...]) -> str: + raw = f"idx_{table}_{'_'.join(columns)}" + return raw[:62] + + +async def _table_indexes( + connection: Any, table: str, dialect: Dialect +) -> set[tuple[str, ...]]: + indexes: set[tuple[str, ...]] = set() + if dialect == "sqlite": + rows = await connection.execute_query_dict( + f"PRAGMA index_list({_quote_identifier(table, dialect)})" + ) + for row in rows: + if bool(row.get("unique")): + continue + index_name = row.get("name") + if not index_name: + continue + info = await connection.execute_query_dict( + f"PRAGMA index_info({_quote_identifier(str(index_name), dialect)})" + ) + columns = tuple( + str(item.get("name")) + for item in sorted(info, key=lambda item: int(item.get("seqno") or 0)) + if item.get("name") + ) + if columns: + indexes.add(columns) + return indexes + if dialect == "postgres": + rows = await connection.execute_query_dict( + "SELECT indexname, indexdef FROM pg_indexes " + "WHERE schemaname = current_schema() AND tablename = $1", + [table], + ) + for row in rows: + indexdef = str(row.get("indexdef") or "") + if " UNIQUE INDEX " in indexdef.upper(): + continue + start = indexdef.rfind("(") + end = indexdef.rfind(")") + if start < 0 or end <= start: + continue + columns = tuple( + part.strip().strip('"') + for part in indexdef[start + 1 : end].split(",") + if part.strip() + ) + if columns: + indexes.add(columns) + return indexes + if dialect == "mysql": + rows = await connection.execute_query_dict( + "SELECT INDEX_NAME, COLUMN_NAME, SEQ_IN_INDEX, NON_UNIQUE " + "FROM INFORMATION_SCHEMA.STATISTICS " + "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = %s", + [table], + ) + grouped: dict[str, list[tuple[int, str]]] = {} + unique_names: set[str] = set() + for row in rows: + name = str(row.get("INDEX_NAME") or "") + column = str(row.get("COLUMN_NAME") or "") + if not name or not column: + continue + if int(row.get("NON_UNIQUE") or 0) == 0: + unique_names.add(name) + continue + grouped.setdefault(name, []).append( + (int(row.get("SEQ_IN_INDEX") or 0), column) + ) + for name, items in grouped.items(): + if name in unique_names: + continue + columns = tuple(column for _, column in sorted(items)) + if columns: + indexes.add(columns) + return indexes + return indexes + + +def _index_sql(table: str, columns: tuple[str, ...], dialect: Dialect) -> str | None: + if not columns: + return None + table_sql = _quote_identifier(table, dialect) + index_sql = _quote_identifier(_index_name(table, columns), dialect) + columns_sql = ", ".join(_quote_identifier(column, dialect) for column in columns) + if dialect in {"sqlite", "postgres"}: + return ( + f"CREATE INDEX IF NOT EXISTS {index_sql} " f"ON {table_sql}({columns_sql})" + ) + if dialect == "mysql": + return f"CREATE INDEX {index_sql} ON {table_sql}({columns_sql})" + return None + + +def _meta_index_columns(index: Any) -> tuple[str, ...]: + if isinstance(index, str): + return (index,) + if isinstance(index, list | tuple): + return tuple(str(column) for column in index if column) + fields = getattr(index, "fields", None) + if fields: + return tuple(str(column) for column in fields if column) + return () + + +def _empty_drift(table: str) -> dict[str, Any]: + return { + "table": table, + "missing_columns": [], + "extra_columns": [], + "type_mismatches": [], + "missing_indexes": [], + "unsafe_changes": [], + } + + +async def _write_schema_report(result: SchemaGuardResult) -> None: + report = { + "created_at": datetime.now().isoformat(timespec="seconds"), + "summary": { + "checked_tables": result.checked_tables, + "repaired_columns": result.repaired_columns, + "repaired_indexes": result.repaired_indexes, + "skipped_columns": result.skipped_columns, + "skipped_indexes": result.skipped_indexes, + "type_mismatches": result.type_mismatches, + "warnings": result.warnings, + }, + "tables": result.drift or [], + } + path = Path() / "data" / "db" / "schema_report.json" + path.parent.mkdir(parents=True, exist_ok=True) + await asyncio.to_thread( + path.write_text, + json.dumps(report, ensure_ascii=False, indent=2), + "utf-8", + ) + + +async def repair_safe_schema_drift() -> SchemaGuardResult: + """Repair low-risk missing columns after Tortoise schema creation. + + This guard intentionally avoids type changes, constraints, unique indexes, foreign + keys, and SQLite table rebuilds. It is a startup-only safety net for model/schema + drift caused by copied databases or skipped legacy script hashes. + """ + result = SchemaGuardResult() + result.drift = [] + connection = Tortoise.get_connection("default") + dialect = _connection_dialect(connection) + if dialect == "unknown": + logger.debug("SchemaGuard 跳过未知数据库方言", LOG_COMMAND) + return result + + app = Tortoise.apps.get("models", {}) + for model in app.values(): + meta = getattr(model, "_meta", None) + table = getattr(meta, "db_table", None) or getattr(meta, "table", None) + if not table: + continue + try: + columns = await _table_columns(connection, table, dialect) + except Exception as exc: + result.warnings += 1 + logger.debug(f"SchemaGuard 检查表 {table} 失败", LOG_COMMAND, e=exc) + continue + if columns is None: + continue + result.checked_tables += 1 + drift = _empty_drift(table) + fields_map = getattr(meta, "fields_map", {}) or {} + expected_sources: set[str] = set() + for field_name, field in fields_map.items(): + if not _is_db_field(field, dialect): + continue + source = str(getattr(field, "source_field", None) or field_name) + expected_sources.add(source) + if source in columns: + if not _field_type_compatible(field, columns[source], dialect): + result.type_mismatches += 1 + drift["type_mismatches"].append( + { + "column": source, + "expected": _field_sql_type(field, dialect), + "actual": columns[source].data_type, + } + ) + continue + drift["missing_columns"].append(source) + if not _is_safe_missing_field(field): + result.skipped_columns += 1 + drift["unsafe_changes"].append( + { + "kind": "missing_required_column", + "column": source, + } + ) + logger.debug( + f"SchemaGuard 跳过非低风险缺字段: {table}.{source}", + LOG_COMMAND, + ) + continue + sql = _missing_column_sql(table, source, field, dialect) + if not sql: + result.skipped_columns += 1 + logger.debug( + f"SchemaGuard 无法生成补字段 SQL: {table}.{source}", + LOG_COMMAND, + ) + continue + try: + await asyncio.wait_for( + connection.execute_query_dict(sql), + timeout=DB_TIMEOUT_SECONDS, + ) + columns[source] = ColumnInfo(name=source, data_type="") + result.repaired_columns += 1 + logger.info(f"SchemaGuard 已补齐字段: {table}.{source}", LOG_COMMAND) + except OperationalError as exc: + err = str(exc).lower() + if any( + text in err + for text in ("duplicate column", "already exists", "已存在") + ): + columns[source] = ColumnInfo(name=source, data_type="") + continue + result.warnings += 1 + logger.warning( + f"SchemaGuard 补齐字段失败: {table}.{source}", + LOG_COMMAND, + e=exc, + ) + except Exception as exc: + result.warnings += 1 + logger.warning( + f"SchemaGuard 补齐字段失败: {table}.{source}", + LOG_COMMAND, + e=exc, + ) + drift["extra_columns"] = sorted( + column for column in columns.keys() if column not in expected_sources + ) + try: + existing_indexes = await _table_indexes(connection, table, dialect) + except Exception as exc: + result.warnings += 1 + logger.debug(f"SchemaGuard 检查索引 {table} 失败", LOG_COMMAND, e=exc) + existing_indexes = set() + indexes = getattr(meta, "indexes", ()) or () + for index in indexes: + index_columns = _meta_index_columns(index) + if not index_columns or index_columns in existing_indexes: + continue + drift["missing_indexes"].append(list(index_columns)) + if not all(column in columns for column in index_columns): + result.skipped_indexes += 1 + logger.debug( + f"SchemaGuard 跳过缺字段索引: {table}.{index_columns}", + LOG_COMMAND, + ) + continue + sql = _index_sql(table, index_columns, dialect) + if not sql: + result.skipped_indexes += 1 + continue + try: + await asyncio.wait_for( + connection.execute_query_dict(sql), + timeout=DB_TIMEOUT_SECONDS, + ) + existing_indexes.add(index_columns) + result.repaired_indexes += 1 + logger.debug( + f"SchemaGuard 已补齐索引: {table}.{index_columns}", + LOG_COMMAND, + ) + except OperationalError as exc: + err = str(exc).lower() + if any( + text in err for text in ("already exists", "duplicate", "已存在") + ): + existing_indexes.add(index_columns) + continue + result.warnings += 1 + logger.warning( + f"SchemaGuard 补齐索引失败: {table}.{index_columns}", + LOG_COMMAND, + e=exc, + ) + if any( + drift[key] + for key in ( + "missing_columns", + "extra_columns", + "type_mismatches", + "missing_indexes", + "unsafe_changes", + ) + ): + result.drift.append(drift) + + logger.info( + "SchemaGuard 完成: " + f"checked={result.checked_tables}, " + f"repaired_columns={result.repaired_columns}, " + f"repaired_indexes={result.repaired_indexes}, " + f"skipped_columns={result.skipped_columns}, " + f"skipped_indexes={result.skipped_indexes}, " + f"type_mismatches={result.type_mismatches}, " + f"warnings={result.warnings}", + LOG_COMMAND, + ) + with contextlib.suppress(Exception): + await _write_schema_report(result) + return result + + +async def repair_table_schema(table_name: str) -> SchemaGuardResult: + """Repair one table by reusing the startup guard path. + + The current guard is cheap enough and already table-scoped internally by model + metadata, so this helper keeps write-path recovery simple and conservative. + """ + result = SchemaGuardResult() + result.drift = [] + connection = Tortoise.get_connection("default") + dialect = _connection_dialect(connection) + if dialect == "unknown": + return result + app = Tortoise.apps.get("models", {}) + for model in app.values(): + meta = getattr(model, "_meta", None) + table = getattr(meta, "db_table", None) or getattr(meta, "table", None) + if table == table_name: + # Keep the implementation conservative: repairing all tables is still + # startup-style work and avoids a second partial code path. + return await repair_safe_schema_drift() + return result diff --git a/zhenxun/services/db_context/schema_ops.py b/zhenxun/services/db_context/schema_ops.py new file mode 100644 index 00000000..daadecce --- /dev/null +++ b/zhenxun/services/db_context/schema_ops.py @@ -0,0 +1,180 @@ +from __future__ import annotations + +from dataclasses import dataclass +from enum import Enum +from typing import Literal, Protocol + +Dialect = Literal["sqlite", "postgres", "mysql", "unknown"] + + +class SchemaOpRisk(str, Enum): + SAFE = "safe" + GUARDED = "guarded" + MANUAL = "manual" + + +class SchemaOp(Protocol): + risk: SchemaOpRisk + + def to_sql(self, dialect: Dialect) -> list[str]: ... + + +def quote_identifier(identifier: str, dialect: Dialect) -> str: + if dialect == "mysql": + return f"`{identifier.replace('`', '``')}`" + return f'"{identifier.replace(chr(34), chr(34) + chr(34))}"' + + +def _column_type(column_type: str | dict[str, str], dialect: Dialect) -> str: + if isinstance(column_type, dict): + return column_type.get(dialect) or column_type.get("default") or "TEXT" + return column_type + + +def _default_sql(default: str | float | bool | None, dialect: Dialect) -> str: + if default is None: + return "" + if isinstance(default, bool): + if dialect == "postgres": + return " DEFAULT TRUE" if default else " DEFAULT FALSE" + return " DEFAULT 1" if default else " DEFAULT 0" + if isinstance(default, int | float): + return f" DEFAULT {default}" + escaped = default.replace("'", "''") + return f" DEFAULT '{escaped}'" + + +@dataclass(frozen=True, slots=True) +class AddColumn: + table: str + column: str + column_type: str | dict[str, str] + nullable: bool = True + default: str | float | bool | None = None + risk: SchemaOpRisk = SchemaOpRisk.SAFE + + def to_sql(self, dialect: Dialect) -> list[str]: + table = quote_identifier(self.table, dialect) + column = quote_identifier(self.column, dialect) + column_type = _column_type(self.column_type, dialect) + null_sql = "" if self.nullable else " NOT NULL" + default_sql = _default_sql(self.default, dialect) + if not self.nullable and not default_sql: + return [] + if dialect == "postgres": + return [ + "ALTER TABLE " + f"{table} ADD COLUMN IF NOT EXISTS {column} " + f"{column_type}{default_sql}{null_sql}" + ] + return [ + "ALTER TABLE " + f"{table} ADD COLUMN {column} {column_type}{default_sql}{null_sql}" + ] + + +@dataclass(frozen=True, slots=True) +class CreateIndex: + table: str + columns: tuple[str, ...] + name: str | None = None + if_not_exists: bool = True + unique: bool = False + where: str | None = None + risk: SchemaOpRisk = SchemaOpRisk.SAFE + + def __init__( + self, + table: str, + columns: tuple[str, ...] | list[str], + name: str | None = None, + if_not_exists: bool = True, + unique: bool = False, + where: str | None = None, + risk: SchemaOpRisk | None = None, + ) -> None: + object.__setattr__(self, "table", table) + object.__setattr__(self, "columns", tuple(columns)) + object.__setattr__(self, "name", name) + object.__setattr__(self, "if_not_exists", if_not_exists) + object.__setattr__(self, "unique", unique) + object.__setattr__(self, "where", where) + resolved_risk = risk or (SchemaOpRisk.MANUAL if unique else SchemaOpRisk.SAFE) + object.__setattr__(self, "risk", resolved_risk) + + def to_sql(self, dialect: Dialect) -> list[str]: + if self.unique or not self.columns: + return [] + name = self.name or f"idx_{self.table}_{'_'.join(self.columns)}"[:62] + if dialect == "mysql" and self.where: + return [] + exists_sql = ( + "IF NOT EXISTS " if self.if_not_exists and dialect != "mysql" else "" + ) + table = quote_identifier(self.table, dialect) + index = quote_identifier(name, dialect) + columns = ", ".join( + quote_identifier(column, dialect) for column in self.columns + ) + where = f" WHERE {self.where}" if self.where and dialect != "mysql" else "" + return [f"CREATE INDEX {exists_sql}{index} ON {table}({columns}){where}"] + + +@dataclass(frozen=True, slots=True) +class RenameColumn: + table: str + old: str + new: str + risk: SchemaOpRisk = SchemaOpRisk.GUARDED + + def to_sql(self, dialect: Dialect) -> list[str]: + table = quote_identifier(self.table, dialect) + old = quote_identifier(self.old, dialect) + new = quote_identifier(self.new, dialect) + return [f"ALTER TABLE {table} RENAME COLUMN {old} TO {new}"] + + +@dataclass(frozen=True, slots=True) +class DropColumn: + table: str + column: str + risk: SchemaOpRisk = SchemaOpRisk.GUARDED + + def to_sql(self, dialect: Dialect) -> list[str]: + table = quote_identifier(self.table, dialect) + column = quote_identifier(self.column, dialect) + return [f"ALTER TABLE {table} DROP COLUMN {column}"] + + +@dataclass(frozen=True, slots=True) +class AlterColumnType: + table: str + column: str + column_type: str | dict[str, str] + nullable: bool | None = None + risk: SchemaOpRisk = SchemaOpRisk.GUARDED + + def to_sql(self, dialect: Dialect) -> list[str]: + column_type = _column_type(self.column_type, dialect) + table = quote_identifier(self.table, dialect) + column = quote_identifier(self.column, dialect) + if dialect == "postgres": + return [f"ALTER TABLE {table} ALTER COLUMN {column} TYPE {column_type}"] + if dialect == "mysql": + null_sql = " NULL" if self.nullable else " NOT NULL" + if self.nullable is None: + null_sql = "" + return [ + f"ALTER TABLE {table} MODIFY COLUMN {column} {column_type}{null_sql}" + ] + return [] + + +def normalize_schema_ops(items: list[str | SchemaOp], dialect: Dialect) -> list[str]: + sql_list: list[str] = [] + for item in items: + if isinstance(item, str): + sql_list.append(item) + else: + sql_list.extend(item.to_sql(dialect)) + return sql_list diff --git a/zhenxun/services/hot_query_cache.py b/zhenxun/services/hot_query_cache.py new file mode 100644 index 00000000..55bf67b4 --- /dev/null +++ b/zhenxun/services/hot_query_cache.py @@ -0,0 +1,482 @@ +from __future__ import annotations + +import asyncio +from collections.abc import Iterable +from dataclasses import dataclass +from datetime import datetime +from typing import Any, Literal + +from tortoise.functions import Count + +from zhenxun.services.cache.bounded_ttl import BoundedTTLCache + + +@dataclass(frozen=True, slots=True) +class GroupMemberSnapshot: + id: int + user_id: str + user_name: str + group_id: str + user_join_time: datetime | None + uid: int | None + platform: str | None + + +def _member_cache_sizeof(members: tuple[GroupMemberSnapshot, ...]) -> int: + size = 0 + for member in members: + size += 96 + size += len(member.user_id) + len(member.user_name) + len(member.group_id) + size += len(member.platform or "") + return size + + +_GROUP_MEMBER_CACHE = BoundedTTLCache[str, tuple[GroupMemberSnapshot, ...]]( + "hot_group_info_users", + ttl_seconds=45, + max_items=512, + max_total_bytes=32 * 1024 * 1024, + sizeof=_member_cache_sizeof, +) +_GROUP_USER_IDS_CACHE = BoundedTTLCache[str, tuple[str, ...]]( + "hot_group_info_user_ids", + ttl_seconds=45, + max_items=2048, +) +_GROUP_MEMBER_BY_ID_CACHE = BoundedTTLCache[str, tuple[GroupMemberSnapshot | None]]( + "hot_group_info_user_by_id", + ttl_seconds=45, + max_items=50000, +) +_USER_GROUP_CACHE = BoundedTTLCache[str, tuple[str, ...]]( + "hot_group_info_user_groups", + ttl_seconds=45, + max_items=4096, +) +_USER_NAME_CACHE = BoundedTTLCache[str, str]( + "hot_group_info_user_names", + ttl_seconds=45, + max_items=20000, +) +_CHAT_RANK_CACHE = BoundedTTLCache[str, tuple[tuple[str, int], ...]]( + "hot_chat_history_rank", + ttl_seconds=20, + max_items=512, +) +_CHAT_FIRST_MSG_CACHE = BoundedTTLCache[str, tuple[datetime | None]]( + "hot_chat_history_first_msg", + ttl_seconds=300, + max_items=2048, +) +_STATISTICS_COUNT_CACHE = BoundedTTLCache[str, tuple[tuple[str, int], ...]]( + "hot_statistics_plugin_counts", + ttl_seconds=20, + max_items=512, +) + +_GROUP_MEMBER_LOCKS: dict[str, asyncio.Lock] = {} +_GROUP_USER_IDS_LOCKS: dict[str, asyncio.Lock] = {} +_USER_GROUP_LOCKS: dict[str, asyncio.Lock] = {} +_CHAT_RANK_LOCKS: dict[str, asyncio.Lock] = {} +_CHAT_FIRST_MSG_LOCKS: dict[str, asyncio.Lock] = {} +_STATISTICS_LOCKS: dict[str, asyncio.Lock] = {} +_MAX_LOCK_POOL_SIZE = 4096 + + +def _get_lock(pool: dict[str, asyncio.Lock], key: str) -> asyncio.Lock: + lock = pool.get(key) + if lock is None: + if len(pool) >= _MAX_LOCK_POOL_SIZE: + for old_key, old_lock in list(pool.items()): + if not old_lock.locked(): + pool.pop(old_key, None) + break + lock = asyncio.Lock() + pool[key] = lock + return lock + + +def _normalize_id(value: object) -> str: + return str(value or "") + + +def _normalize_ids(values: Iterable[object] | None) -> tuple[str, ...] | None: + if values is None: + return None + return tuple(dict.fromkeys(v for value in values if (v := _normalize_id(value)))) + + +async def get_group_members( + group_id: str | int | None, +) -> tuple[GroupMemberSnapshot, ...]: + """Return lightweight group-member snapshots with a short runtime TTL.""" + group_key = _normalize_id(group_id) + if not group_key: + return () + + cached = await _GROUP_MEMBER_CACHE.get(group_key) + if cached is not None: + return cached + + lock = _get_lock(_GROUP_MEMBER_LOCKS, group_key) + async with lock: + cached = await _GROUP_MEMBER_CACHE.get(group_key) + if cached is not None: + return cached + + from zhenxun.models.group_member_info import GroupInfoUser + + rows = await GroupInfoUser.filter(group_id=group_key).values_list( + "id", + "user_id", + "user_name", + "user_join_time", + "uid", + "platform", + ) + members = tuple( + GroupMemberSnapshot( + id=int(row[0] or 0), + user_id=str(row[1] or ""), + user_name=str(row[2] or ""), + group_id=group_key, + user_join_time=row[3], + uid=int(row[4]) if row[4] is not None else None, + platform=str(row[5]) if row[5] else None, + ) + for row in rows + if row[1] + ) + await _GROUP_MEMBER_CACHE.set(group_key, members) + await _GROUP_USER_IDS_CACHE.set( + group_key, tuple(member.user_id for member in members) + ) + return members + + +async def get_group_member_map( + group_id: str | int | None, + user_ids: Iterable[object] | None = None, +) -> dict[str, GroupMemberSnapshot]: + group_key = _normalize_id(group_id) + if not group_key: + return {} + wanted = _normalize_ids(user_ids) + if wanted is not None and not wanted: + return {} + if wanted is None: + members = await get_group_members(group_key) + return {member.user_id: member for member in members} + + cached_members = await _GROUP_MEMBER_CACHE.get(group_key) + if cached_members is not None: + wanted_set = set(wanted) + return { + member.user_id: member + for member in cached_members + if member.user_id in wanted_set + } + + result: dict[str, GroupMemberSnapshot] = {} + missing: list[str] = [] + for user_id in wanted: + cache_key = f"{group_key}:{user_id}" + cached = await _GROUP_MEMBER_BY_ID_CACHE.get(cache_key) + if cached is None: + missing.append(user_id) + else: + member = cached[0] + if member is not None: + result[user_id] = member + + if missing: + from zhenxun.models.group_member_info import GroupInfoUser + + rows = await GroupInfoUser.filter( + group_id=group_key, user_id__in=missing + ).values_list( + "id", + "user_id", + "user_name", + "user_join_time", + "uid", + "platform", + ) + found: set[str] = set() + for row in rows: + if not row[1]: + continue + member = GroupMemberSnapshot( + id=int(row[0] or 0), + user_id=str(row[1] or ""), + user_name=str(row[2] or ""), + group_id=group_key, + user_join_time=row[3], + uid=int(row[4]) if row[4] is not None else None, + platform=str(row[5]) if row[5] else None, + ) + result[member.user_id] = member + found.add(member.user_id) + await _GROUP_MEMBER_BY_ID_CACHE.set( + f"{group_key}:{member.user_id}", (member,) + ) + for user_id in missing: + if user_id not in found: + await _GROUP_MEMBER_BY_ID_CACHE.set(f"{group_key}:{user_id}", (None,)) + return result + + +async def get_group_member( + group_id: str | int | None, + user_id: str | int | None, +) -> GroupMemberSnapshot | None: + user_key = _normalize_id(user_id) + if not user_key: + return None + return (await get_group_member_map(group_id, [user_key])).get(user_key) + + +async def get_group_user_ids(group_id: str | int | None) -> set[str]: + group_key = _normalize_id(group_id) + if not group_key: + return set() + cached = await _GROUP_USER_IDS_CACHE.get(group_key) + if cached is not None: + return set(cached) + + cached_members = await _GROUP_MEMBER_CACHE.get(group_key) + if cached_members is not None: + user_ids = tuple(member.user_id for member in cached_members) + await _GROUP_USER_IDS_CACHE.set(group_key, user_ids) + return set(user_ids) + + lock = _get_lock(_GROUP_USER_IDS_LOCKS, group_key) + async with lock: + cached = await _GROUP_USER_IDS_CACHE.get(group_key) + if cached is not None: + return set(cached) + + from zhenxun.models.group_member_info import GroupInfoUser + + rows = await GroupInfoUser.filter(group_id=group_key).values_list( + "user_id", flat=True + ) + user_ids = tuple(str(user_id) for user_id in rows if user_id) + await _GROUP_USER_IDS_CACHE.set(group_key, user_ids) + return set(user_ids) + + +async def get_user_group_ids(user_id: str | int | None) -> list[str]: + user_key = _normalize_id(user_id) + if not user_key: + return [] + + cached = await _USER_GROUP_CACHE.get(user_key) + if cached is not None: + return list(cached) + + lock = _get_lock(_USER_GROUP_LOCKS, user_key) + async with lock: + cached = await _USER_GROUP_CACHE.get(user_key) + if cached is not None: + return list(cached) + + from zhenxun.models.group_member_info import GroupInfoUser + + rows = await GroupInfoUser.filter(user_id=user_key).values_list( + "group_id", flat=True + ) + group_ids = tuple(str(group_id) for group_id in rows if group_id) + await _USER_GROUP_CACHE.set(user_key, group_ids) + return list(group_ids) + + +async def get_member_names( + user_ids: Iterable[object], + group_id: str | int | None = None, +) -> dict[str, str]: + user_keys = _normalize_ids(user_ids) or () + if not user_keys: + return {} + if group_id: + members = await get_group_member_map(group_id, user_keys) + return {user_id: members[user_id].user_name for user_id in members} + + result: dict[str, str] = {} + missing: list[str] = [] + for user_id in user_keys: + cached = await _USER_NAME_CACHE.get(user_id) + if cached is None: + missing.append(user_id) + else: + result[user_id] = cached + + if missing: + from zhenxun.models.group_member_info import GroupInfoUser + + rows = await GroupInfoUser.filter(user_id__in=missing).values_list( + "user_id", "user_name" + ) + for user_id, user_name in rows: + user_key = str(user_id) + if user_key not in result: + result[user_key] = str(user_name or "") + for user_id in missing: + await _USER_NAME_CACHE.set(user_id, result.get(user_id, "")) + return result + + +async def get_member_name( + user_id: str | int | None, + group_id: str | int | None = None, +) -> str | None: + user_key = _normalize_id(user_id) + if not user_key: + return None + return (await get_member_names([user_key], group_id)).get(user_key) or None + + +async def invalidate_group_members( + group_id: str | int | None = None, + user_ids: Iterable[object] | None = None, +) -> None: + if group_id is None: + await _GROUP_MEMBER_CACHE.clear() + await _GROUP_USER_IDS_CACHE.clear() + await _GROUP_MEMBER_BY_ID_CACHE.clear() + return + group_key = _normalize_id(group_id) + await _GROUP_MEMBER_CACHE.delete(group_key) + await _GROUP_USER_IDS_CACHE.delete(group_key) + normalized_ids = _normalize_ids(user_ids) + if normalized_ids is None: + await _GROUP_MEMBER_BY_ID_CACHE.clear() + return + for user_id in normalized_ids: + await _GROUP_MEMBER_BY_ID_CACHE.delete(f"{group_key}:{user_id}") + + +async def invalidate_member_names(user_ids: Iterable[object] | None = None) -> None: + if user_ids is None: + await _USER_NAME_CACHE.clear() + await _USER_GROUP_CACHE.clear() + return + for user_id in _normalize_ids(user_ids) or (): + await _USER_NAME_CACHE.delete(user_id) + await _USER_GROUP_CACHE.delete(user_id) + + +def _datetime_key(value: datetime | None) -> str: + return value.isoformat(" ", timespec="seconds") if value else "" + + +def _date_scope_key(date_scope: tuple[datetime, datetime] | None) -> str: + if not date_scope: + return "" + end_bucket = int(date_scope[1].timestamp() // 20) + return f"{_datetime_key(date_scope[0])}..bucket:{end_bucket}" + + +async def get_chat_history_rank_cached( + model: Any, + gid: str | None, + limit: int = 10, + order: str = "DESC", + date_scope: tuple[datetime, datetime] | None = None, +) -> list[tuple[str, int]]: + key = f"{gid or '*'}:{limit}:{order}:{_date_scope_key(date_scope)}" + cached = await _CHAT_RANK_CACHE.get(key) + if cached is not None: + return list(cached) + + lock = _get_lock(_CHAT_RANK_LOCKS, key) + async with lock: + cached = await _CHAT_RANK_CACHE.get(key) + if cached is not None: + return list(cached) + + order_prefix = "-" if order == "DESC" else "" + query: Any = model.filter(group_id=gid) if gid else model + if date_scope: + filter_scope = ( + date_scope[0].isoformat(" "), + date_scope[1].isoformat(" "), + ) + query = query.filter(create_time__range=filter_scope) + rows = ( + await query.annotate(count=Count("user_id")) + .order_by(f"{order_prefix}count") + .group_by("user_id") + .limit(limit) + .values_list("user_id", "count") + ) + result = tuple((str(user_id), int(count)) for user_id, count in rows) + await _CHAT_RANK_CACHE.set(key, result) + return list(result) + + +async def get_chat_history_first_msg_datetime_cached( + model: Any, + group_id: str | None, +) -> datetime | None: + key = group_id or "*" + cached = await _CHAT_FIRST_MSG_CACHE.get(key) + if cached is not None: + return cached[0] + + lock = _get_lock(_CHAT_FIRST_MSG_LOCKS, key) + async with lock: + cached = await _CHAT_FIRST_MSG_CACHE.get(key) + if cached is not None: + return cached[0] + + query: Any = model.filter(group_id=group_id) if group_id else model.all() + message = await query.order_by("create_time").first() + result = getattr(message, "create_time", None) if message else None + await _CHAT_FIRST_MSG_CACHE.set(key, (result,)) + return result + + +async def get_statistics_plugin_counts_cached( + scope: Literal["global", "user", "group"], + *, + plugin_name: str | None, + start_time: datetime | None, + user_id: str | None = None, + group_id: str | None = None, +) -> list[tuple[str, int]]: + key = ( + f"{scope}:{plugin_name or ''}:{_datetime_key(start_time)}:" + f"{user_id or ''}:{group_id or ''}" + ) + cached = await _STATISTICS_COUNT_CACHE.get(key) + if cached is not None: + return list(cached) + + lock = _get_lock(_STATISTICS_LOCKS, key) + async with lock: + cached = await _STATISTICS_COUNT_CACHE.get(key) + if cached is not None: + return list(cached) + + from zhenxun.models.statistics import Statistics + + query: Any = Statistics + if scope == "user": + query = Statistics.filter(user_id=user_id) + if group_id: + query = query.filter(group_id=group_id) + elif scope == "group": + query = Statistics.filter(group_id=group_id) + if plugin_name: + query = query.filter(plugin_name=plugin_name) + if start_time: + query = query.filter(create_time__gte=start_time) + rows = ( + await query.annotate(count=Count("id")) + .group_by("plugin_name") + .values_list("plugin_name", "count") + ) + result = tuple((str(plugin), int(count)) for plugin, count in rows) + await _STATISTICS_COUNT_CACHE.set(key, result) + return list(result) diff --git a/zhenxun/services/runtime_bootstrap.py b/zhenxun/services/runtime_bootstrap.py index a813fea9..91cdd993 100644 --- a/zhenxun/services/runtime_bootstrap.py +++ b/zhenxun/services/runtime_bootstrap.py @@ -140,6 +140,9 @@ 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 diff --git a/zhenxun/services/send_queue.py b/zhenxun/services/send_queue.py index e381da2d..5c872529 100644 --- a/zhenxun/services/send_queue.py +++ b/zhenxun/services/send_queue.py @@ -1,30 +1,124 @@ import asyncio +from collections import defaultdict +from contextlib import contextmanager +from contextvars import ContextVar +from dataclasses import dataclass import time -from typing import Any +from typing import Any, ClassVar, cast from nonebot.adapters import Bot +from nonebot.adapters.onebot.v11 import Adapter as OneBotV11Adapter +from nonebot.adapters.onebot.v11 import Bot as OneBotV11Bot from zhenxun.services.log import logger _SEND_APIS = {"send_msg", "send_group_msg", "send_private_msg", "send_like"} +_OBSERVED_SEND_APIS = {"send_msg", "send_group_msg", "send_private_msg"} _WORKERS = 3 _MIN_INTERVAL = 0.05 _QUEUE_MAXSIZE = 2000 _SHUTDOWN_DRAIN_TIMEOUT_SECONDS = 3.0 _QUEUE_PRESSURE_LOG_INTERVAL = 10.0 -_QUEUE: asyncio.Queue[tuple[Bot, str, dict[str, Any], asyncio.Future[Any]]] = ( - asyncio.Queue(maxsize=_QUEUE_MAXSIZE) -) +_QUEUE: asyncio.Queue[ + tuple[Bot, str, dict[str, Any], asyncio.Future[Any], str | None] +] = asyncio.Queue(maxsize=_QUEUE_MAXSIZE) _SEND_LOCK = asyncio.Lock() _LAST_SEND_TS = 0.0 _API_SEMAPHORE = asyncio.Semaphore(3) -_ORIG_CALL_API = Bot.call_api +_ORIG_CALL_API = OneBotV11Adapter._call_api _PATCHED = False _WORKER_TASKS: list[asyncio.Task] = [] _QUEUE_TIMEOUT_COUNT = 0 _SEND_LIKE_DROP_COUNT = 0 _LAST_QUEUE_PRESSURE_LOG = 0.0 _STOPPING = False +_CURRENT_SEND_TRACE_ID: ContextVar[str | None] = ContextVar( + "zhenxun_send_trace_id", + default=None, +) +_MAX_OBSERVED_RECORDS_PER_TRACE = 12 +_MAX_OBSERVED_TEXT_LEN = 900 + + +def _send_platform_scope(adapter: Any) -> str: + if adapter is None: + return "unknown" + if isinstance(adapter, OneBotV11Adapter): + return "qq_client" + get_name = getattr(adapter, "get_name", None) + if callable(get_name): + try: + name = str(get_name()).lower() + except Exception: + name = "" + else: + name = adapter.__class__.__name__.lower() + if name == "qq" or "qq" in name: + return "qq_api" + return name or "unknown" + + +@dataclass(frozen=True) +class SendObservation: + trace_id: str + api: str + text: str + raw_message: str + result: Any + timestamp: float + + +class SendObserver: + _records: ClassVar[dict[str, list[SendObservation]]] = defaultdict(list) + + @classmethod + @contextmanager + def activate(cls, trace_id: str): + trace_key = str(trace_id or "").strip() + token = _CURRENT_SEND_TRACE_ID.set(trace_key or None) + try: + yield + finally: + _CURRENT_SEND_TRACE_ID.reset(token) + + @classmethod + def record( + cls, + *, + trace_id: str | None, + api: str, + data: dict[str, Any], + result: Any, + ) -> None: + trace_key = str(trace_id or "").strip() + if not trace_key or api not in _OBSERVED_SEND_APIS: + return + target = cls._records[trace_key] + if len(target) >= _MAX_OBSERVED_RECORDS_PER_TRACE: + return + raw_message = _message_to_text(data.get("message")) + target.append( + SendObservation( + trace_id=trace_key, + api=api, + text=_compact_text(raw_message), + raw_message=raw_message[:_MAX_OBSERVED_TEXT_LEN], + result=result, + timestamp=time.time(), + ) + ) + + @classmethod + def pop(cls, trace_id: str) -> list[SendObservation]: + return cls._records.pop(str(trace_id or "").strip(), []) + + +def observe_send_trace(trace_id: str): + return SendObserver.activate(trace_id) + + +def pop_send_observations(trace_id: str) -> list[SendObservation]: + return SendObserver.pop(trace_id) async def _rate_limit(): @@ -50,17 +144,45 @@ def _log_queue_pressure(reason: str) -> None: ) -async def _direct_call_api(bot: Bot, api: str, data: dict[str, Any]) -> Any: +async def _direct_call_api( + adapter: OneBotV11Adapter, + bot: Bot, + api: str, + data: dict[str, Any], + trace_id: str | None = None, +) -> Any: await _rate_limit() async with _API_SEMAPHORE: - return await _ORIG_CALL_API(bot, api, **data) + try: + result = await _ORIG_CALL_API( + adapter, + cast(OneBotV11Bot, bot), + api, + **data, + ) + except Exception as exc: + SendObserver.record( + trace_id=trace_id, + api=api, + data=data, + result={"ok": False, "error": str(exc)}, + ) + raise + SendObserver.record(trace_id=trace_id, api=api, data=data, result=result) + return result async def _worker(worker_id: int): while True: - bot, api, data, future = await _QUEUE.get() + bot, api, data, future, trace_id = await _QUEUE.get() try: - result = await _direct_call_api(bot, api, data) + result = await _direct_call_api( + cast(OneBotV11Adapter, bot.adapter), + bot, + api, + data, + trace_id=trace_id, + ) if not future.done(): future.set_result(result) except asyncio.CancelledError: @@ -80,15 +202,28 @@ async def _worker(worker_id: int): _QUEUE.task_done() -async def _queued_call_api(self: Bot, api: str, **data: Any): +async def _queued_call_api( + adapter: OneBotV11Adapter, + bot: Bot, + api: str, + **data: Any, +): + if _send_platform_scope(adapter) != "qq_client": + return await _ORIG_CALL_API(adapter, cast(OneBotV11Bot, bot), api, **data) if api not in _SEND_APIS: - return await _ORIG_CALL_API(self, api, **data) + return await _ORIG_CALL_API(adapter, cast(OneBotV11Bot, bot), api, **data) if _STOPPING: - return await _direct_call_api(self, api, data) + return await _direct_call_api( + adapter, + bot, + api, + data, + trace_id=_CURRENT_SEND_TRACE_ID.get(), + ) loop = asyncio.get_running_loop() future: asyncio.Future[Any] = loop.create_future() - queue_item = (self, api, data, future) + queue_item = (bot, api, data, future, _CURRENT_SEND_TRACE_ID.get()) try: _QUEUE.put_nowait(queue_item) except asyncio.QueueFull: @@ -100,7 +235,13 @@ async def _queued_call_api(self: Bot, api: str, **data: Any): global _QUEUE_TIMEOUT_COUNT _QUEUE_TIMEOUT_COUNT += 1 _log_queue_pressure(f"{api} fallback to direct send because queue is full") - return await _direct_call_api(self, api, data) + return await _direct_call_api( + adapter, + bot, + api, + data, + trace_id=_CURRENT_SEND_TRACE_ID.get(), + ) return await future @@ -108,7 +249,7 @@ def _drain_pending_futures(reason: str) -> int: drained = 0 while True: try: - _, _, _, future = _QUEUE.get_nowait() + _, _, _, future, _ = _QUEUE.get_nowait() except asyncio.QueueEmpty: break if not future.done(): @@ -122,7 +263,7 @@ def patch_send_queue() -> None: global _PATCHED if _PATCHED: return - Bot.call_api = _queued_call_api # type: ignore[assignment] + OneBotV11Adapter._call_api = _queued_call_api # type: ignore[assignment] _PATCHED = True @@ -130,7 +271,7 @@ def unpatch_send_queue() -> None: global _PATCHED if not _PATCHED: return - Bot.call_api = _ORIG_CALL_API # type: ignore[assignment] + OneBotV11Adapter._call_api = _ORIG_CALL_API # type: ignore[assignment] _PATCHED = False @@ -138,12 +279,37 @@ async def start_send_queue() -> None: global _STOPPING patch_send_queue() _STOPPING = False + _WORKER_TASKS[:] = [task for task in _WORKER_TASKS if not task.done()] if _WORKER_TASKS: return for idx in range(_WORKERS): _WORKER_TASKS.append(asyncio.create_task(_worker(idx))) +def _message_to_text(message: Any) -> str: + if message is None: + return "" + if hasattr(message, "extract_plain_text"): + try: + text = str(message.extract_plain_text()) + if text.strip(): + return text + except Exception: + pass + try: + return str(message) + except Exception as exc: + logger.debug(f"send observation stringify failed: {exc}") + return "" + + +def _compact_text(text: str) -> str: + normalized = " ".join(str(text or "").split()) + if len(normalized) <= _MAX_OBSERVED_TEXT_LEN: + return normalized + return normalized[: _MAX_OBSERVED_TEXT_LEN - 1].rstrip() + "…" + + async def stop_send_queue() -> None: global _STOPPING _STOPPING = True diff --git a/zhenxun/utils/platform.py b/zhenxun/utils/platform.py index 16b07cff..416b37a5 100644 --- a/zhenxun/utils/platform.py +++ b/zhenxun/utils/platform.py @@ -24,6 +24,19 @@ from zhenxun.utils.message import MessageUtils driver = nonebot.get_driver() +def _adapter_name(bot: Bot) -> str: + adapter = getattr(bot, "adapter", None) + if adapter is None: + return "" + get_name = getattr(adapter, "get_name", None) + if callable(get_name): + try: + return str(get_name()).lower() + except Exception: + return "" + return adapter.__class__.__name__.lower() + + class UserData(BaseModel): name: str """昵称""" @@ -380,6 +393,47 @@ class PlatformUtils: return "qq" if platform.startswith("qq") else platform return "unknown" + @classmethod + def get_platform_scope(cls, t: Bot | Uninfo | object) -> str: + """获取细粒度平台作用域,不改变旧 get_platform 返回值。""" + if isinstance(t, Bot): + adapter_name = _adapter_name(t) + if "onebot" in adapter_name: + return "qq_client" + if adapter_name == "qq" or "qq" in adapter_name: + return "qq_api" + if BotConfig.get_qbot_uid(t.self_id): + return "qq_api" + return adapter_name or cls.get_platform(t) + + scope = str(getattr(t, "scope", "") or "").lower() + 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" + + adapter = getattr(t, "adapter", None) + if adapter is not None: + name = ( + adapter.get_name().lower() + if callable(getattr(adapter, "get_name", None)) + else adapter.__class__.__name__.lower() + ) + if "onebot" in name: + return "qq_client" + if name == "qq" or "qq" in name: + return "qq_api" + return name + + platform = str(getattr(t, "platform", "") or "").lower() + return "qq_client" if platform == "qq" else platform or "unknown" + @classmethod def is_forward_merge_supported(cls, t: Bot | Uninfo) -> bool: """是否支持转发消息