refactor: enhance UserConsole UID management and remove mute plugin

This commit is contained in:
HibiKier
2025-12-25 09:47:01 +08:00
parent 52b32915cc
commit f86beb928f
8 changed files with 190 additions and 384 deletions
+156 -28
View File
@@ -1,12 +1,11 @@
from typing import ClassVar from tortoise import BaseDBAsyncClient, Tortoise, fields
from tortoise import fields
from tortoise.exceptions import IntegrityError from tortoise.exceptions import IntegrityError
from zhenxun.configs.config import BotConfig
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, DbLockType, GoldHandle from zhenxun.services.log import logger
from zhenxun.utils.enum import CacheType, GoldHandle
from zhenxun.utils.exception import GoodsNotFound, InsufficientGold from zhenxun.utils.exception import GoodsNotFound, InsufficientGold
from .user_gold_log import UserGoldLog from .user_gold_log import UserGoldLog
@@ -18,7 +17,7 @@ class UserConsole(Model):
user_id = fields.CharField(255, unique=True, description="用户id") user_id = fields.CharField(255, unique=True, description="用户id")
"""用户id""" """用户id"""
uid = fields.IntField(description="UID", unique=True) uid = fields.IntField(description="UID", unique=True)
"""UID""" """UID,用户可修改"""
gold = fields.IntField(default=100, description="金币数量") gold = fields.IntField(default=100, description="金币数量")
"""金币数量""" """金币数量"""
sign = fields.ReverseRelation["SignUser"] # type: ignore sign = fields.ReverseRelation["SignUser"] # type: ignore
@@ -39,39 +38,107 @@ class UserConsole(Model):
"""缓存类型""" """缓存类型"""
cache_key_field = "user_id" cache_key_field = "user_id"
"""缓存键字段""" """缓存键字段"""
lock_fields: ClassVar[dict[DbLockType, str]] = {DbLockType.CREATE: "user_id"}
@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): if user := await cls.get_or_none(user_id=user_id):
return user return user
# 使用数据库序列获取 uid,原子操作无竞争
uid = await cls._next_uid_from_sequence()
try: try:
uid = await cls.get_new_uid()
return await cls.create(user_id=user_id, uid=uid, platform=platform) return await cls.create(user_id=user_id, uid=uid, platform=platform)
except IntegrityError: except IntegrityError:
# 并发竞争下,依赖 user_id 唯一约束兜底查询 # user_id 冲突(并发创建同一用户)
return await cls.get(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 @classmethod
async def get_new_uid(cls) -> int: async def _next_uid_from_sequence(cls) -> int:
"""获取最新uid """获取下一个 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") return await cls.all().count()
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
@classmethod @classmethod
async def add_gold( async def add_gold(
@@ -194,7 +261,68 @@ class UserConsole(Model):
@classmethod @classmethod
async def _run_script(cls): async def _run_script(cls):
return [ """初始化脚本,根据数据库类型创建序列/表"""
"CREATE INDEX idx_user_console_user_id ON user_console(user_id);", db_type = BotConfig.get_sql_type()
"CREATE INDEX idx_user_console_uid ON user_console(uid);",
# 通用索引
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
-20
View File
@@ -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()))
-129
View File
@@ -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()
-74
View File
@@ -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)
-116
View File
@@ -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()
+32 -15
View File
@@ -7,7 +7,6 @@ from typing_extensions import Self
from tortoise.backends.base.client import BaseDBAsyncClient from tortoise.backends.base.client import BaseDBAsyncClient
from tortoise.exceptions import IntegrityError, MultipleObjectsReturned from tortoise.exceptions import IntegrityError, MultipleObjectsReturned
from tortoise.models import Model as TortoiseModel from tortoise.models import Model as TortoiseModel
from tortoise.transactions import in_transaction
from zhenxun.services.cache import CacheRoot from zhenxun.services.cache import CacheRoot
from zhenxun.services.log import logger from zhenxun.services.log import logger
@@ -204,7 +203,7 @@ class Model(TortoiseModel):
using_db: BaseDBAsyncClient | None = None, using_db: BaseDBAsyncClient | None = None,
**kwargs: Any, **kwargs: Any,
) -> tuple[Self, bool]: ) -> tuple[Self, bool]:
"""更新或创建数据(使用UPSERT锁)""" """更新或创建数据(优化版本,减少锁等待)"""
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {} lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
lock_key = None lock_key = None
if field := lock_fields.get(DbLockType.UPSERT): if field := lock_fields.get(DbLockType.UPSERT):
@@ -216,21 +215,39 @@ class Model(TortoiseModel):
async with cls._lock_context(DbLockType.UPSERT, lock_key): async with cls._lock_context(DbLockType.UPSERT, lock_key):
try: try:
# 先尝试更新(带行锁) # 优化:先尝试无锁查询,大部分情况数据已存在
async with in_transaction(): if obj := await cls.get_or_none(**kwargs):
if obj := await cls.filter(**kwargs).select_for_update().first(): if defaults:
await obj.update_from_dict(defaults or {}) await obj.update_from_dict(defaults)
await obj.save() # 只更新指定字段,减少写操作
result = (obj, False) await obj.save(update_fields=list(defaults.keys()))
else: if cache_type := cls.get_cache_type():
# 创建时不重复加锁 await CacheRoot.invalidate_cache(
result = await cls.create(**kwargs, **(defaults or {})), True cache_type, cls.get_cache_key(obj)
)
return obj, False
if cache_type := cls.get_cache_type(): # 数据不存在,尝试创建(依赖数据库唯一约束)
await CacheRoot.invalidate_cache( try:
cache_type, cls.get_cache_key(result[0]) 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: except IntegrityError:
# 处理极端情况下的唯一约束冲突 # 处理极端情况下的唯一约束冲突
obj = await cls.get(**kwargs) obj = await cls.get(**kwargs)
+1 -1
View File
@@ -3,7 +3,7 @@ from collections.abc import Callable
from pydantic import BaseModel from pydantic import BaseModel
# 数据库操作超时设置(秒) # 数据库操作超时设置(秒)
DB_TIMEOUT_SECONDS = 3.0 DB_TIMEOUT_SECONDS = 5.0
# 性能监控阈值(秒) # 性能监控阈值(秒)
SLOW_QUERY_THRESHOLD = 0.5 SLOW_QUERY_THRESHOLD = 0.5
+1 -1
View File
@@ -141,7 +141,7 @@ class BotProfileManager:
"""构建BOT自我介绍图片""" """构建BOT自我介绍图片"""
profile, service_count, call_count = await asyncio.gather( profile, service_count, call_count = await asyncio.gather(
cls.get_bot_profile(bot_id), cls.get_bot_profile(bot_id),
UserConsole.get_new_uid(), UserConsole.get_user_count(),
Statistics.filter(bot_id=bot_id).count(), Statistics.filter(bot_id=bot_id).count(),
) )
if not profile: if not profile: