mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-28 16:20:56 +08:00
refactor: optimize MuteManager data source
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user