mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-10 14:20:04 +08:00
bugfix:修复notice事件扩散问题 (#2132)
* bugfix:修复notice事件扩散问题 * 优化并发调度 * bugfix:修复签到样式 * bugfix:功能调用统计修复 * bugfix:修复私聊时功能调用统计显示已退群问题 * 提高插件适配兼容性 * 优化发送队列 * 修改权限检查设计 * 继续修改权限检查设计 * 完善权限检查设计 * 优化sqlite配置 * 优化数据库初始化 * 代码整理,无用代码清理 * bugfix:修复启动时数据库校验问题 * bugfix:修复预算裁剪过于激进问题
This commit is contained in:
@@ -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 "群组成员信息更新完成!"
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -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"<m>{escape_tag(bot.type)} {escape_tag(bot.self_id)}</m> | "
|
||||
try:
|
||||
log_msg += event.get_log_string()
|
||||
except no_log_exception:
|
||||
show_log = False
|
||||
if show_log:
|
||||
logger_.opt(colors=True).success(log_msg)
|
||||
|
||||
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(
|
||||
"<r><bg #f8bbd0>Error when checking Matcher.</bg #f8bbd0></r>"
|
||||
),
|
||||
}
|
||||
):
|
||||
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",
|
||||
]
|
||||
@@ -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)
|
||||
|
||||
@@ -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"]
|
||||
@@ -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)
|
||||
@@ -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),
|
||||
),
|
||||
]
|
||||
)
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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"]
|
||||
@@ -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",
|
||||
)
|
||||
@@ -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
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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,
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user