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"➤该群{self.data.group_id}未开启群聊学习,跳过") 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"➤发言人{self.data.user_id}在屏蔽列表中,跳过") 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'群{self.data.group_id}{"开启" 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"➤➤达到复读阈值,复读{messages[0].message}") 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"➤➤本次回复阈值为{answer_count_threshold},跨群阈值为{cross_group_threshold}", ) 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"➤➤将回复{result_message}") 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"待禁用消息{message_id}尝试撤回失败") 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"待禁用消息{last_reply.message_id}尝试撤回失败") 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"学习词{keywords}将被全局禁用") await session.execute(delete(ChatAnswer).where(ChatAnswer.keywords == keywords)) else: log_info("群聊学习", f"群{self.data.group_id}禁用了学习词{keywords}") await session.execute( delete(ChatAnswer).where( ChatAnswer.keywords == keywords, ChatAnswer.group_id == self.data.group_id, ) ) else: log_info("群聊学习", f"群{self.data.group_id}禁用了学习词{keywords}") 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"学习词{keywords}将被全局禁用") await session.execute(delete(ChatAnswer).where(ChatAnswer.keywords == keywords)) else: log_info("群聊学习", f"群{group_id}禁用了学习词{keywords}") await session.execute( delete(ChatAnswer).where( ChatAnswer.keywords == keywords, ChatAnswer.group_id == group_id, ) ) else: ban_word.global_ban = True log_info("群聊学习", f"学习词{keywords}将被全局禁用") await session.execute(delete(ChatAnswer).where(ChatAnswer.keywords == keywords)) else: if is_message and group_id: log_info("群聊学习", f"群{group_id}禁用了学习词{keywords}") 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"学习词{keywords}将被全局禁用") 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'主动发言:群热度排行{">>".join([str(g[0]) for g in popularity])}') for group_id, messages in popularity: if len(messages) < 30: log_info("群聊学习", f"主动发言:群{group_id}消息小于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"主动发言:群{group_id}未开启,不发言") 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"主动发言:群{group_id}最后一条消息是{NICKNAME}发的,不发言") continue elif cur_time - last_reply_time < config.speak_min_interval: log_info("群聊学习", f"主动发言:群{group_id}上次主动发言时间小于主动发言最小间隔,不发言") 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"主动发言:群{group_id}已沉默时间({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"主动发言:群{group_id}没有找到符合条件的发言,不发言") 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"➤将被学习为{message.message}的回答,已学次数为{answer_count}", ) 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