780 lines
36 KiB
Python
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("[") and message.endswith("]"):
|
|
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("[") and message.endswith("]"):
|
|
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("[") and raw_message.endswith("]"):
|
|
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("[") and raw_message.endswith("]"):
|
|
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 |