bugfix:修复notice事件扩散问题 (#2132)

* bugfix:修复notice事件扩散问题

* 优化并发调度

* bugfix:修复签到样式

* bugfix:功能调用统计修复

* bugfix:修复私聊时功能调用统计显示已退群问题

* 提高插件适配兼容性

* 优化发送队列

* 修改权限检查设计

* 继续修改权限检查设计

* 完善权限检查设计

* 优化sqlite配置

* 优化数据库初始化

* 代码整理,无用代码清理

* bugfix:修复启动时数据库校验问题

* bugfix:修复预算裁剪过于激进问题
This commit is contained in:
Copaan
2026-05-28 22:57:28 +08:00
committed by GitHub
parent 12fc5663fb
commit 5596497947
57 changed files with 7557 additions and 1522 deletions
@@ -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)
+129 -13
View File
@@ -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",
]
+39 -3
View File
@@ -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",
]
+1 -1
View File
@@ -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)
+3 -9
View File
@@ -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):
+1 -1
View File
@@ -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)