From f86beb928fd0cb043acfe4891f120785f44e4348 Mon Sep 17 00:00:00 2001 From: HibiKier <775757368@qq.com> Date: Thu, 25 Dec 2025 09:47:01 +0800 Subject: [PATCH] refactor: enhance UserConsole UID management and remove mute plugin --- zhenxun/models/user_console.py | 184 ++++++++++++++++--- zhenxun/plugins/mute/__init__.py | 20 -- zhenxun/plugins/mute/_data_source.py | 129 ------------- zhenxun/plugins/mute/mute_message.py | 74 -------- zhenxun/plugins/mute/mute_setting.py | 116 ------------ zhenxun/services/db_context/base_model.py | 47 +++-- zhenxun/services/db_context/config.py | 2 +- zhenxun/utils/manager/bot_profile_manager.py | 2 +- 8 files changed, 190 insertions(+), 384 deletions(-) delete mode 100644 zhenxun/plugins/mute/__init__.py delete mode 100644 zhenxun/plugins/mute/_data_source.py delete mode 100644 zhenxun/plugins/mute/mute_message.py delete mode 100644 zhenxun/plugins/mute/mute_setting.py diff --git a/zhenxun/models/user_console.py b/zhenxun/models/user_console.py index c2b4dce9..08d2857f 100644 --- a/zhenxun/models/user_console.py +++ b/zhenxun/models/user_console.py @@ -1,12 +1,11 @@ -from typing import ClassVar - -from tortoise import fields +from tortoise import BaseDBAsyncClient, Tortoise, fields from tortoise.exceptions import IntegrityError +from zhenxun.configs.config import BotConfig from zhenxun.models.goods_info import GoodsInfo -from zhenxun.services.cache import CacheRoot from zhenxun.services.db_context import Model -from zhenxun.utils.enum import CacheType, DbLockType, GoldHandle +from zhenxun.services.log import logger +from zhenxun.utils.enum import CacheType, GoldHandle from zhenxun.utils.exception import GoodsNotFound, InsufficientGold from .user_gold_log import UserGoldLog @@ -18,7 +17,7 @@ class UserConsole(Model): user_id = fields.CharField(255, unique=True, description="用户id") """用户id""" uid = fields.IntField(description="UID", unique=True) - """UID""" + """UID,用户可修改""" gold = fields.IntField(default=100, description="金币数量") """金币数量""" sign = fields.ReverseRelation["SignUser"] # type: ignore @@ -39,39 +38,107 @@ class UserConsole(Model): """缓存类型""" cache_key_field = "user_id" """缓存键字段""" - lock_fields: ClassVar[dict[DbLockType, str]] = {DbLockType.CREATE: "user_id"} @classmethod 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 + # 使用数据库序列获取 uid,原子操作无竞争 + uid = await cls._next_uid_from_sequence() + try: - uid = await cls.get_new_uid() return await cls.create(user_id=user_id, uid=uid, platform=platform) except IntegrityError: - # 并发竞争下,依赖 user_id 唯一约束兜底查询 - return await cls.get(user_id=user_id) + # user_id 冲突(并发创建同一用户) + if user := await cls.get_or_none(user_id=user_id): + return user + # uid 冲突(极罕见,用户手动修改了 uid),重试 + for _ in range(3): + try: + uid = await cls._next_uid_from_sequence() + return await cls.create(user_id=user_id, uid=uid, platform=platform) + except IntegrityError: + if user := await cls.get_or_none(user_id=user_id): + return user + raise @classmethod - async def get_new_uid(cls) -> int: - """获取最新uid + async def _next_uid_from_sequence(cls) -> int: + """获取下一个 UID(原子操作,支持 PostgreSQL/MySQL/SQLite)""" + conn = Tortoise.get_connection("default") + db_type = BotConfig.get_sql_type() + + try: + if db_type == "postgresql": + return await cls._next_uid_postgresql(conn) + elif db_type == "mysql": + return await cls._next_uid_mysql(conn) + else: # sqlite + return await cls._next_uid_sqlite(conn) + except Exception as e: + logger.debug(f"序列获取失败,使用备用方案: {e}") + return await cls._get_max_uid() + 1 + + @classmethod + async def _next_uid_postgresql(cls, conn: BaseDBAsyncClient) -> int: + """PostgreSQL: 使用序列""" + result = await conn.execute_query_dict( + "SELECT nextval('user_console_uid_seq') as uid" + ) + return result[0]["uid"] + + @classmethod + async def _next_uid_mysql(cls, conn: BaseDBAsyncClient) -> int: + """MySQL: 使用序列表实现原子自增""" + # 原子更新并获取新值 + await conn.execute_query( + """ + INSERT INTO user_console_sequence (id, current_value) + VALUES (1, 1) + ON DUPLICATE KEY UPDATE current_value = current_value + 1 + """ + ) + result = await conn.execute_query_dict( + "SELECT current_value as uid FROM user_console_sequence WHERE id = 1" + ) + return result[0]["uid"] + + @classmethod + async def _next_uid_sqlite(cls, conn: BaseDBAsyncClient) -> int: + """SQLite: 使用序列表实现原子自增""" + # SQLite 使用 INSERT OR REPLACE 实现原子操作 + await conn.execute_query( + """ + INSERT OR REPLACE INTO user_console_sequence (id, current_value) + VALUES (1, COALESCE( + (SELECT current_value + 1 FROM user_console_sequence WHERE id = 1), + (SELECT COALESCE(MAX(uid), 0) + 1 FROM user_console) + )) + """ + ) + result = await conn.execute_query_dict( + "SELECT current_value as uid FROM user_console_sequence WHERE id = 1" + ) + return result[0]["uid"] + + @classmethod + async def _get_max_uid(cls) -> int: + """获取当前最大 uid(备用方案)""" + data: list[int] = ( # pyright: ignore[reportAssignmentType] + await cls.annotate().order_by("-uid").limit(1).values_list("uid", flat=True) + ) + return data[0] if data else 0 + + @classmethod + async def get_user_count(cls) -> int: + """获取用户总数 返回: - int: 最新uid + int: 用户总数 """ - uid: int | None = await CacheRoot.get(CacheType.TEMP, "USER_CONSOLE_UID") - if uid is None: - 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 or 1 + return await cls.all().count() @classmethod async def add_gold( @@ -194,7 +261,68 @@ class UserConsole(Model): @classmethod async def _run_script(cls): - return [ - "CREATE INDEX idx_user_console_user_id ON user_console(user_id);", - "CREATE INDEX idx_user_console_uid ON user_console(uid);", + """初始化脚本,根据数据库类型创建序列/表""" + db_type = BotConfig.get_sql_type() + + # 通用索引 + scripts = [ + "CREATE INDEX IF NOT EXISTS idx_user_console_user_id " + "ON user_console(user_id);", + "CREATE INDEX IF NOT EXISTS idx_user_console_uid ON user_console(uid);", ] + + # 根据数据库类型添加序列初始化脚本 + if db_type == "postgresql": + scripts.append( + """ + DO $$ + BEGIN + IF NOT EXISTS ( + SELECT 1 FROM pg_sequences + WHERE schemaname = 'public' + AND sequencename = 'user_console_uid_seq' + ) THEN + CREATE SEQUENCE user_console_uid_seq; + PERFORM setval( + 'user_console_uid_seq', + COALESCE((SELECT MAX(uid) FROM user_console), 0) + 1, + false + ); + END IF; + END $$; + """ + ) + elif db_type == "mysql": + # MySQL: 创建序列表 + scripts.extend( + [ + """ + CREATE TABLE IF NOT EXISTS user_console_sequence ( + id INT PRIMARY KEY, + current_value BIGINT NOT NULL DEFAULT 0 + ); + """, + """ + INSERT IGNORE INTO user_console_sequence (id, current_value) + SELECT 1, COALESCE(MAX(uid), 0) FROM user_console; + """, + ] + ) + else: # sqlite + # SQLite: 创建序列表 + scripts.extend( + [ + """ + CREATE TABLE IF NOT EXISTS user_console_sequence ( + id INTEGER PRIMARY KEY, + current_value INTEGER NOT NULL DEFAULT 0 + ); + """, + """ + INSERT OR IGNORE INTO user_console_sequence (id, current_value) + SELECT 1, COALESCE(MAX(uid), 0) FROM user_console; + """, + ] + ) + + return scripts diff --git a/zhenxun/plugins/mute/__init__.py b/zhenxun/plugins/mute/__init__.py deleted file mode 100644 index f846f4b6..00000000 --- a/zhenxun/plugins/mute/__init__.py +++ /dev/null @@ -1,20 +0,0 @@ -from pathlib import Path - -import nonebot -from nonebot.plugin import PluginMetadata - -from zhenxun.configs.utils import PluginExtraData -from zhenxun.utils.enum import PluginType - -__plugin_meta__ = PluginMetadata( - name="刷屏禁言检测", - description="", - usage="", - extra=PluginExtraData( - author="HibiKier", - version="0.1-473ecd8", - plugin_type=PluginType.PARENT, - ).to_dict(), -) - -nonebot.load_plugins(str(Path(__file__).parent.resolve())) diff --git a/zhenxun/plugins/mute/_data_source.py b/zhenxun/plugins/mute/_data_source.py deleted file mode 100644 index 5d356f53..00000000 --- a/zhenxun/plugins/mute/_data_source.py +++ /dev/null @@ -1,129 +0,0 @@ -import time - -from pydantic import BaseModel, Field -import ujson as json - -from zhenxun.configs.config import Config -from zhenxun.configs.path_config import DATA_PATH - -base_config = Config.get("mute_setting") - - -class GroupData(BaseModel): - count: int - """次数""" - time: int - """检测时长""" - duration: int - """禁言时长""" - message_data: dict = Field(default_factory=dict) - """消息存储""" - - -class MuteManager: - file = DATA_PATH / "group_mute_data.json" - - def __init__(self) -> None: - self._group_data: dict[str, GroupData] = {} - if self.file.exists(): - with open(self.file, encoding="utf-8") as f: - _data = json.load(f) - for gid, gdata in _data.items(): - self._group_data[gid] = GroupData( - count=gdata["count"], - time=gdata["time"], - duration=gdata["duration"], - ) - - def get_group_data(self, group_id: str) -> GroupData: - """获取群组数据 - - 参数: - group_id: 群组id - - 返回: - GroupData: GroupData - """ - if group_id not in self._group_data: - self._group_data[group_id] = GroupData( - count=base_config.get("MUTE_DEFAULT_COUNT", 10) or 10, - time=base_config.get("MUTE_DEFAULT_TIME", 7) or 7, - duration=base_config.get("MUTE_DEFAULT_DURATION", 10) or 10, - ) - return self._group_data[group_id] - - def reset(self, user_id: str, group_id: str): - """重置用户检查次数 - - 参数: - user_id: 用户id - group_id: 群组id - """ - if group_data := self._group_data.get(group_id): - if user_id in group_data.message_data: - group_data.message_data[user_id]["count"] = 0 - - def save_data(self): - """保存数据""" - data = { - gid: { - "count": gdata.count, - "time": gdata.time, - "duration": gdata.duration, - } - for gid, gdata in self._group_data.items() - } - with open(self.file, "w", encoding="utf-8") as f: - json.dump(data, f, indent=4, ensure_ascii=False) - - def add_message(self, user_id: str, group_id: str, message: str) -> int: - """添加消息 - - 参数: - user_id: 用户id - group_id: 群组id - message: 消息内容 - - 返回: - int: 禁言时长 - """ - group_data = self.get_group_data(group_id) - if group_data.duration == 0: - return 0 - - message_data = group_data.message_data - user_data = message_data.get(user_id) - now = time.time() - - if not user_data: - message_data[user_id] = { - "time": now, - "count": 1, - "message": message, - } - return 0 - - # 超过检测时间窗口,重置计数 - if now - user_data["time"] > group_data.time: - user_data["time"] = now - user_data["count"] = 1 - user_data["message"] = message - return 0 - - # 消息内容相似(包含之前的消息),累加计数 - if user_data["message"] in message: - user_data["count"] += 1 - else: - user_data["time"] = now - user_data["count"] = 1 - - user_data["message"] = message - - # 检测是否触发刷屏 - if user_data["count"] > group_data.count: - return group_data.duration - - return 0 - - -mute_manager = MuteManager() diff --git a/zhenxun/plugins/mute/mute_message.py b/zhenxun/plugins/mute/mute_message.py deleted file mode 100644 index 9c185bbf..00000000 --- a/zhenxun/plugins/mute/mute_message.py +++ /dev/null @@ -1,74 +0,0 @@ -from nonebot import on_message -from nonebot.adapters import Bot -from nonebot.plugin import PluginMetadata -from nonebot_plugin_alconna import Image, UniMsg -from nonebot_plugin_uninfo import Uninfo - -from zhenxun.configs.config import BotConfig -from zhenxun.configs.utils import PluginExtraData -from zhenxun.models.ban_console import BanConsole -from zhenxun.services.log import logger -from zhenxun.utils.enum import PluginType -from zhenxun.utils.image_utils import get_download_image_hash -from zhenxun.utils.message import MessageUtils -from zhenxun.utils.platform import PlatformUtils -from zhenxun.utils.utils import FreqLimiter, get_entity_ids - -from ._data_source import mute_manager - -__plugin_meta__ = PluginMetadata( - name="刷屏监听", - description="这是刷屏检测的监听器,用于检测用户是否在规定时间内发送了相同的消息", - usage="无", - extra=PluginExtraData( - author="HibiKier", - version="0.1-473ecd8", - menu_type="其他", - plugin_type=PluginType.DEPENDANT, - ).to_dict(), -) - - -async def rule(session: Uninfo) -> bool: - entity_ids = get_entity_ids(session) - if not session.group: - return False - if mute_manager.get_group_data(entity_ids.group_id or "0").duration == 0: - return False - if await BanConsole.is_ban_cached(entity_ids.user_id, entity_ids.group_id): - return False - return True - - -_matcher = on_message(rule=rule, priority=1, block=False) - -_flmt = FreqLimiter(30) - - -@_matcher.handle() -async def _(bot: Bot, session: Uninfo, message: UniMsg): - entity_ids = get_entity_ids(session) - plain_text = message.extract_plain_text() - image_list = [m.url for m in message if isinstance(m, Image) and m.url] - img_hash = "" - for url in image_list: - img_hash += await get_download_image_hash(url, "_mute_") - _message = plain_text + img_hash - if duration := mute_manager.add_message( - entity_ids.user_id, entity_ids.group_id or "0", _message - ): - try: - if _flmt.check(entity_ids.user_id): - _flmt.start_cd(entity_ids.user_id) - await PlatformUtils.ban_user( - bot, entity_ids.user_id, entity_ids.group_id or "0", duration - ) - await MessageUtils.build_message( - f"检测到恶意刷屏,{BotConfig.self_nickname}要把你关进小黑屋!" - ).send(at_sender=True) - mute_manager.reset(entity_ids.user_id, entity_ids.group_id or "0") - logger.info( - f"检测刷屏 被禁言 {duration} 分钟", "禁言检查", session=session - ) - except Exception as e: - logger.error("禁言发送错误", "禁言检测", session=session, e=e) diff --git a/zhenxun/plugins/mute/mute_setting.py b/zhenxun/plugins/mute/mute_setting.py deleted file mode 100644 index a98eb3be..00000000 --- a/zhenxun/plugins/mute/mute_setting.py +++ /dev/null @@ -1,116 +0,0 @@ -from nonebot.plugin import PluginMetadata -from nonebot_plugin_alconna import Alconna, Args, Arparma, Match, Option, on_alconna -from nonebot_plugin_uninfo import Uninfo - -from zhenxun.configs.config import BotConfig -from zhenxun.configs.utils import PluginExtraData, RegisterConfig -from zhenxun.services.log import logger -from zhenxun.utils.enum import PluginType -from zhenxun.utils.message import MessageUtils -from zhenxun.utils.rules import ensure_group -from zhenxun.utils.utils import get_entity_ids - -from ._data_source import base_config, mute_manager - -__plugin_meta__ = PluginMetadata( - name="刷屏禁言", - description="刷屏禁言相关操作", - usage=f""" - 刷屏禁言相关操作,需要 {BotConfig.self_nickname} 有群管理员权限 - 指令: - 刷屏设置: 查看当前设置 - -c [count]: 检测最大次数 - -t [time]: 规定时间内 - -d [duration]: 禁言时长 - 示例: - 刷屏设置 -c 10: 设置最大次数为10 - 刷屏设置 -t 100 -d 20: 设置规定时间和禁言时长 - 刷屏设置 -d 10: 设置禁言时长为10 - * 即 X 秒内发送同样消息 N 次,禁言 M 分钟 * - """.strip(), - extra=PluginExtraData( - author="HibiKier", - version="0.1-473ecd8", - menu_type="其他", - plugin_type=PluginType.ADMIN, - admin_level=base_config.get("MUTE_LEVEL", 5), - configs=[ - RegisterConfig( - key="MUTE_LEVEL", - value=5, - help="更改禁言设置的管理权限", - default_value=5, - type=int, - ), - RegisterConfig( - key="MUTE_DEFAULT_COUNT", - value=10, - help="刷屏禁言默认检测次数", - default_value=10, - type=int, - ), - RegisterConfig( - key="MUTE_DEFAULT_TIME", - value=7, - help="刷屏检测默认规定时间", - default_value=7, - type=int, - ), - RegisterConfig( - key="MUTE_DEFAULT_DURATION", - value=10, - help="刷屏检测默禁言时长(分钟)", - default_value=10, - type=int, - ), - ], - ).to_dict(), -) - - -_setting_matcher = on_alconna( - Alconna( - "刷屏设置", - Option("-t|--time", Args["time", int], help_text="检测时长"), - Option("-c|--count", Args["count", int], help_text="检测次数"), - Option("-d|--duration", Args["duration", int], help_text="禁言时长"), - ), - rule=ensure_group, - block=True, - priority=5, -) - - -@_setting_matcher.handle() -async def _( - session: Uninfo, - arparma: Arparma, - time: Match[int], - count: Match[int], - duration: Match[int], -): - entity_ids = get_entity_ids(session) - _time = time.result if time.available else None - _count = count.result if count.available else None - _duration = duration.result if duration.available else None - group_data = mute_manager.get_group_data(entity_ids.group_id or "0") - if _time is None and _count is None and _duration is None: - await MessageUtils.build_message( - f"最大次数:{group_data.count} 次\n" - f"规定时间:{group_data.time} 秒\n" - f"禁言时长:{group_data.duration:.2f} 分钟\n" - f"【在规定时间内发送相同消息超过最大次数则禁言\n当禁言时长为0时关闭此功能】" - ).finish(reply_to=True) - if _time is not None: - group_data.time = _time - if _count is not None: - group_data.count = _count - if _duration is not None: - group_data.duration = _duration - await MessageUtils.build_message("设置成功!").send(reply_to=True) - logger.info( - f"设置禁言配置 time: {_time}, count: {_count}, duration: {_duration}", - arparma.header_result, - session=session, - ) - mute_manager.save_data() diff --git a/zhenxun/services/db_context/base_model.py b/zhenxun/services/db_context/base_model.py index 5e3bd19b..35019a64 100644 --- a/zhenxun/services/db_context/base_model.py +++ b/zhenxun/services/db_context/base_model.py @@ -7,7 +7,6 @@ from typing_extensions import Self from tortoise.backends.base.client import BaseDBAsyncClient from tortoise.exceptions import IntegrityError, MultipleObjectsReturned from tortoise.models import Model as TortoiseModel -from tortoise.transactions import in_transaction from zhenxun.services.cache import CacheRoot from zhenxun.services.log import logger @@ -204,7 +203,7 @@ class Model(TortoiseModel): using_db: BaseDBAsyncClient | None = None, **kwargs: Any, ) -> tuple[Self, bool]: - """更新或创建数据(使用UPSERT锁)""" + """更新或创建数据(优化版本,减少锁等待)""" lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {} lock_key = None if field := lock_fields.get(DbLockType.UPSERT): @@ -216,21 +215,39 @@ class Model(TortoiseModel): async with cls._lock_context(DbLockType.UPSERT, lock_key): try: - # 先尝试更新(带行锁) - async with in_transaction(): - if obj := await cls.filter(**kwargs).select_for_update().first(): - await obj.update_from_dict(defaults or {}) - await obj.save() - result = (obj, False) - else: - # 创建时不重复加锁 - result = await cls.create(**kwargs, **(defaults or {})), True + # 优化:先尝试无锁查询,大部分情况数据已存在 + if obj := await cls.get_or_none(**kwargs): + if defaults: + await obj.update_from_dict(defaults) + # 只更新指定字段,减少写操作 + await obj.save(update_fields=list(defaults.keys())) + if cache_type := cls.get_cache_type(): + await CacheRoot.invalidate_cache( + cache_type, cls.get_cache_key(obj) + ) + return obj, False - if cache_type := cls.get_cache_type(): - await CacheRoot.invalidate_cache( - cache_type, cls.get_cache_key(result[0]) + # 数据不存在,尝试创建(依赖数据库唯一约束) + try: + obj = await super().create( + using_db=using_db, **kwargs, **(defaults or {}) ) - return result + if cache_type := cls.get_cache_type(): + await CacheRoot.invalidate_cache( + cache_type, cls.get_cache_key(obj) + ) + return obj, True + except IntegrityError: + # 并发创建冲突,重新获取并更新 + obj = await cls.get(**kwargs) + if defaults: + await obj.update_from_dict(defaults) + await obj.save(update_fields=list(defaults.keys())) + if cache_type := cls.get_cache_type(): + await CacheRoot.invalidate_cache( + cache_type, cls.get_cache_key(obj) + ) + return obj, False except IntegrityError: # 处理极端情况下的唯一约束冲突 obj = await cls.get(**kwargs) diff --git a/zhenxun/services/db_context/config.py b/zhenxun/services/db_context/config.py index ae6d6b8c..35fdd411 100644 --- a/zhenxun/services/db_context/config.py +++ b/zhenxun/services/db_context/config.py @@ -3,7 +3,7 @@ from collections.abc import Callable from pydantic import BaseModel # 数据库操作超时设置(秒) -DB_TIMEOUT_SECONDS = 3.0 +DB_TIMEOUT_SECONDS = 5.0 # 性能监控阈值(秒) SLOW_QUERY_THRESHOLD = 0.5 diff --git a/zhenxun/utils/manager/bot_profile_manager.py b/zhenxun/utils/manager/bot_profile_manager.py index 22eded44..28e19637 100644 --- a/zhenxun/utils/manager/bot_profile_manager.py +++ b/zhenxun/utils/manager/bot_profile_manager.py @@ -141,7 +141,7 @@ class BotProfileManager: """构建BOT自我介绍图片""" profile, service_count, call_count = await asyncio.gather( cls.get_bot_profile(bot_id), - UserConsole.get_new_uid(), + UserConsole.get_user_count(), Statistics.filter(bot_id=bot_id).count(), ) if not profile: