Files
HeXi/hexi/plugins/nonebot_plugin_learning_chat/services/learn.py
T
sansenhoshi 131b92b319 结构调整
视频解析多图/多媒体结构 消息体适配
2026-09-08 14:25:32 +08:00

780 lines
36 KiB
Python

import asyncio
import datetime
import random
import re
import time
from dataclasses import dataclass
from functools import cached_property, cmp_to_key
# from ..core.MessageUtils import send_poke
try:
import jieba_fast.analyse as jieba_analyse
except ImportError:
import jieba.analyse as jieba_analyse
from typing import List, Union, Optional, Tuple
from enum import IntEnum, auto
from nonebot import get_adapter
from nonebot.adapters.onebot.v11 import GroupMessageEvent, MessageSegment, ActionFailed, Adapter
from sqlalchemy import select, delete, func
from nonebot_plugin_orm import get_session
from ..models import ChatBlackList, ChatContext, ChatAnswer, ChatMessage
from ..config import (
config_manager,
SUPERUSERS,
NICKNAME,
COMMAND_START,
log_info,
log_info,
)
chat_config = config_manager.config
NO_PERMISSION_WORDS = [f"{NICKNAME}就喜欢说这个,哼!", f"你管得着{NICKNAME}吗!"]
ENABLE_WORDS = [f"{NICKNAME}会尝试学你们说怪话!", f"好的呢,让{NICKNAME}学学你们的说话方式~"]
DISABLE_WORDS = [f"好好好,{NICKNAME}不学说话就是了!", f"果面呐噻,{NICKNAME}以后不学了..."]
SORRY_WORDS = [
f"{NICKNAME}知道错了...达咩!",
f"{NICKNAME}不会再这么说了...",
f"果面呐噻,{NICKNAME}说错话了...",
]
DOUBT_WORDS = [f"{NICKNAME}有说什么奇怪的话吗?"]
BREAK_REPEAT_WORDS = ["打断复读", "打断!"]
ALL_WORDS = (
NO_PERMISSION_WORDS
+ SORRY_WORDS
+ DOUBT_WORDS
+ ENABLE_WORDS
+ DISABLE_WORDS
+ BREAK_REPEAT_WORDS
)
class Result(IntEnum):
Learn = auto()
Pass = auto()
Repeat = auto()
Ban = auto()
SetEnable = auto()
SetDisable = auto()
@dataclass
class MessageData:
"""用于在 session 外安全传递消息数据的普通数据类"""
group_id: int
user_id: int
message_id: int
message: str
raw_message: str
plain_text: str
time: int
@cached_property
def is_plain_text(self) -> bool:
return "[CQ:" not in self.message
@cached_property
def keyword_list(self) -> List[str]:
if not self.is_plain_text and not len(self.plain_text):
return []
return jieba_analyse.extract_tags(self.plain_text, topK=chat_config.KEYWORDS_SIZE)
@cached_property
def keywords(self) -> str:
if not self.is_plain_text and not len(self.plain_text):
return self.message
return (
self.message if len(self.keyword_list) < 2 else " ".join(self.keyword_list)
)
def _db_to_message_data(m: ChatMessage) -> MessageData:
"""将数据库 ChatMessage 对象转为普通 MessageData"""
return MessageData(
group_id=m.group_id,
user_id=m.user_id,
message_id=m.message_id,
message=m.message,
raw_message=m.raw_message,
plain_text=m.plain_text,
time=m.time,
)
class LearningChat:
def __init__(self, event: GroupMessageEvent):
if event.reply:
self.reply = event.reply
self.data = MessageData(
group_id=event.group_id,
user_id=event.user_id,
message_id=event.message_id,
message=re.sub(
r"(\[CQ:at,qq=.+])|(\[CQ:reply,id=.+])",
"",
re.sub(r"(,subType=\d+,url=.+])", r"]", event.raw_message),
).strip(),
raw_message=event.raw_message,
plain_text=event.get_plaintext(),
time=event.time,
)
else:
self.data = MessageData(
group_id=event.group_id,
user_id=event.user_id,
message_id=event.message_id,
message=re.sub(
r"(\[CQ:at,qq=.+])",
"",
re.sub(r"(,subType=\d+,url=.+])", r"]", event.raw_message),
).strip(),
raw_message=event.raw_message,
plain_text=event.get_plaintext(),
time=event.time,
)
self.reply = None
self.bot_id = event.self_id
self.to_me = event.to_me or NICKNAME in self.data.message
self.role = "superuser" if event.user_id in SUPERUSERS else event.sender.role
self.config = config_manager.get_group_config(self.data.group_id)
self.ban_users = set(chat_config.ban_users + self.config.ban_users)
self.ban_words = set(chat_config.ban_words + self.config.ban_words)
async def _learn(self) -> Result:
if self.to_me and any(w in self.data.message for w in {"学说话", "快学", "开启学习"}):
return Result.SetEnable
elif self.to_me and any(w in self.data.message for w in {"闭嘴", "别学", "关闭学习"}):
return Result.SetDisable
elif not chat_config.total_enable or not self.config.enable:
log_info("群聊学习", f"➤该群<m>{self.data.group_id}</m>未开启群聊学习,跳过")
return Result.Pass
elif COMMAND_START and self.data.message.startswith(tuple(COMMAND_START)):
log_info("群聊学习", "➤该消息以命令前缀开头,跳过")
return Result.Pass
elif self.data.user_id in self.ban_users:
log_info("群聊学习", f"➤发言人<m>{self.data.user_id}</m>在屏蔽列表中,跳过")
return Result.Pass
elif self.to_me and any(w in self.data.message for w in {"不可以", "达咩", "不能说这"}):
return Result.Ban
async with get_session() as session:
if not await self._check_allow(session, self.data):
log_info("群聊学习", "➤消息未通过校验,跳过")
return Result.Pass
if self.reply:
result = await session.execute(
select(ChatMessage).where(ChatMessage.message_id == self.reply.message_id).limit(1)
)
db_message = result.scalar_one_or_none()
if not db_message:
log_info("群聊学习", "➤回复的消息不在数据库中,跳过")
return Result.Pass
if db_message.user_id in self.ban_users:
log_info("群聊学习", "➤回复的人在屏蔽列表中,跳过")
return Result.Pass
if not await self._check_allow(session, _db_to_message_data(db_message)):
log_info("群聊学习", "➤回复的消息未通过校验,跳过")
return Result.Pass
await self._set_answer(session, _db_to_message_data(db_message))
return Result.Learn
else:
result = await session.execute(
select(ChatMessage)
.where(
ChatMessage.group_id == self.data.group_id,
ChatMessage.time >= self.data.time - 3600,
)
.order_by(ChatMessage.time.desc())
.limit(5)
)
messages = [_db_to_message_data(m) for m in result.scalars().all()]
if messages:
if messages[0].message == self.data.message:
log_info("群聊学习", "➤复读中,跳过")
return Result.Repeat
for message in messages:
if (
message.user_id not in self.ban_users
and set(self.data.keyword_list) & set(message.keyword_list)
and self.data.keyword_list != message.keyword_list
and await self._check_allow(session, message)
):
await self._set_answer(session, message)
return Result.Learn
if messages[0].user_id in self.ban_users or not await self._check_allow(session, messages[0]):
log_info("群聊学习", "➤最后一条消息未通过校验,跳过")
return Result.Pass
await self._set_answer(session, messages[0])
return Result.Learn
else:
return Result.Pass
async def answer(self) -> Optional[List[Union[MessageSegment, str]]]:
"""获取这句话的回复"""
result = await self._learn()
# 保存本条消息到数据库
async with get_session() as session:
session.add(ChatMessage(
group_id=self.data.group_id,
user_id=self.data.user_id,
message_id=self.data.message_id,
message=self.data.message,
raw_message=self.data.raw_message,
plain_text=self.data.plain_text,
time=self.data.time,
))
await session.commit()
if result == Result.Ban:
if self.role not in {"superuser", "admin", "owner"}:
return [random.choice(NO_PERMISSION_WORDS)]
if self.reply:
ban_result = await self._ban(message_id=self.reply.message_id)
else:
ban_result = await self._ban()
if ban_result:
return [random.choice(SORRY_WORDS)]
else:
return [random.choice(DOUBT_WORDS)]
elif result in [Result.SetEnable, Result.SetDisable]:
if self.role not in {"superuser", "admin", "owner"}:
return [random.choice(NO_PERMISSION_WORDS)]
self.config.update(enable=(result == Result.SetEnable))
config_manager.config.group_config[self.data.group_id] = self.config
config_manager.save()
log_info(
"群聊学习",
f'群<m>{self.data.group_id}</m>{"开启" if result == Result.SetEnable else "关闭"}学习功能',
)
return [random.choice(ENABLE_WORDS if result == Result.SetEnable else DISABLE_WORDS)]
elif result == Result.Pass:
return None
elif result == Result.Repeat:
async with get_session() as session:
already_result = await session.execute(
select(ChatMessage).where(
ChatMessage.group_id == self.data.group_id,
ChatMessage.time >= self.data.time - 3600,
ChatMessage.user_id == self.bot_id,
ChatMessage.message == self.data.message,
).limit(self.config.repeat_threshold + 5)
)
if already_result.scalar_one_or_none():
log_info("群聊学习", "➤➤已经复读过了,跳过")
return None
msgs_result = await session.execute(
select(ChatMessage)
.where(
ChatMessage.group_id == self.data.group_id,
ChatMessage.time >= self.data.time - 3600,
)
.order_by(ChatMessage.time.desc())
.limit(self.config.repeat_threshold)
)
messages = [_db_to_message_data(m) for m in msgs_result.scalars().all()]
if not messages:
return None
if (
len(messages) >= self.config.repeat_threshold
and all(message.message == self.data.message for message in messages)
and any(message.user_id != self.data.user_id for message in messages)
):
if random.random() < self.config.break_probability:
log_info("群聊学习", "➤➤达到复读阈值,打断复读!")
return [random.choice(BREAK_REPEAT_WORDS)]
else:
log_info("群聊学习", f"➤➤达到复读阈值,复读<m>{messages[0].message}</m>")
return [self.data.message]
return None
else:
if self.data.is_plain_text and len(self.data.plain_text) <= 1:
log_info("群聊学习", "➤➤消息过短,不回复")
return None
async with get_session() as session:
ctx_result = await session.execute(
select(ChatContext).where(ChatContext.keywords == self.data.keywords).limit(1)
)
context = ctx_result.scalar_one_or_none()
if not context:
log_info("群聊学习", "➤➤尚未有已学习的回复,不回复")
return None
if not self.to_me:
answer_choices = list(
range(
self.config.answer_threshold - len(self.config.answer_threshold_weights) + 1,
self.config.answer_threshold + 1,
)
)
answer_count_threshold = random.choices(
answer_choices, weights=self.config.answer_threshold_weights
)[0]
if len(self.data.keyword_list) == chat_config.KEYWORDS_SIZE:
answer_count_threshold -= 1
cross_group_threshold = chat_config.cross_group_threshold
else:
answer_count_threshold = 1
cross_group_threshold = 1
log_info(
"群聊学习",
f"➤➤本次回复阈值为<m>{answer_count_threshold}</m>,跨群阈值为<m>{cross_group_threshold}</m>",
)
cross_kw_result = await session.execute(
select(ChatAnswer.keywords)
.where(ChatAnswer.context_id == context.id)
.group_by(ChatAnswer.keywords)
.having(func.count(ChatAnswer.keywords) >= cross_group_threshold)
)
cross_keywords = [row[0] for row in cross_kw_result.all()]
answers_cross_result = await session.execute(
select(ChatAnswer).where(
ChatAnswer.context_id == context.id,
ChatAnswer.count >= answer_count_threshold,
ChatAnswer.keywords.in_(cross_keywords),
)
)
answers_cross = answers_cross_result.scalars().all()
answer_same_result = await session.execute(
select(ChatAnswer).where(
ChatAnswer.context_id == context.id,
ChatAnswer.count >= answer_count_threshold,
ChatAnswer.group_id == self.data.group_id,
)
)
answer_same_group = answer_same_result.scalars().all()
# 提前把所有需要的字段读出来,避免 session 关闭后访问报错
candidate_answers = []
for answer in set(answers_cross) | set(answer_same_group):
a_keywords = answer.keywords
a_group_id = answer.group_id
a_count = answer.count
a_messages = list(answer.messages)
if not await self._check_allow_answer(session, a_keywords, a_group_id, a_messages):
continue
candidate_answers.append({
"count": a_count,
"messages": a_messages,
"keywords": a_keywords,
})
if not candidate_answers:
log_info("群聊学习", "➤➤没有符合条件的候选回复")
return None
sum_count = sum(a["count"] for a in candidate_answers)
per_list = [
a["count"] / sum_count * (1 - 1 / a["count"])
for a in candidate_answers
]
per_list.append(1 - sum(per_list))
result_answer = random.choices(candidate_answers + [None], weights=per_list)[0]
if result_answer is None:
log_info("群聊学习", "➤➤但不进行回复")
return None
result_message = random.choice(result_answer["messages"])
log_info("群聊学习", f"➤➤将回复<m>{result_message}</m>")
await asyncio.sleep(random.random() + 0.5)
return [result_message]
async def _ban(self, message_id: Optional[int] = None) -> bool:
bots = get_adapter(Adapter).bots
if len(bots) == 0:
return False
bot = list(bots.values())[0]
async with get_session() as session:
if message_id:
msg_result = await session.execute(
select(ChatMessage).where(ChatMessage.message_id == message_id).limit(1)
)
db_message = msg_result.scalar_one_or_none()
if not db_message or db_message.message in ALL_WORDS:
return False
keywords = db_message.keywords
try:
await bot.delete_msg(message_id=message_id)
except ActionFailed:
log_info("群聊学习", f"待禁用消息<m>{message_id}</m>尝试撤回<r>失败</r>")
else:
lr_result = await session.execute(
select(ChatMessage)
.where(
ChatMessage.group_id == self.data.group_id,
ChatMessage.user_id == self.bot_id,
)
.order_by(ChatMessage.time.desc()).limit(1)
)
last_reply = lr_result.scalar_one_or_none()
if not last_reply or last_reply.message in ALL_WORDS:
return False
keywords = last_reply.keywords
try:
await bot.delete_msg(message_id=last_reply.message_id)
except ActionFailed:
log_info("群聊学习", f"待禁用消息<m>{last_reply.message_id}</m>尝试撤回<r>失败</r>")
bl_result = await session.execute(
select(ChatBlackList).where(ChatBlackList.keywords == keywords).limit(1)
)
ban_word = bl_result.scalar_one_or_none()
if ban_word:
if self.data.group_id not in ban_word.ban_group_id:
ban_word.ban_group_id = ban_word.ban_group_id + [self.data.group_id]
if len(ban_word.ban_group_id) >= 2:
ban_word.global_ban = True
log_info("群聊学习", f"学习词<m>{keywords}</m>将被全局禁用")
await session.execute(delete(ChatAnswer).where(ChatAnswer.keywords == keywords))
else:
log_info("群聊学习", f"群<m>{self.data.group_id}</m>禁用了学习词<m>{keywords}</m>")
await session.execute(
delete(ChatAnswer).where(
ChatAnswer.keywords == keywords,
ChatAnswer.group_id == self.data.group_id,
)
)
else:
log_info("群聊学习", f"群<m>{self.data.group_id}</m>禁用了学习词<m>{keywords}</m>")
ban_word = ChatBlackList(keywords=keywords, ban_group_id=[self.data.group_id])
session.add(ban_word)
await session.execute(
delete(ChatAnswer).where(
ChatAnswer.keywords == keywords,
ChatAnswer.group_id == self.data.group_id,
)
)
await session.execute(delete(ChatContext).where(ChatContext.keywords == keywords))
await session.commit()
return True
@staticmethod
async def add_ban(data: Union[ChatMessage, ChatContext, ChatAnswer]):
# 提前读出需要的字段
keywords = data.keywords
group_id = getattr(data, "group_id", None)
is_message = isinstance(data, ChatMessage)
async with get_session() as session:
bl_result = await session.execute(
select(ChatBlackList).where(ChatBlackList.keywords == keywords).limit(1)
)
ban_word = bl_result.scalar_one_or_none()
if ban_word:
if is_message and group_id:
if group_id not in ban_word.ban_group_id:
ban_word.ban_group_id = ban_word.ban_group_id + [group_id]
if len(ban_word.ban_group_id) >= 2:
ban_word.global_ban = True
log_info("群聊学习", f"学习词<m>{keywords}</m>将被全局禁用")
await session.execute(delete(ChatAnswer).where(ChatAnswer.keywords == keywords))
else:
log_info("群聊学习", f"群<m>{group_id}</m>禁用了学习词<m>{keywords}</m>")
await session.execute(
delete(ChatAnswer).where(
ChatAnswer.keywords == keywords,
ChatAnswer.group_id == group_id,
)
)
else:
ban_word.global_ban = True
log_info("群聊学习", f"学习词<m>{keywords}</m>将被全局禁用")
await session.execute(delete(ChatAnswer).where(ChatAnswer.keywords == keywords))
else:
if is_message and group_id:
log_info("群聊学习", f"群<m>{group_id}</m>禁用了学习词<m>{keywords}</m>")
ban_word = ChatBlackList(keywords=keywords, ban_group_id=[group_id])
await session.execute(
delete(ChatAnswer).where(
ChatAnswer.keywords == keywords,
ChatAnswer.group_id == group_id,
)
)
else:
log_info("群聊学习", f"学习词<m>{keywords}</m>将被全局禁用")
ban_word = ChatBlackList(keywords=keywords, global_ban=True)
await session.execute(delete(ChatAnswer).where(ChatAnswer.keywords == keywords))
session.add(ban_word)
await session.execute(delete(ChatContext).where(ChatContext.keywords == keywords))
await session.commit()
@staticmethod
async def speak(self_id: int) -> Optional[Tuple[int, List[Union[str, MessageSegment]]]]:
cur_time = int(time.time())
today_time = time.mktime(datetime.date.today().timetuple())
async with get_session() as session:
groups_result = await session.execute(
select(ChatMessage.group_id, func.count(ChatMessage.id).label("count"))
.where(ChatMessage.time >= today_time)
.group_by(ChatMessage.group_id)
.having(func.count(ChatMessage.id) >= 10)
)
groups = [row[0] for row in groups_result.all()]
if not groups:
return None
total_messages = {}
for group_id in groups:
msgs_result = await session.execute(
select(ChatMessage)
.where(ChatMessage.group_id == group_id, ChatMessage.time >= today_time)
.order_by(ChatMessage.time.desc())
)
msgs = [_db_to_message_data(m) for m in msgs_result.scalars().all()]
if msgs:
total_messages[group_id] = msgs
if not total_messages:
return None
def group_popularity_cmp(left_group, right_group):
def cmp(a, b):
return (a > b) - (a < b)
left_group_id, left_messages = left_group
right_group_id, right_messages = right_group
left_duration = left_messages[0].time - left_messages[-1].time
right_duration = right_messages[0].time - right_messages[-1].time
return cmp(
len(left_messages) / left_duration if left_duration else 0,
len(right_messages) / right_duration if right_duration else 0,
)
popularity = sorted(total_messages.items(), key=cmp_to_key(group_popularity_cmp), reverse=True)
log_info("群聊学习", f'主动发言:群热度排行<m>{">>".join([str(g[0]) for g in popularity])}</m>')
for group_id, messages in popularity:
if len(messages) < 30:
log_info("群聊学习", f"主动发言:群<m>{group_id}</m>消息小于30条,不发言")
continue
config = config_manager.get_group_config(group_id)
ban_words = set(
chat_config.ban_words + config.ban_words + [
"[CQ:xml", "[CQ:json", "[CQ:at", "[CQ:video", "[CQ:record", "[CQ:share",
]
)
if not config.speak_enable or not config.enable:
log_info("群聊学习", f"主动发言:群<m>{group_id}</m>未开启,不发言")
continue
lr_result = await session.execute(
select(ChatMessage)
.where(ChatMessage.group_id == group_id, ChatMessage.user_id == self_id)
.order_by(ChatMessage.time.desc()).limit(1)
)
last_reply = lr_result.scalar_one_or_none()
if last_reply:
last_reply_time = last_reply.time
if last_reply_time >= messages[0].time:
log_info("群聊学习", f"主动发言:群<m>{group_id}</m>最后一条消息是{NICKNAME}发的,不发言")
continue
elif cur_time - last_reply_time < config.speak_min_interval:
log_info("群聊学习", f"主动发言:群<m>{group_id}</m>上次主动发言时间小于主动发言最小间隔,不发言")
continue
avg_interval = (messages[0].time - messages[-1].time) / len(messages)
silent_time = cur_time - messages[0].time
threshold = avg_interval * config.speak_threshold
if silent_time < threshold:
log_info("群聊学习", f"主动发言:群<m>{group_id}</m>已沉默时间({silent_time})小于阈值({int(threshold)}),不发言")
continue
ctx_result = await session.execute(
select(ChatContext).where(ChatContext.count >= config.answer_threshold)
)
contexts = ctx_result.scalars().all()
if not contexts:
continue
speak_list = []
contexts_list = list(contexts)
random.shuffle(contexts_list)
for context in contexts_list:
if (
not speak_list or random.random() < config.speak_continuously_probability
) and len(speak_list) < config.speak_continuously_max_len:
ans_result = await session.execute(
select(ChatAnswer).where(
ChatAnswer.context_id == context.id,
ChatAnswer.group_id == group_id,
ChatAnswer.count >= config.answer_threshold,
)
)
answers = ans_result.scalars().all()
if answers:
answer_data = [(a.count, a.time, list(a.messages), a.keywords) for a in answers]
answer = random.choices(
answer_data,
weights=[a[0] + 1 if a[1] >= today_time else a[0] for a in answer_data],
)[0]
message = random.choice(answer[2])
if len(message) < 2:
continue
if message.startswith("&#91;") and message.endswith("&#93;"):
continue
if any(word in message for word in ban_words):
continue
speak_list.append(message)
follow_keywords = answer[3]
while (
random.random() < config.speak_continuously_probability
and len(speak_list) < config.speak_continuously_max_len
):
fc_result = await session.execute(
select(ChatContext).where(ChatContext.keywords == follow_keywords).limit(1)
)
follow_context = fc_result.scalar_one_or_none()
if follow_context:
fa_result = await session.execute(
select(ChatAnswer).where(
ChatAnswer.group_id == group_id,
ChatAnswer.context_id == follow_context.id,
ChatAnswer.count >= config.answer_threshold,
)
)
follow_answers = fa_result.scalars().all()
if follow_answers:
fa_data = [(a.count, a.time, list(a.messages), a.keywords) for a in follow_answers]
follow_answer = random.choices(
fa_data,
weights=[a[0] + 1 if a[1] >= today_time else a[0] for a in fa_data],
)[0]
message = random.choice(follow_answer[2])
follow_keywords = follow_answer[3]
if len(message) < 2:
continue
if message.startswith("&#91;") and message.endswith("&#93;"):
continue
if all(word not in message for word in ban_words):
speak_list.append(message)
else:
break
else:
break
else:
break
if speak_list:
if random.random() < config.speak_poke_probability:
last_speak_users = {m.user_id for m in messages[:5] if m.user_id != self_id}
if last_speak_users:
select_user = random.choice(list(last_speak_users))
speak_list.append(MessageSegment("poke", {"qq": select_user}))
# 手动戳
# await send_poke(group_id=group_id, user_id=select_user)
return group_id, speak_list
else:
log_info("群聊学习", f"主动发言:群<m>{group_id}</m>没有找到符合条件的发言,不发言")
log_info("群聊学习", "主动发言:没有符合条件的群,不主动发言")
return None
async def _set_answer(self, session, message: MessageData):
"""在已有 session 内设置回复,不新开 session"""
ctx_result = await session.execute(
select(ChatContext).where(ChatContext.keywords == message.keywords).limit(1)
)
context = ctx_result.scalar_one_or_none()
if context:
if context.count < chat_config.learn_max_count:
context.count += 1
context.time = self.data.time
ans_result = await session.execute(
select(ChatAnswer).where(
ChatAnswer.keywords == self.data.keywords,
ChatAnswer.group_id == self.data.group_id,
ChatAnswer.context_id == context.id,
).limit(1)
)
answer = ans_result.scalar_one_or_none()
if answer:
if answer.count < chat_config.learn_max_count:
answer.count += 1
answer.time = self.data.time
if self.data.message not in answer.messages:
answer.messages = answer.messages + [self.data.message]
answer_count = answer.count
else:
answer = ChatAnswer(
keywords=self.data.keywords,
group_id=self.data.group_id,
time=self.data.time,
context=context,
messages=[self.data.message],
)
session.add(answer)
answer_count = 1
else:
context = ChatContext(keywords=message.keywords, time=self.data.time)
session.add(context)
await session.flush()
answer = ChatAnswer(
keywords=self.data.keywords,
group_id=self.data.group_id,
time=self.data.time,
context=context,
messages=[self.data.message],
)
session.add(answer)
answer_count = 1
await session.commit()
log_info(
"群聊学习",
f"➤将被学习为<m>{message.message}</m>的回答,已学次数为<m>{answer_count}</m>",
)
async def _check_allow(self, session, message: MessageData) -> bool:
"""检查 MessageData 是否合法"""
raw_message = message.message
if any(
i in raw_message
for i in {"[CQ:xml", "[CQ:json", "[CQ:at", "[CQ:video", "[CQ:record", "[CQ:share"}
):
return False
if any(i in raw_message for i in self.ban_words):
return False
if raw_message.startswith("&#91;") and raw_message.endswith("&#93;"):
return False
bl_result = await session.execute(
select(ChatBlackList).where(ChatBlackList.keywords == message.keywords).limit(1)
)
ban_word = bl_result.scalar_one_or_none()
if ban_word:
if ban_word.global_ban or message.group_id in ban_word.ban_group_id:
return False
return True
async def _check_allow_answer(self, session, keywords: str, group_id: int, messages: List[str]) -> bool:
"""检查 ChatAnswer 内容是否合法(使用原始字段而非 ORM 对象)"""
if not messages:
return False
raw_message = messages[0]
if any(
i in raw_message
for i in {"[CQ:xml", "[CQ:json", "[CQ:at", "[CQ:video", "[CQ:record", "[CQ:share"}
):
return False
if any(i in raw_message for i in self.ban_words):
return False
if raw_message.startswith("&#91;") and raw_message.endswith("&#93;"):
return False
bl_result = await session.execute(
select(ChatBlackList).where(ChatBlackList.keywords == keywords).limit(1)
)
ban_word = bl_result.scalar_one_or_none()
if ban_word:
if ban_word.global_ban or group_id in ban_word.ban_group_id:
return False
return True