mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-28 16:20:56 +08:00
refactor: enhance UserConsole UID management and remove mute plugin
This commit is contained in:
+156
-28
@@ -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
|
||||
|
||||
@@ -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()))
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user