mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-02 18:20:01 +08:00
363 lines
12 KiB
Python
363 lines
12 KiB
Python
"""
|
|
快照构建器
|
|
|
|
负责从多个数据源聚合数据构建权限快照
|
|
"""
|
|
|
|
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
|