mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-10 14:20:04 +08:00
feat: add permission snapshot system to optimize auth checks
This commit is contained in:
@@ -0,0 +1,15 @@
|
||||
"""
|
||||
权限快照服务模块
|
||||
|
||||
提供预聚合的权限检查数据,将多次数据库/缓存查询优化为1-2次
|
||||
"""
|
||||
|
||||
from .models import AuthSnapshot, PluginSnapshot
|
||||
from .service import AuthSnapshotService, PluginSnapshotService
|
||||
|
||||
__all__ = [
|
||||
"AuthSnapshot",
|
||||
"AuthSnapshotService",
|
||||
"PluginSnapshot",
|
||||
"PluginSnapshotService",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user