feat: add TEMP cache type and enhance user and bot management with improved error handling and caching mechanisms

This commit is contained in:
HibiKier
2025-12-22 10:16:02 +08:00
parent 632dff3bad
commit a3cbfefaa1
6 changed files with 108 additions and 75 deletions
+1 -1
View File
@@ -33,7 +33,7 @@ def register_cache_types():
CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}" CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}"
) )
CacheRegistry.register(CacheType.BAN, BanConsole, key_format="{user_id}_{group_id}") CacheRegistry.register(CacheType.BAN, BanConsole, key_format="{user_id}_{group_id}")
CacheRegistry.register(CacheType.TEMP, None, 3600)
if cache_config.cache_mode == CacheMode.NONE: if cache_config.cache_mode == CacheMode.NONE:
logger.info("缓存功能已禁用,将直接从数据库获取数据") logger.info("缓存功能已禁用,将直接从数据库获取数据")
else: else:
@@ -3,6 +3,7 @@ from typing import cast
import nonebot import nonebot
from nonebot.adapters import Bot from nonebot.adapters import Bot
from nonebot.plugin import PluginMetadata from nonebot.plugin import PluginMetadata
from tortoise.exceptions import IntegrityError
from zhenxun.configs.utils import PluginExtraData from zhenxun.configs.utils import PluginExtraData
from zhenxun.models.bot_console import BotConsole from zhenxun.models.bot_console import BotConsole
@@ -72,9 +73,17 @@ async def init_bot_console(bot: Bot):
list[str], await TaskInfo.filter(status=True).values_list("module", flat=True) list[str], await TaskInfo.filter(status=True).values_list("module", flat=True)
) )
platform = PlatformUtils.get_platform(bot) platform = PlatformUtils.get_platform(bot)
bot_data, created = await BotConsole.get_or_create(
bot_id=bot.self_id, platform=platform try:
) bot_data = await BotConsole.create(
bot_id=bot.self_id,
platform=platform,
)
created = True
except IntegrityError:
bot_data = await BotConsole.get(bot_id=bot.self_id)
created = False
if not created: if not created:
task_list = await _filter_blocked_items( task_list = await _filter_blocked_items(
+41 -39
View File
@@ -1,3 +1,4 @@
import asyncio
import time import time
from typing import ClassVar from typing import ClassVar
from typing_extensions import Self from typing_extensions import Self
@@ -28,6 +29,7 @@ class BanConsole(Model):
"""ban时长""" """ban时长"""
operator = fields.CharField(255) operator = fields.CharField(255)
"""使用Ban命令的用户""" """使用Ban命令的用户"""
_inflight: ClassVar[dict[tuple[str | None, str | None], asyncio.Future]] = {}
class Meta: # pyright: ignore [reportIncompatibleVariableOverride] class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "ban_console" table = "ban_console"
@@ -44,29 +46,38 @@ class BanConsole(Model):
@classmethod @classmethod
async def _get_data(cls, user_id: str | None, group_id: str | None) -> Self | None: async def _get_data(cls, user_id: str | None, group_id: str | None) -> Self | None:
"""获取数据
参数:
user_id: 用户id
group_id: 群组id
异常:
UserAndGroupIsNone: 用户id和群组id都为空
返回:
Self | None: Self
"""
if not user_id and not group_id: if not user_id and not group_id:
raise UserAndGroupIsNone() raise UserAndGroupIsNone()
dao = DataAccess(cls)
if user_id: key = (user_id, group_id)
return ( future = cls._inflight.get(key)
await dao.safe_get_or_none(user_id=user_id, group_id=group_id) if future:
if group_id return await future
else await dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
) loop = asyncio.get_running_loop()
else: future = loop.create_future()
return await dao.safe_get_or_none(user_id="", group_id=group_id) cls._inflight[key] = future
try:
dao = DataAccess(cls)
if user_id:
result = (
await dao.safe_get_or_none(user_id=user_id, group_id=group_id)
if group_id
else await dao.safe_get_or_none(
user_id=user_id, group_id__isnull=True
)
)
else:
result = await dao.safe_get_or_none(user_id="", group_id=group_id)
future.set_result(result)
return result
except Exception as e:
future.set_exception(e)
raise
finally:
cls._inflight.pop(key, None)
@classmethod @classmethod
async def check_ban_level( async def check_ban_level(
@@ -143,30 +154,21 @@ class BanConsole(Model):
duration: int, duration: int,
operator: str | None = None, operator: str | None = None,
): ):
"""ban掉目标用户
参数:
user_id: 用户id
group_id: 群组id
ban_level: 使用命令者的权限等级
duration: 时长,分钟,-1时为永久
operator: 操作者id
"""
logger.debug( logger.debug(
f"封禁用户/群组,等级:{ban_level},时长: {duration}", f"封禁用户/群组,等级:{ban_level},时长: {duration}",
target=f"{group_id}:{user_id}", target=f"{group_id}:{user_id}",
) )
target = await cls._get_data(user_id, group_id)
if target: await cls.update_or_create(
await cls.unban(user_id, group_id)
await cls.create(
user_id=user_id, user_id=user_id,
group_id=group_id, group_id=group_id,
ban_level=ban_level, defaults={
ban_time=int(time.time()), "ban_level": ban_level,
ban_reason=reason, "ban_time": int(time.time()),
duration=duration, "ban_reason": reason,
operator=operator or 0, "duration": duration,
"operator": operator or 0,
},
) )
@classmethod @classmethod
+51 -31
View File
@@ -1,6 +1,8 @@
from tortoise import fields from tortoise import fields
from tortoise.exceptions import IntegrityError
from zhenxun.models.goods_info import GoodsInfo from zhenxun.models.goods_info import GoodsInfo
from zhenxun.services.cache import CacheRoot
from zhenxun.services.db_context import Model from zhenxun.services.db_context import Model
from zhenxun.utils.enum import CacheType, GoldHandle from zhenxun.utils.enum import CacheType, GoldHandle
from zhenxun.utils.exception import GoodsNotFound, InsufficientGold from zhenxun.utils.exception import GoodsNotFound, InsufficientGold
@@ -38,24 +40,23 @@ class UserConsole(Model):
@classmethod @classmethod
async def get_user(cls, user_id: str, platform: str | None = None) -> "UserConsole": async def get_user(cls, user_id: str, platform: str | None = None) -> "UserConsole":
"""获取用户 if user := await cls.get_or_none(user_id=user_id):
return user
参数: try:
user_id: 用户id user, created = await cls.get_or_create(
platform: 平台. user_id=user_id,
defaults={
返回: "platform": platform,
UserConsole: UserConsole "uid": 0,
""" },
if not await cls.exists(user_id=user_id):
await cls.create(
user_id=user_id, platform=platform, uid=await cls.get_new_uid()
) )
# user, _ = await UserConsole.get_or_create( if created:
# user_id=user_id, user.uid = await cls.get_new_uid()
# defaults={"platform": platform, "uid": await cls.get_new_uid()}, await user.save(update_fields=["uid"])
# ) return user
return await cls.get(user_id=user_id) except IntegrityError:
return await cls.get(user_id=user_id)
@classmethod @classmethod
async def get_new_uid(cls) -> int: async def get_new_uid(cls) -> int:
@@ -64,9 +65,18 @@ class UserConsole(Model):
返回: 返回:
int: 最新uid int: 最新uid
""" """
if user := await cls.annotate().order_by("-uid").first(): uid: int | None = await CacheRoot.get(CacheType.TEMP, "USER_CONSOLE_UID")
return user.uid + 1 if uid is None:
return 1 data: list[int] = ( # pyright: ignore[reportAssignmentType]
await cls.annotate()
.order_by("-uid")
.limit(1)
.values_list("uid", flat=True)
)
uid = data[0] if data else 0
uid = uid + 1
await CacheRoot.set(CacheType.TEMP, "USER_CONSOLE_UID", uid)
return uid
@classmethod @classmethod
async def add_gold( async def add_gold(
@@ -80,12 +90,14 @@ class UserConsole(Model):
source: 来源 source: 来源
platform: 平台. platform: 平台.
""" """
user, _ = await cls.get_or_create( user, created = await cls.get_or_create(
user_id=user_id, user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()}, defaults={"platform": platform, "uid": 0},
) )
if created:
user.uid = await cls.get_new_uid()
user.gold += gold user.gold += gold
await user.save(update_fields=["gold"]) await user.save(update_fields=["gold", "uid"])
await UserGoldLog.create( await UserGoldLog.create(
user_id=user_id, gold=gold, handle=GoldHandle.GET, source=source user_id=user_id, gold=gold, handle=GoldHandle.GET, source=source
) )
@@ -111,14 +123,16 @@ class UserConsole(Model):
异常: 异常:
InsufficientGold: 金币不足 InsufficientGold: 金币不足
""" """
user, _ = await cls.get_or_create( user, created = await cls.get_or_create(
user_id=user_id, user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()}, defaults={"platform": platform, "uid": 0},
) )
if created:
user.uid = await cls.get_new_uid()
if user.gold < gold: if user.gold < gold:
raise InsufficientGold() raise InsufficientGold()
user.gold -= gold user.gold -= gold
await user.save(update_fields=["gold"]) await user.save(update_fields=["gold", "uid"])
await UserGoldLog.create( await UserGoldLog.create(
user_id=user_id, gold=gold, handle=handle, source=plugin_module user_id=user_id, gold=gold, handle=handle, source=plugin_module
) )
@@ -135,14 +149,16 @@ class UserConsole(Model):
num: 道具数量. num: 道具数量.
platform: 平台. platform: 平台.
""" """
user, _ = await cls.get_or_create( user, created = await cls.get_or_create(
user_id=user_id, user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()}, defaults={"platform": platform, "uid": 0},
) )
if created:
user.uid = await cls.get_new_uid()
if goods_uuid not in user.props: if goods_uuid not in user.props:
user.props[goods_uuid] = 0 user.props[goods_uuid] = 0
user.props[goods_uuid] += num user.props[goods_uuid] += num
await user.save(update_fields=["props"]) await user.save(update_fields=["props", "uid"])
@classmethod @classmethod
async def add_props_by_name( async def add_props_by_name(
@@ -172,17 +188,21 @@ class UserConsole(Model):
num: 道具数量. num: 道具数量.
platform: 平台. platform: 平台.
""" """
user, _ = await cls.get_or_create( user, created = await cls.get_or_create(
user_id=user_id, user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()}, defaults={"platform": platform, "uid": 0},
) )
if created:
user.uid = await cls.get_new_uid()
if goods_uuid not in user.props or user.props[goods_uuid] < num: if goods_uuid not in user.props or user.props[goods_uuid] < num:
if created:
await user.save(update_fields=["uid"])
raise GoodsNotFound("未找到商品或道具数量不足...") raise GoodsNotFound("未找到商品或道具数量不足...")
user.props[goods_uuid] -= num user.props[goods_uuid] -= num
if user.props[goods_uuid] <= 0: if user.props[goods_uuid] <= 0:
del user.props[goods_uuid] del user.props[goods_uuid]
await user.save(update_fields=["props"]) await user.save(update_fields=["props", "uid"])
@classmethod @classmethod
async def use_props_by_name( async def use_props_by_name(
+1 -1
View File
@@ -708,7 +708,7 @@ class CacheManager:
if self._cache_backend: if self._cache_backend:
try: try:
await self._cache_backend.close() # type: ignore await self._cache_backend.close() # type: ignore
except (AttributeError, Exception) as e: except Exception as e:
logger.warning(f"关闭缓存连接失败: {e}", LOG_COMMAND) logger.warning(f"关闭缓存连接失败: {e}", LOG_COMMAND)
self._cache_backend = None self._cache_backend = None
+2
View File
@@ -65,6 +65,8 @@ class CacheType(StrEnum):
"""用户权限""" """用户权限"""
LIMIT = "GLOBAL_LIMIT" LIMIT = "GLOBAL_LIMIT"
"""插件限制""" """插件限制"""
TEMP = "TEMP"
"""临时缓存"""
class DbLockType(StrEnum): class DbLockType(StrEnum):