refactor: optimize MuteManager data source

This commit is contained in:
HibiKier
2025-12-24 15:09:24 +08:00
parent be316a5caf
commit 52b32915cc
4 changed files with 56 additions and 57 deletions
+1 -4
View File
@@ -6,6 +6,7 @@ from typing_extensions import Self
from tortoise import fields from tortoise import fields
from tortoise.expressions import Q from tortoise.expressions import Q
from zhenxun.services.cache import CacheRoot
from zhenxun.services.data_access import DataAccess from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import Model from zhenxun.services.db_context import Model
from zhenxun.services.log import logger from zhenxun.services.log import logger
@@ -194,10 +195,6 @@ class BanConsole(Model):
返回: 返回:
list[Self]: ban记录列表,空列表表示未被ban list[Self]: ban记录列表,空列表表示未被ban
""" """
from zhenxun.services.cache import CacheRoot
from zhenxun.services.data_access import DataAccess
from zhenxun.utils.enum import CacheType
cache_key = f"{user_id}_{group_id}" cache_key = f"{user_id}_{group_id}"
results = await CacheRoot.get(CacheType.BAN, cache_key) results = await CacheRoot.get(CacheType.BAN, cache_key)
+44 -37
View File
@@ -1,6 +1,6 @@
import time import time
from pydantic import BaseModel from pydantic import BaseModel, Field
import ujson as json import ujson as json
from zhenxun.configs.config import Config from zhenxun.configs.config import Config
@@ -16,7 +16,7 @@ class GroupData(BaseModel):
"""检测时长""" """检测时长"""
duration: int duration: int
"""禁言时长""" """禁言时长"""
message_data: dict = {} message_data: dict = Field(default_factory=dict)
"""消息存储""" """消息存储"""
@@ -26,12 +26,13 @@ class MuteManager:
def __init__(self) -> None: def __init__(self) -> None:
self._group_data: dict[str, GroupData] = {} self._group_data: dict[str, GroupData] = {}
if self.file.exists(): if self.file.exists():
_data = json.load(open(self.file)) with open(self.file, encoding="utf-8") as f:
for gid in _data: _data = json.load(f)
for gid, gdata in _data.items():
self._group_data[gid] = GroupData( self._group_data[gid] = GroupData(
count=_data[gid]["count"], count=gdata["count"],
time=_data[gid]["time"], time=gdata["time"],
duration=_data[gid]["duration"], duration=gdata["duration"],
) )
def get_group_data(self, group_id: str) -> GroupData: def get_group_data(self, group_id: str) -> GroupData:
@@ -64,14 +65,15 @@ class MuteManager:
def save_data(self): def save_data(self):
"""保存数据""" """保存数据"""
data = {} data = {
for gid in self._group_data: gid: {
data[gid] = { "count": gdata.count,
"count": self._group_data[gid].count, "time": gdata.time,
"time": self._group_data[gid].time, "duration": gdata.duration,
"duration": self._group_data[gid].duration,
} }
with open(self.file, "w") as f: 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) json.dump(data, f, indent=4, ensure_ascii=False)
def add_message(self, user_id: str, group_id: str, message: str) -> int: def add_message(self, user_id: str, group_id: str, message: str) -> int:
@@ -85,37 +87,42 @@ class MuteManager:
返回: 返回:
int: 禁言时长 int: 禁言时长
""" """
if group_id not in self._group_data: group_data = self.get_group_data(group_id)
self._group_data[group_id] = GroupData(
count=base_config.get("MUTE_DEFAULT_COUNT"),
time=base_config.get("MUTE_DEFAULT_TIME"),
duration=base_config.get("MUTE_DEFAULT_DURATION"),
)
group_data = self._group_data[group_id]
if group_data.duration == 0: if group_data.duration == 0:
return 0 return 0
message_data = group_data.message_data message_data = group_data.message_data
if not message_data.get(user_id): user_data = message_data.get(user_id)
now = time.time()
if not user_data:
message_data[user_id] = { message_data[user_id] = {
"time": time.time(), "time": now,
"count": 1, "count": 1,
"message": message, "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: else:
if message.find(message_data[user_id]["message"]) != -1: user_data["time"] = now
message_data[user_id]["count"] += 1 user_data["count"] = 1
else:
message_data[user_id]["time"] = time.time() user_data["message"] = message
message_data[user_id]["count"] = 1
message_data[user_id]["message"] = message # 检测是否触发刷屏
if time.time() - message_data[user_id]["time"] > group_data.time: if user_data["count"] > group_data.count:
message_data[user_id]["time"] = time.time() return group_data.duration
message_data[user_id]["count"] = 1
if (
message_data[user_id]["count"] > group_data.count
and time.time() - message_data[user_id]["time"] < group_data.time
):
return group_data.duration
return 0 return 0
+4 -8
View File
@@ -7,7 +7,6 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import BotConfig from zhenxun.configs.config import BotConfig
from zhenxun.configs.utils import PluginExtraData from zhenxun.configs.utils import PluginExtraData
from zhenxun.models.ban_console import BanConsole from zhenxun.models.ban_console import BanConsole
from zhenxun.services.data_access import DataAccess
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType from zhenxun.utils.enum import PluginType
from zhenxun.utils.image_utils import get_download_image_hash from zhenxun.utils.image_utils import get_download_image_hash
@@ -19,8 +18,8 @@ from ._data_source import mute_manager
__plugin_meta__ = PluginMetadata( __plugin_meta__ = PluginMetadata(
name="刷屏监听", name="刷屏监听",
description="", description="这是刷屏检测的监听器,用于检测用户是否在规定时间内发送了相同的消息",
usage="", usage="无",
extra=PluginExtraData( extra=PluginExtraData(
author="HibiKier", author="HibiKier",
version="0.1-473ecd8", version="0.1-473ecd8",
@@ -34,12 +33,9 @@ async def rule(session: Uninfo) -> bool:
entity_ids = get_entity_ids(session) entity_ids = get_entity_ids(session)
if not session.group: if not session.group:
return False return False
ban_dao = DataAccess(BanConsole) if mute_manager.get_group_data(entity_ids.group_id or "0").duration == 0:
if not await ban_dao.safe_get_or_none(
user_id=entity_ids.user_id, group_id=entity_ids.group_id
):
return False return False
if not await ban_dao.safe_get_or_none(user_id="", group_id=entity_ids.group_id): if await BanConsole.is_ban_cached(entity_ids.user_id, entity_ids.group_id):
return False return False
return True return True
+7 -8
View File
@@ -1,6 +1,6 @@
from nonebot.plugin import PluginMetadata from nonebot.plugin import PluginMetadata
from nonebot_plugin_alconna import Alconna, Args, Arparma, Match, Option, on_alconna from nonebot_plugin_alconna import Alconna, Args, Arparma, Match, Option, on_alconna
from nonebot_plugin_session import EventSession from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import BotConfig from zhenxun.configs.config import BotConfig
from zhenxun.configs.utils import PluginExtraData, RegisterConfig from zhenxun.configs.utils import PluginExtraData, RegisterConfig
@@ -8,8 +8,9 @@ from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType from zhenxun.utils.enum import PluginType
from zhenxun.utils.message import MessageUtils from zhenxun.utils.message import MessageUtils
from zhenxun.utils.rules import ensure_group from zhenxun.utils.rules import ensure_group
from zhenxun.utils.utils import get_entity_ids
from ._data_source import base_config, mute_manage from ._data_source import base_config, mute_manager
__plugin_meta__ = PluginMetadata( __plugin_meta__ = PluginMetadata(
name="刷屏禁言", name="刷屏禁言",
@@ -82,19 +83,17 @@ _setting_matcher = on_alconna(
@_setting_matcher.handle() @_setting_matcher.handle()
async def _( async def _(
session: EventSession, session: Uninfo,
arparma: Arparma, arparma: Arparma,
time: Match[int], time: Match[int],
count: Match[int], count: Match[int],
duration: Match[int], duration: Match[int],
): ):
group_id = session.id2 entity_ids = get_entity_ids(session)
if not session.id1 or not group_id:
return
_time = time.result if time.available else None _time = time.result if time.available else None
_count = count.result if count.available else None _count = count.result if count.available else None
_duration = duration.result if duration.available else None _duration = duration.result if duration.available else None
group_data = mute_manage.get_group_data(group_id) 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: if _time is None and _count is None and _duration is None:
await MessageUtils.build_message( await MessageUtils.build_message(
f"最大次数:{group_data.count} 次\n" f"最大次数:{group_data.count} 次\n"
@@ -114,4 +113,4 @@ async def _(
arparma.header_result, arparma.header_result,
session=session, session=session,
) )
mute_manage.save_data() mute_manager.save_data()