feat: add permission snapshot system to optimize auth checks

This commit is contained in:
HibiKier
2025-12-29 09:20:07 +08:00
parent f86beb928f
commit 96ba8d5a21
8 changed files with 1816 additions and 0 deletions
@@ -0,0 +1,15 @@
"""
权限快照服务模块
提供预聚合的权限检查数据,将多次数据库/缓存查询优化为1-2次
"""
from .models import AuthSnapshot, PluginSnapshot
from .service import AuthSnapshotService, PluginSnapshotService
__all__ = [
"AuthSnapshot",
"AuthSnapshotService",
"PluginSnapshot",
"PluginSnapshotService",
]
+362
View File
@@ -0,0 +1,362 @@
"""
快照构建器
负责从多个数据源聚合数据构建权限快照
"""
import asyncio
import time
from typing import Any
from zhenxun.models.ban_console import BanConsole
from zhenxun.models.bot_console import BotConsole
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.level_user import LevelUser
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.user_console import UserConsole
from zhenxun.services.log import logger
from .models import AuthSnapshot, PluginSnapshot
LOG_COMMAND = "auth_snapshot"
BUILD_TIMEOUT = 5.0 # 构建超时时间(秒)
class SnapshotBuilder:
"""快照构建器
从多个数据源并行获取数据,聚合成权限快照
"""
@classmethod
async def build_auth_snapshot(
cls,
user_id: str,
group_id: str | None,
bot_id: str,
) -> AuthSnapshot:
"""构建权限快照
并行获取所有需要的数据,聚合成一个快照对象
参数:
user_id: 用户ID
group_id: 群组ID(可为None表示私聊)
bot_id: Bot ID
返回:
AuthSnapshot: 权限快照对象
"""
start_time = time.time()
try:
# 创建所有查询任务
tasks: dict[str, asyncio.Task] = {}
# 用户信息
tasks["user"] = asyncio.create_task(
cls._get_user(user_id), name="snapshot:user"
)
# Ban状态(用户+群组)
tasks["ban"] = asyncio.create_task(
cls._get_ban_status(user_id, group_id), name="snapshot:ban"
)
# 用户权限等级(全局+群组)
tasks["level"] = asyncio.create_task(
cls._get_user_levels(user_id, group_id), name="snapshot:level"
)
# 群组信息(如果有)
if group_id:
tasks["group"] = asyncio.create_task(
cls._get_group(group_id), name="snapshot:group"
)
# Bot信息
tasks["bot"] = asyncio.create_task(
cls._get_bot(bot_id), name="snapshot:bot"
)
# 并行执行所有查询
results: dict[str, Any] = {}
try:
await asyncio.wait_for(
cls._gather_results(tasks, results), timeout=BUILD_TIMEOUT
)
except asyncio.TimeoutError:
logger.warning(
f"构建权限快照超时: user={user_id}, group={group_id}",
LOG_COMMAND,
)
# 取消未完成的任务
for task in tasks.values():
if not task.done():
task.cancel()
# 聚合结果
snapshot = cls._aggregate_results(user_id, group_id, bot_id, results)
elapsed = time.time() - start_time
if elapsed > 0.5:
logger.warning(
f"构建权限快照耗时较长: {elapsed:.3f}s, "
f"user={user_id}, group={group_id}",
LOG_COMMAND,
)
return snapshot
except Exception as e:
logger.error(
f"构建权限快照失败: user={user_id}, group={group_id}",
LOG_COMMAND,
e=e,
)
# 返回一个默认快照
return AuthSnapshot(user_id=user_id, group_id=group_id, bot_id=bot_id)
@classmethod
async def _gather_results(
cls, tasks: dict[str, asyncio.Task], results: dict[str, Any]
):
"""收集所有任务结果
参数:
tasks: 任务字典
results: 结果字典(会被修改)
"""
done, _ = await asyncio.wait(tasks.values(), return_when=asyncio.ALL_COMPLETED)
for name, task in tasks.items():
if task in done:
try:
results[name] = task.result()
except Exception as e:
logger.warning(f"获取 {name} 数据失败: {e}", LOG_COMMAND)
results[name] = None
@classmethod
async def _get_user(cls, user_id: str) -> UserConsole | None:
"""获取用户信息"""
try:
return await UserConsole.get_or_none(user_id=user_id)
except Exception as e:
logger.warning(f"获取用户信息失败: {user_id}", LOG_COMMAND, e=e)
return None
@classmethod
async def _get_ban_status(
cls, user_id: str, group_id: str | None
) -> dict[str, Any]:
"""获取ban状态
返回:
dict: {
"user_banned": int, # 0/时间戳/-1
"user_ban_duration": int,
"group_banned": int
}
"""
result = {
"user_banned": 0,
"user_ban_duration": 0,
"group_banned": 0,
}
try:
# 获取所有相关的ban记录
ban_records = await BanConsole.is_ban(user_id, group_id)
for record in ban_records:
if record.user_id and not record.group_id:
# 用户级别的ban(全局)
if record.duration == -1:
result["user_banned"] = -1
result["user_ban_duration"] = -1
else:
result["user_banned"] = int(record.ban_time + record.duration)
result["user_ban_duration"] = record.duration
elif record.user_id and record.group_id:
# 用户在特定群组的ban
if record.duration == -1:
result["user_banned"] = -1
result["user_ban_duration"] = -1
else:
result["user_banned"] = int(record.ban_time + record.duration)
result["user_ban_duration"] = record.duration
elif not record.user_id and record.group_id:
# 群组级别的ban
if record.duration == -1:
result["group_banned"] = -1
else:
result["group_banned"] = int(record.ban_time + record.duration)
except Exception as e:
logger.warning(
f"获取ban状态失败: user={user_id}, group={group_id}",
LOG_COMMAND,
e=e,
)
return result
@classmethod
async def _get_user_levels(
cls, user_id: str, group_id: str | None
) -> dict[str, int]:
"""获取用户权限等级
返回:
dict: {"global": int, "group": int}
"""
result = {"global": 0, "group": 0}
try:
# 并行查询全局和群组权限
tasks = []
# 全局权限
tasks.append(LevelUser.get_or_none(user_id=user_id, group_id__isnull=True))
# 群组权限
if group_id:
tasks.append(LevelUser.get_or_none(user_id=user_id, group_id=group_id))
results = await asyncio.gather(*tasks, return_exceptions=True)
# 处理全局权限
if len(results) > 0 and isinstance(results[0], LevelUser):
result["global"] = results[0].user_level
# 处理群组权限
if len(results) > 1 and isinstance(results[1], LevelUser):
result["group"] = results[1].user_level
except Exception as e:
logger.warning(
f"获取用户权限等级失败: user={user_id}, group={group_id}",
LOG_COMMAND,
e=e,
)
return result
@classmethod
async def _get_group(cls, group_id: str) -> GroupConsole | None:
"""获取群组信息"""
try:
return await GroupConsole.get_or_none(
group_id=group_id, channel_id__isnull=True
)
except Exception as e:
logger.warning(f"获取群组信息失败: {group_id}", LOG_COMMAND, e=e)
return None
@classmethod
async def _get_bot(cls, bot_id: str) -> BotConsole | None:
"""获取Bot信息"""
try:
return await BotConsole.get_or_none(bot_id=bot_id)
except Exception as e:
logger.warning(f"获取Bot信息失败: {bot_id}", LOG_COMMAND, e=e)
return None
@classmethod
def _aggregate_results(
cls,
user_id: str,
group_id: str | None,
bot_id: str,
results: dict[str, Any],
) -> AuthSnapshot:
"""聚合查询结果为快照
参数:
user_id: 用户ID
group_id: 群组ID
bot_id: Bot ID
results: 查询结果字典
返回:
AuthSnapshot: 权限快照
"""
snapshot = AuthSnapshot(
user_id=user_id,
group_id=group_id,
bot_id=bot_id,
)
# 用户信息
if user := results.get("user"):
snapshot.user_gold = user.gold
# Ban状态
if ban_status := results.get("ban"):
snapshot.user_banned = ban_status.get("user_banned", 0)
snapshot.user_ban_duration = ban_status.get("user_ban_duration", 0)
snapshot.group_banned = ban_status.get("group_banned", 0)
# 用户权限等级
if levels := results.get("level"):
snapshot.user_level_global = levels.get("global", 0)
snapshot.user_level_group = levels.get("group", 0)
# 群组信息
if group := results.get("group"):
snapshot.group_exists = True
snapshot.group_status = group.status
snapshot.group_level = group.level
snapshot.group_is_super = group.is_super
snapshot.group_block_plugins = group.block_plugin or ""
snapshot.group_superuser_block_plugins = group.superuser_block_plugin or ""
elif group_id:
# 有 group_id 但没有群组数据,可能是新群
snapshot.group_exists = False
# Bot信息
if bot := results.get("bot"):
snapshot.bot_status = bot.status
# BotConsole 的 block_plugins 是一个列表
if hasattr(bot, "block_plugins") and bot.block_plugins:
if isinstance(bot.block_plugins, list):
snapshot.bot_block_plugins = "".join(
f"<{p}," for p in bot.block_plugins
)
else:
snapshot.bot_block_plugins = bot.block_plugins
return snapshot
@classmethod
async def build_plugin_snapshot(cls, module: str) -> PluginSnapshot | None:
"""构建插件快照
参数:
module: 插件模块名
返回:
PluginSnapshot | None: 插件快照,不存在时返回None
"""
try:
plugin = await PluginInfo.get_or_none(module=module)
if not plugin:
return None
return PluginSnapshot(
module=plugin.module,
name=plugin.name,
status=plugin.status,
block_type=plugin.block_type,
plugin_type=plugin.plugin_type,
admin_level=plugin.admin_level or 0,
cost_gold=plugin.cost_gold,
level=plugin.level,
limit_superuser=plugin.limit_superuser,
ignore_prompt=plugin.ignore_prompt,
)
except Exception as e:
logger.error(f"构建插件快照失败: {module}", LOG_COMMAND, e=e)
return None
+389
View File
@@ -0,0 +1,389 @@
"""
优化后的权限检查器
使用预聚合的权限快照进行权限检查,将查询次数从6-10次降低到1-2次
"""
import asyncio
import time
from nonebot.adapters import Bot, Event
from nonebot.exception import IgnoredException
from nonebot.matcher import Matcher
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo
from zhenxun.models.user_console import UserConsole
from zhenxun.services.log import logger
from zhenxun.utils.enum import BlockType, GoldHandle
from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.utils import get_entity_ids
from .models import AuthSnapshot, PluginSnapshot
from .service import AuthSnapshotService, PluginSnapshotService
LOG_COMMAND = "auth_checker_v2"
WARNING_THRESHOLD = 0.5 # 警告阈值(秒)
class AuthCheckResult:
"""权限检查结果"""
def __init__(self):
self.passed: bool = True
self.skip_reason: str = ""
self.cost_gold: int = 0
self.is_superuser: bool = False
def fail(self, reason: str):
"""标记检查失败"""
self.passed = False
self.skip_reason = reason
class OptimizedAuthChecker:
"""优化后的权限检查器
核心优化:
1. 使用预聚合的权限快照,将多次查询合并为1-2次
2. 所有检查基于内存中的快照数据,无额外I/O
3. 保持与原有系统相同的检查逻辑和结果
"""
async def check(
self,
matcher: Matcher,
event: Event,
bot: Bot,
session: Uninfo,
message: UniMsg,
):
"""执行权限检查
参数:
matcher: Matcher
event: Event
bot: Bot
session: Uninfo
message: UniMsg
"""
start_time = time.time()
result = AuthCheckResult()
hook_times: dict[str, str] = {}
try:
# 1. 获取基础信息
entity = get_entity_ids(session)
module = matcher.plugin_name or ""
if not module:
result.fail("Matcher插件名称不存在...")
raise IgnoredException(result.skip_reason)
# 2. 获取权限快照(第一次查询)
snapshot_start = time.time()
auth_snapshot = await AuthSnapshotService.get_snapshot(
user_id=entity.user_id,
group_id=entity.group_id,
bot_id=bot.self_id,
)
hook_times["get_auth_snapshot"] = f"{time.time() - snapshot_start:.3f}s"
# 3. 获取插件快照(第二次查询,通常命中内存缓存)
plugin_start = time.time()
plugin_snapshot = await PluginSnapshotService.get_plugin(module)
hook_times["get_plugin_snapshot"] = f"{time.time() - plugin_start:.3f}s"
if not plugin_snapshot:
result.fail(f"插件:{module} 数据不存在...")
raise IgnoredException(result.skip_reason)
# 4. 检查是否为隐藏插件
if plugin_snapshot.is_hidden():
result.fail(f"插件: {plugin_snapshot.name}:{module} 为HIDDEN...")
raise IgnoredException(result.skip_reason)
# 5. 检查超级用户
is_superuser = session.user.id in bot.config.superusers
result.is_superuser = is_superuser
# 6. 执行所有权限检查(纯内存计算)
check_start = time.time()
await self._run_all_checks(
result=result,
auth_snapshot=auth_snapshot,
plugin_snapshot=plugin_snapshot,
message=message,
session=session,
is_superuser=is_superuser,
)
hook_times["run_checks"] = f"{time.time() - check_start:.3f}s"
# 7. 处理检查结果
if not result.passed:
logger.info(result.skip_reason, LOG_COMMAND, session=session)
raise IgnoredException(result.skip_reason)
# 8. 超级用户跳过后续限制
if is_superuser:
if plugin_snapshot.is_superuser_plugin():
logger.debug(
"超级用户访问超级用户插件,跳过权限检测...",
LOG_COMMAND,
session=session,
)
return
if not plugin_snapshot.limit_superuser:
logger.debug(
"超级用户跳过权限检测...", LOG_COMMAND, session=session
)
return
# 9. 扣除金币(如果需要)
if result.cost_gold > 0:
try:
gold_start = time.time()
await asyncio.wait_for(
UserConsole.reduce_gold(
entity.user_id,
result.cost_gold,
GoldHandle.PLUGIN,
module,
PlatformUtils.get_platform(session),
),
timeout=5.0,
)
hook_times["reduce_gold"] = f"{time.time() - gold_start:.3f}s"
# 扣除金币后失效用户快照缓存
await AuthSnapshotService.invalidate_user(entity.user_id)
except asyncio.TimeoutError:
logger.error(
f"扣除金币超时,模块: {module}", LOG_COMMAND, session=session
)
except IgnoredException:
raise
except Exception as e:
logger.error(f"权限检查异常: {e}", LOG_COMMAND, session=session, e=e)
raise IgnoredException("权限检查异常") from e
finally:
# 记录总执行时间
total_time = time.time() - start_time
if total_time > WARNING_THRESHOLD:
logger.warning(
f"权限检查耗时过长: {total_time:.3f}s, "
f"模块: {matcher.plugin_name}, 详情: {hook_times}",
LOG_COMMAND,
session=session,
)
async def _run_all_checks(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
message: UniMsg,
session: Uninfo,
is_superuser: bool,
):
"""执行所有权限检查
所有检查都基于内存中的快照数据,无I/O操作
"""
# 1. Ban检查(关键优先级)
self._check_ban(result, auth_snapshot, plugin_snapshot, is_superuser)
if not result.passed:
return
# 2. Bot状态检查(关键优先级)
self._check_bot_status(result, auth_snapshot, plugin_snapshot)
if not result.passed:
return
# 3. 插件全局状态检查(高优先级)
self._check_plugin_global_status(result, auth_snapshot, plugin_snapshot)
if not result.passed:
return
# 4. 群组状态检查(高优先级)
if auth_snapshot.group_id:
self._check_group_status(result, auth_snapshot, plugin_snapshot, message)
if not result.passed:
return
else:
# 私聊检查
self._check_private_status(result, plugin_snapshot)
if not result.passed:
return
# 5. 管理员权限检查(中优先级)
self._check_admin_level(result, auth_snapshot, plugin_snapshot)
if not result.passed:
return
# 6. 金币检查(低优先级)
self._check_gold(result, auth_snapshot, plugin_snapshot)
def _check_ban(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
is_superuser: bool,
):
"""检查ban状态"""
# 超级用户不受ban限制
if is_superuser:
return
# 检查群组ban
if auth_snapshot.is_group_banned():
result.fail(f"群组: {auth_snapshot.group_id} 处于黑名单中...")
return
# 检查用户ban
if auth_snapshot.is_user_banned():
remaining = auth_snapshot.get_user_ban_remaining()
if remaining == -1:
result.fail("用户处于永久黑名单中...")
else:
result.fail(f"用户处于黑名单中,剩余 {remaining} 秒...")
def _check_bot_status(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
):
"""检查Bot状态"""
if not auth_snapshot.bot_status:
result.fail("Bot不存在或休眠中阻断权限检测...")
return
if auth_snapshot.is_plugin_blocked_by_bot(plugin_snapshot.module):
result.fail(
f"Bot插件 {plugin_snapshot.name}({plugin_snapshot.module}) "
"权限检查结果为关闭..."
)
def _check_plugin_global_status(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
):
"""检查插件全局状态"""
# 全局禁用检查
if not plugin_snapshot.status and plugin_snapshot.block_type == BlockType.ALL:
# 超级群组可以使用全局关闭的功能
if auth_snapshot.group_is_super:
return
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) 全局未开启此功能..."
)
def _check_group_status(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
message: UniMsg,
):
"""检查群组状态"""
# 群组不存在
if not auth_snapshot.group_exists:
result.fail("群组信息不存在...")
return
# 群组黑名单
if auth_snapshot.group_level < 0:
result.fail("群组黑名单, 目标群组群权限权限-1...")
return
# 群组休眠状态(除非是开启命令)
text = message.extract_plain_text().strip()
if text != "开启" and not auth_snapshot.group_status:
result.fail("群组休眠状态...")
return
# 插件等级检查
if plugin_snapshot.level > auth_snapshot.group_level:
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) 群等级限制,"
f"该功能需要的群等级: {plugin_snapshot.level}..."
)
return
# 超级用户禁用检查
if auth_snapshot.is_plugin_blocked_by_superuser(plugin_snapshot.module):
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
"超级管理员禁用了该群此功能..."
)
return
# 普通禁用检查
if auth_snapshot.is_plugin_blocked_by_group(plugin_snapshot.module):
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) 未开启此功能..."
)
return
# 群组禁用类型检查
if plugin_snapshot.block_type == BlockType.GROUP:
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
"该插件在群组中已被禁用..."
)
def _check_private_status(
self,
result: AuthCheckResult,
plugin_snapshot: PluginSnapshot,
):
"""检查私聊状态"""
if plugin_snapshot.block_type == BlockType.PRIVATE:
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
"该插件在私聊中已被禁用..."
)
def _check_admin_level(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
):
"""检查管理员权限"""
if not plugin_snapshot.admin_level:
return
user_level = auth_snapshot.get_user_level()
if user_level < plugin_snapshot.admin_level:
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
f"管理员权限不足,需要等级: {plugin_snapshot.admin_level}..."
)
def _check_gold(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
):
"""检查金币"""
if plugin_snapshot.cost_gold <= 0:
return
if auth_snapshot.user_gold < plugin_snapshot.cost_gold:
result.fail(f"金币不足..该功能需要{plugin_snapshot.cost_gold}金币..")
return
# 记录需要扣除的金币
result.cost_gold = plugin_snapshot.cost_gold
# 全局实例
optimized_auth_checker = OptimizedAuthChecker()
+263
View File
@@ -0,0 +1,263 @@
"""
权限快照数据模型
定义 AuthSnapshot 和 PluginSnapshot 的数据结构
"""
import time
from typing import ClassVar
from pydantic import BaseModel, Field
from zhenxun.utils.enum import BlockType, PluginType
class AuthSnapshot(BaseModel):
"""权限快照数据模型
聚合了权限检查所需的所有用户、群组、Bot相关数据
"""
# 快照标识
user_id: str
group_id: str | None = None
bot_id: str
# === 用户信息 ===
user_gold: int = 100
"""用户金币"""
user_banned: int = 0
"""0=未ban, -1=永久ban, >0=ban结束时间戳"""
user_ban_duration: int = 0
"""ban时长(秒),-1为永久"""
# === 用户权限等级 ===
user_level_global: int = 0
"""全局权限等级"""
user_level_group: int = 0
"""群组内权限等级"""
# === 群组信息 ===
group_exists: bool = False
"""群组是否存在(用于区分私聊和未知群组)"""
group_status: bool = True
"""群组状态 (True=开启, False=休眠)"""
group_level: int = 5
"""群组等级"""
group_is_super: bool = False
"""是否超级群组(可以使用全局关闭的功能)"""
group_block_plugins: str = ""
"""禁用插件列表,格式: "<plugin1,<plugin2," """
group_superuser_block_plugins: str = ""
"""超级用户禁用插件列表"""
# === 群组ban状态 ===
group_banned: int = 0
"""0=未ban, -1=永久ban, >0=ban结束时间戳"""
# === Bot信息 ===
bot_status: bool = True
"""Bot状态"""
bot_block_plugins: str = ""
"""Bot禁用插件列表,格式: "<plugin1,<plugin2," """
# === 元数据 ===
version: int = 1
"""快照版本"""
created_at: float = Field(default_factory=time.time)
"""创建时间戳"""
# === 类变量 ===
DEFAULT_TTL: ClassVar[int] = 60
"""默认过期时间(秒)"""
def is_expired(self, ttl: int | None = None) -> bool:
"""检查快照是否过期
参数:
ttl: 过期时间(秒),为None时使用默认值
返回:
bool: 是否过期
"""
expire_ttl = ttl if ttl is not None else self.DEFAULT_TTL
return time.time() - self.created_at > expire_ttl
def is_user_banned(self) -> bool:
"""检查用户是否被ban
返回:
bool: 用户是否被ban
"""
if self.user_banned == 0:
return False
if self.user_banned == -1:
return True
# 检查ban是否过期
return time.time() < self.user_banned
def is_group_banned(self) -> bool:
"""检查群组是否被ban
返回:
bool: 群组是否被ban
"""
if self.group_banned == 0:
return False
if self.group_banned == -1:
return True
return time.time() < self.group_banned
def get_user_ban_remaining(self) -> int:
"""获取用户ban剩余时间
返回:
int: 剩余时间(秒),-1表示永久,0表示未被ban
"""
if self.user_banned == 0:
return 0
if self.user_banned == -1:
return -1
remaining = int(self.user_banned - time.time())
return max(remaining, 0)
def get_user_level(self) -> int:
"""获取用户有效权限等级(取全局和群组的最大值)
返回:
int: 用户权限等级
"""
return max(self.user_level_global, self.user_level_group)
def is_plugin_blocked_by_group(self, module: str) -> bool:
"""检查插件是否被群组禁用
参数:
module: 插件模块名
返回:
bool: 是否被禁用
"""
marker = f"<{module},"
return marker in self.group_block_plugins
def is_plugin_blocked_by_superuser(self, module: str) -> bool:
"""检查插件是否被超级用户禁用
参数:
module: 插件模块名
返回:
bool: 是否被禁用
"""
marker = f"<{module},"
return marker in self.group_superuser_block_plugins
def is_plugin_blocked_by_bot(self, module: str) -> bool:
"""检查插件是否被Bot禁用
参数:
module: 插件模块名
返回:
bool: 是否被禁用
"""
marker = f"<{module},"
return marker in self.bot_block_plugins
class PluginSnapshot(BaseModel):
"""插件快照数据模型
包含插件权限检查所需的所有配置信息
"""
# 插件标识
module: str
"""模块名"""
name: str = ""
"""插件名称"""
# === 插件状态 ===
status: bool = True
"""全局开关状态"""
block_type: BlockType | None = None
"""禁用类型 (PRIVATE/GROUP/ALL/None)"""
plugin_type: PluginType | None = None
"""插件类型"""
# === 权限要求 ===
admin_level: int = 0
"""调用所需权限等级"""
cost_gold: int = 0
"""调用所需金币"""
level: int = 5
"""所需群权限等级"""
limit_superuser: bool = False
"""是否限制超级用户"""
# === 显示配置 ===
ignore_prompt: bool = False
"""是否忽略阻断提示"""
# === 元数据 ===
created_at: float = Field(default_factory=time.time)
"""创建时间戳"""
# === 类变量 ===
DEFAULT_TTL: ClassVar[int] = 300
"""默认过期时间(秒)"""
MEMORY_TTL: ClassVar[int] = 30
"""本地内存缓存过期时间(秒)"""
def is_expired(self, ttl: int | None = None) -> bool:
"""检查快照是否过期
参数:
ttl: 过期时间(秒),为None时使用默认值
返回:
bool: 是否过期
"""
expire_ttl = ttl if ttl is not None else self.DEFAULT_TTL
return time.time() - self.created_at > expire_ttl
def is_hidden(self) -> bool:
"""检查是否为隐藏插件
返回:
bool: 是否隐藏
"""
return self.plugin_type == PluginType.HIDDEN
def is_superuser_plugin(self) -> bool:
"""检查是否为超级用户插件
返回:
bool: 是否为超级用户插件
"""
return self.plugin_type == PluginType.SUPERUSER
def is_globally_disabled(self) -> bool:
"""检查是否全局禁用
返回:
bool: 是否全局禁用
"""
return not self.status and self.block_type == BlockType.ALL
def is_disabled_in_group(self) -> bool:
"""检查是否在群组中禁用
返回:
bool: 是否在群组中禁用
"""
return self.block_type == BlockType.GROUP
def is_disabled_in_private(self) -> bool:
"""检查是否在私聊中禁用
返回:
bool: 是否在私聊中禁用
"""
return self.block_type == BlockType.PRIVATE
+409
View File
@@ -0,0 +1,409 @@
"""
快照服务
提供权限快照的获取、缓存、失效等功能
"""
import asyncio
import time
from typing import ClassVar
from zhenxun.services.cache import CacheRoot, cache_config
from zhenxun.services.cache.config import CacheMode
from zhenxun.services.log import logger
from zhenxun.utils.enum import CacheType
from .builder import SnapshotBuilder
from .models import AuthSnapshot, PluginSnapshot
LOG_COMMAND = "auth_snapshot"
# 缓存键前缀
AUTH_SNAPSHOT_PREFIX = "AUTH_SNAPSHOT"
PLUGIN_SNAPSHOT_PREFIX = "PLUGIN_SNAPSHOT"
class AuthSnapshotService:
"""权限快照服务
提供权限快照的获取、缓存和失效管理
"""
# 本地内存缓存(用于热点数据)
_memory_cache: ClassVar[dict[str, tuple[float, AuthSnapshot]]] = {}
_memory_cache_ttl: ClassVar[int] = 10 # 内存缓存TTL(秒)
_cache_ttl: ClassVar[int] = 60 # Redis缓存TTL(秒)
# 正在构建中的快照(防止并发重复构建)
_building: ClassVar[dict[str, asyncio.Future]] = {}
@classmethod
def _build_cache_key(cls, user_id: str, group_id: str | None, bot_id: str) -> str:
"""构建缓存键"""
group_part = group_id or "PRIVATE"
return f"{AUTH_SNAPSHOT_PREFIX}:{user_id}:{group_part}:{bot_id}"
@classmethod
async def get_snapshot(
cls,
user_id: str,
group_id: str | None,
bot_id: str,
force_refresh: bool = False,
) -> AuthSnapshot:
"""获取权限快照
优先从缓存获取,缓存未命中时构建新快照
参数:
user_id: 用户ID
group_id: 群组ID(可为None)
bot_id: Bot ID
force_refresh: 是否强制刷新
返回:
AuthSnapshot: 权限快照
"""
cache_key = cls._build_cache_key(user_id, group_id, bot_id)
# 1. 尝试从内存缓存获取
if not force_refresh:
if snapshot := cls._get_from_memory(cache_key):
return snapshot
# 2. 尝试从Redis获取
if not force_refresh and cache_config.cache_mode != CacheMode.NONE:
try:
cached = await CacheRoot.get(CacheType.TEMP, cache_key)
if cached and isinstance(cached, dict):
snapshot = AuthSnapshot.model_validate(cached)
if not snapshot.is_expired(cls._cache_ttl):
# 更新内存缓存
cls._set_to_memory(cache_key, snapshot)
return snapshot
except Exception as e:
logger.debug(f"从Redis获取权限快照失败: {cache_key}", LOG_COMMAND, e=e)
# 3. 检查是否正在构建中(防止并发)
if cache_key in cls._building:
try:
return await cls._building[cache_key]
except Exception:
pass
# 4. 构建新快照
return await cls._build_and_cache(user_id, group_id, bot_id, cache_key)
@classmethod
def _get_from_memory(cls, cache_key: str) -> AuthSnapshot | None:
"""从内存缓存获取"""
if cache_key in cls._memory_cache:
created_at, snapshot = cls._memory_cache[cache_key]
if time.time() - created_at < cls._memory_cache_ttl:
return snapshot
# 过期,删除
del cls._memory_cache[cache_key]
return None
@classmethod
def _set_to_memory(cls, cache_key: str, snapshot: AuthSnapshot):
"""设置内存缓存"""
cls._memory_cache[cache_key] = (time.time(), snapshot)
# 清理过期的内存缓存(简单策略:超过1000条时清理)
if len(cls._memory_cache) > 1000:
cls._cleanup_memory_cache()
@classmethod
def _cleanup_memory_cache(cls):
"""清理过期的内存缓存"""
now = time.time()
expired_keys = [
k
for k, (created_at, _) in cls._memory_cache.items()
if now - created_at > cls._memory_cache_ttl
]
for key in expired_keys:
del cls._memory_cache[key]
@classmethod
async def _build_and_cache(
cls,
user_id: str,
group_id: str | None,
bot_id: str,
cache_key: str,
) -> AuthSnapshot:
"""构建并缓存快照"""
loop = asyncio.get_running_loop()
future: asyncio.Future[AuthSnapshot] = loop.create_future()
cls._building[cache_key] = future
try:
# 构建快照
snapshot = await SnapshotBuilder.build_auth_snapshot(
user_id, group_id, bot_id
)
# 存入Redis缓存(异步,不阻塞)
if cache_config.cache_mode != CacheMode.NONE:
asyncio.create_task( # noqa: RUF006
cls._cache_to_redis(cache_key, snapshot)
)
# 存入内存缓存
cls._set_to_memory(cache_key, snapshot)
future.set_result(snapshot)
return snapshot
except Exception as e:
future.set_exception(e)
raise
finally:
cls._building.pop(cache_key, None)
@classmethod
async def _cache_to_redis(cls, cache_key: str, snapshot: AuthSnapshot):
"""异步存入Redis"""
try:
await CacheRoot.set(
CacheType.TEMP,
cache_key,
snapshot.model_dump(),
expire=cls._cache_ttl,
)
except Exception as e:
logger.debug(f"缓存权限快照到Redis失败: {cache_key}", LOG_COMMAND, e=e)
@classmethod
async def invalidate_user(cls, user_id: str):
"""失效用户相关的所有快照
参数:
user_id: 用户ID
"""
# 清理内存缓存
keys_to_delete = [k for k in cls._memory_cache if f":{user_id}:" in k]
for key in keys_to_delete:
del cls._memory_cache[key]
logger.debug(f"已失效用户 {user_id} 的权限快照缓存", LOG_COMMAND)
@classmethod
async def invalidate_group(cls, group_id: str):
"""失效群组相关的所有快照
参数:
group_id: 群组ID
"""
# 清理内存缓存
keys_to_delete = [k for k in cls._memory_cache if f":{group_id}:" in k]
for key in keys_to_delete:
del cls._memory_cache[key]
logger.debug(f"已失效群组 {group_id} 的权限快照缓存", LOG_COMMAND)
@classmethod
async def invalidate_bot(cls, bot_id: str):
"""失效Bot相关的所有快照
参数:
bot_id: Bot ID
"""
# 清理内存缓存
keys_to_delete = [k for k in cls._memory_cache if k.endswith(f":{bot_id}")]
for key in keys_to_delete:
del cls._memory_cache[key]
logger.debug(f"已失效Bot {bot_id} 的权限快照缓存", LOG_COMMAND)
@classmethod
def clear_all_cache(cls):
"""清空所有缓存"""
cls._memory_cache.clear()
logger.info("已清空所有权限快照缓存", LOG_COMMAND)
class PluginSnapshotService:
"""插件快照服务
提供插件信息的获取和缓存,支持本地内存缓存 + Redis 双层缓存
"""
# 本地内存缓存
_memory_cache: ClassVar[dict[str, tuple[float, PluginSnapshot]]] = {}
_memory_cache_ttl: ClassVar[int] = 30 # 内存缓存TTL(秒)
_cache_ttl: ClassVar[int] = 300 # Redis缓存TTL(秒)
# 正在构建中的快照
_building: ClassVar[dict[str, asyncio.Future]] = {}
@classmethod
def _build_cache_key(cls, module: str) -> str:
"""构建缓存键"""
return f"{PLUGIN_SNAPSHOT_PREFIX}:{module}"
@classmethod
async def get_plugin(
cls, module: str, force_refresh: bool = False
) -> PluginSnapshot | None:
"""获取插件快照
参数:
module: 插件模块名
force_refresh: 是否强制刷新
返回:
PluginSnapshot | None: 插件快照,不存在时返回None
"""
cache_key = cls._build_cache_key(module)
# 1. 尝试从内存缓存获取(最快)
if not force_refresh:
if snapshot := cls._get_from_memory(cache_key):
return snapshot
# 2. 尝试从Redis获取
if not force_refresh and cache_config.cache_mode != CacheMode.NONE:
try:
cached = await CacheRoot.get(CacheType.PLUGINS, cache_key)
if cached and isinstance(cached, dict):
snapshot = PluginSnapshot.model_validate(cached)
if not snapshot.is_expired(cls._cache_ttl):
cls._set_to_memory(cache_key, snapshot)
return snapshot
except Exception as e:
logger.debug(f"从Redis获取插件快照失败: {module}", LOG_COMMAND, e=e)
# 3. 检查是否正在构建中
if cache_key in cls._building:
try:
return await cls._building[cache_key]
except Exception:
pass
# 4. 从数据库构建
return await cls._build_and_cache(module, cache_key)
@classmethod
def _get_from_memory(cls, cache_key: str) -> PluginSnapshot | None:
"""从内存缓存获取"""
if cache_key in cls._memory_cache:
created_at, snapshot = cls._memory_cache[cache_key]
if time.time() - created_at < cls._memory_cache_ttl:
return snapshot
del cls._memory_cache[cache_key]
return None
@classmethod
def _set_to_memory(cls, cache_key: str, snapshot: PluginSnapshot):
"""设置内存缓存"""
cls._memory_cache[cache_key] = (time.time(), snapshot)
@classmethod
async def _build_and_cache(
cls, module: str, cache_key: str
) -> PluginSnapshot | None:
"""构建并缓存插件快照"""
loop = asyncio.get_running_loop()
future: asyncio.Future[PluginSnapshot | None] = loop.create_future()
cls._building[cache_key] = future
try:
snapshot = await SnapshotBuilder.build_plugin_snapshot(module)
if snapshot:
# 存入Redis缓存
if cache_config.cache_mode != CacheMode.NONE:
asyncio.create_task( # noqa: RUF006
cls._cache_to_redis(cache_key, snapshot)
)
# 存入内存缓存
cls._set_to_memory(cache_key, snapshot)
future.set_result(snapshot)
return snapshot
except Exception as e:
future.set_exception(e)
raise
finally:
cls._building.pop(cache_key, None)
@classmethod
async def _cache_to_redis(cls, cache_key: str, snapshot: PluginSnapshot):
"""异步存入Redis"""
try:
await CacheRoot.set(
CacheType.PLUGINS,
cache_key,
snapshot.model_dump(),
expire=cls._cache_ttl,
)
except Exception as e:
logger.debug(f"缓存插件快照到Redis失败: {cache_key}", LOG_COMMAND, e=e)
@classmethod
async def invalidate_plugin(cls, module: str):
"""失效指定插件的缓存
参数:
module: 插件模块名
"""
cache_key = cls._build_cache_key(module)
# 清理内存缓存
if cache_key in cls._memory_cache:
del cls._memory_cache[cache_key]
# 清理Redis缓存
if cache_config.cache_mode != CacheMode.NONE:
try:
await CacheRoot.delete(CacheType.PLUGINS, cache_key)
except Exception:
pass
logger.debug(f"已失效插件 {module} 的快照缓存", LOG_COMMAND)
@classmethod
async def warmup(cls):
"""预热所有插件缓存
在启动时调用,预加载所有插件信息到缓存
"""
from zhenxun.models.plugin_info import PluginInfo
try:
plugins = await PluginInfo.filter(load_status=True).all()
count = 0
for plugin in plugins:
snapshot = PluginSnapshot(
module=plugin.module,
name=plugin.name,
status=plugin.status,
block_type=plugin.block_type,
plugin_type=plugin.plugin_type,
admin_level=plugin.admin_level or 0,
cost_gold=plugin.cost_gold,
level=plugin.level,
limit_superuser=plugin.limit_superuser,
ignore_prompt=plugin.ignore_prompt,
)
cache_key = cls._build_cache_key(plugin.module)
cls._set_to_memory(cache_key, snapshot)
count += 1
logger.info(f"已预热 {count} 个插件的快照缓存", LOG_COMMAND)
except Exception as e:
logger.error("预热插件缓存失败", LOG_COMMAND, e=e)
@classmethod
def clear_all_cache(cls):
"""清空所有缓存"""
cls._memory_cache.clear()
logger.info("已清空所有插件快照缓存", LOG_COMMAND)