diff --git a/make_dataset/qa_generator.py b/make_dataset/qa_generator.py index 2e0efb5..184d4cb 100644 --- a/make_dataset/qa_generator.py +++ b/make_dataset/qa_generator.py @@ -21,6 +21,18 @@ class DataProcessor: self.data = None self.processed_data = 1 self.csv_folder = "./data/csv" + self.type_list = [ + "文本", + "图片", + "卡片式链接", + "合并转发的聊天记录", + "视频", + "语言", + "未知", + "分享的小程序", + ] + self.skip_type_list = self.type_list.copy() + self.skip_type_list.remove("文本") # 根据self.config.make_dataset_args.conversation_strategy 判断初始化哪一个策略类 if self.config["conversation_strategy"] == "time_window": self.conversation_strategy = TimeWindowStrategy( @@ -47,7 +59,7 @@ class DataProcessor: for csv_file in csv_files: chat_messages = self.load_csv(csv_file) # 第一次预处理后 将chat_message 加入rcsv_df_list - message_list.append(self.group_consecutive_messages(chat_messages)) + message_list.append(self.group_consecutive_messages(messages=chat_messages)) # self.process_by_msgtype(chat_message) def group_consecutive_messages( @@ -65,68 +77,24 @@ class DataProcessor: if not messages: return [] - grouped_messages = [] - current_group = [messages[0]] + def _combine_messages(messages: List[ChatMessage]) -> ChatMessage: + """ + 合并多条消息为一条 - for i in range(1, len(messages)): - current_msg = messages[i] - last_msg = current_group[-1] + Args: + messages: 要合并的消息列表 - # 判断是否是同一个人的连续消息 - if ( - current_msg.is_sender == last_msg.is_sender - and current_msg.talker == last_msg.talker - and (current_msg.CreateTime - last_msg.CreateTime).total_seconds() - < self.config.get("message_time_window", 3600) - ): + Returns: + ChatMessage: 合并后的消息 + """ + base_msg = messages[0] + combined_content = messages[0].msg.strip() - # 同一个人的连续消息,添加到当前组 - current_group.append(current_msg) - else: - # 不是同一个人的消息,处理当前组并开始新组 - if len(current_group) > 1: - # 合并消息内容 - combined_msg = self._combine_messages(current_group) - grouped_messages.append(combined_msg) - else: - # 只有一条消息,直接添加 - grouped_messages.append(current_group[0]) + for i in messages[1:]: + content = i.msg.strip() + if not content: + continue - # 开始新组 - current_group = [current_msg] - - # 处理最后一组消息 - if current_group: - if len(current_group) > 1: - combined_msg = self._combine_messages(current_group) - grouped_messages.append(combined_msg) - else: - grouped_messages.append(current_group[0]) - - return grouped_messages - - def _combine_messages(self, messages: List[ChatMessage]) -> ChatMessage: - """ - 合并多条消息为一条 - - Args: - messages: 要合并的消息列表 - - Returns: - ChatMessage: 合并后的消息 - """ - # 以第一条消息为基础 - base_msg = messages[0] - - # 合并消息内容 - combined_content = "" - for i, msg in enumerate(messages): - content = msg.msg.strip() - if not content: - continue - - if i > 0: - # 如果前一条消息没有以标点符号结尾,添加一个句号 if combined_content and combined_content[-1] not in [ "。", "!", @@ -137,32 +105,64 @@ class DataProcessor: ]: combined_content += "," - combined_content += content + combined_content += content - # 确保最后一条消息以句号结尾 - if combined_content and combined_content[-1] not in [ - "。", - "!", - "?", - "…", - ".", - ]: - combined_content += "。" + combined_message = ChatMessage( + id=base_msg.id, + MsgSvrID=base_msg.MsgSvrID, + type_name=base_msg.type_name, + is_sender=base_msg.is_sender, + talker=base_msg.talker, + room_name=base_msg.room_name, + msg=combined_content, + src=base_msg.src, + CreateTime=messages[-1].CreateTime, # 使用最后一条消息的时间 + ) - # 创建新的合并消息 - combined_message = ChatMessage( - id=base_msg.id, - MsgSvrID=base_msg.MsgSvrID, - type_name=base_msg.type_name, - is_sender=base_msg.is_sender, - talker=base_msg.talker, - room_name=base_msg.room_name, - msg=combined_content, - src=base_msg.src, - CreateTime=messages[-1].CreateTime, # 使用最后一条消息的时间 - ) + return combined_message - return combined_message + grouped_messages = [] + current_group = [] + + for _, current_msg in enumerate(messages): + + if current_msg.type_name in self.skip_type_list: + continue + + if not current_group: + current_group = [current_msg] + continue + + last_msg = current_group[-1] + + # 判断是否是同一个人的连续消息 + if ( + current_msg.is_sender == last_msg.is_sender + and current_msg.talker == last_msg.talker + and (current_msg.CreateTime - last_msg.CreateTime).total_seconds() + < 3600 + ): + current_group.append(current_msg) + else: + # 不是同一个人的消息,处理当前组并开始新组 + if len(current_group) > 1: + combined_msg = _combine_messages(current_group) + grouped_messages.append(combined_msg) + else: + grouped_messages.append(current_group[0]) + + # 开始新组 + current_group = [current_msg] + + # 处理最后一组消息 + if current_group: + if len(current_group) > 1: + combined_msg = _combine_messages(current_group) + grouped_messages.append(combined_msg) + else: + grouped_messages.append(current_group[0]) + + return grouped_messages def create_conversation_data(self, messages: List[ChatMessage]) -> dict: """ @@ -180,8 +180,8 @@ class DataProcessor: conversation.append( { "role": "user" if msg.is_sender == 0 else "assistant", - "content": msg.content, - "timestamp": msg.timestamp, + "content": msg.msg, + "timestamp": msg.CreateTime, } ) @@ -195,8 +195,8 @@ class DataProcessor: return { "conversation": conversation, "metadata": { - "conversation_id": messages[0].conversation_id if messages else None, - "timestamp": messages[-1].timestamp if messages else None, + "conversation_id": str(messages[0].id) if messages else None, + "timestamp": messages[-1].CreateTime if messages else None, }, } diff --git a/settings.json b/settings.json index b37a363..ab419f6 100644 --- a/settings.json +++ b/settings.json @@ -49,7 +49,6 @@ "图片" ], "history_length": 10, - "prefer_comma": false, //聊天时是不是喜欢使用逗号, "conversation_strategy": "time_window", // 基于时间窗口的判断策略 "time_window": 10, // 时间窗口(分钟), "prompt_with_history": false // 是否在prompt中包含历史对话 diff --git a/tests/test_qa_generator.py b/tests/test_qa_generator.py new file mode 100644 index 0000000..dfc573e --- /dev/null +++ b/tests/test_qa_generator.py @@ -0,0 +1,315 @@ +import sys +import os +import pytest +from datetime import datetime, timedelta + +# 添加项目根目录到sys.path +current_dir = os.path.dirname(os.path.abspath(__file__)) +root_dir = os.path.dirname(current_dir) +sys.path.append(root_dir) + +from make_dataset.models import ChatMessage +from make_dataset.qa_generator import DataProcessor + +# 将当前工作目录更改为项目根目录 +os.chdir(root_dir) + +# # 测试数据处理器类的初始化和配置加载 +# def test_data_processor_init(): +# """测试DataProcessor初始化""" +# processor = DataProcessor() +# assert processor.csv_folder == "./data/csv" +# assert "文本" not in processor.skip_type_list +# assert len(processor.type_list) == 8 + + +class MockDataProcessor(DataProcessor): + def __init__(self): + super().__init__() + + +@pytest.fixture +def processor(): + """创建一个测试用的处理器实例""" + return MockDataProcessor() + + +def test_empty_messages(processor): + """测试空消息列表的情况""" + messages = [] + result = processor.group_consecutive_messages(messages) + assert result == [] + + +def test_single_message(processor): + """测试单条消息的情况""" + now = datetime.now() + message = ChatMessage( + id=1, + MsgSvrID=1001, + type_name="文本", + is_sender=0, + talker="user1", + room_name="testroom", + msg="你好", + src="", + CreateTime=now, + ) + + result = processor.group_consecutive_messages([message]) + assert len(result) == 1 + assert result[0].msg == "你好" + + +def test_consecutive_messages_same_sender(processor): + """测试同一发送者的连续消息""" + now = datetime.now() + messages = [ + ChatMessage( + id=1, + MsgSvrID=1001, + type_name="文本", + is_sender=0, + talker="user1", + room_name="testroom", + msg="你好", + src="", + CreateTime=now, + ), + ChatMessage( + id=2, + MsgSvrID=1002, + type_name="文本", + is_sender=0, + talker="user1", + room_name="testroom", + msg="最近怎么样", + src="", + CreateTime=now + timedelta(minutes=5), + ), + ChatMessage( + id=3, + MsgSvrID=1003, + type_name="文本", + is_sender=0, + talker="user1", + room_name="testroom", + msg="我想问个问题", + src="", + CreateTime=now + timedelta(minutes=10), + ), + ] + + result = processor.group_consecutive_messages(messages) + assert len(result) == 1 + assert result[0].msg == "你好,最近怎么样,我想问个问题" + + +def test_messages_different_senders(processor): + """测试不同发送者的消息""" + now = datetime.now() + messages = [ + ChatMessage( + id=1, + MsgSvrID=1001, + type_name="文本", + is_sender=0, + talker="user1", + room_name="testroom", + msg="你好", + src="", + CreateTime=now, + ), + ChatMessage( + id=2, + MsgSvrID=1002, + type_name="文本", + is_sender=1, + talker="user2", + room_name="testroom", + msg="你好,有什么可以帮你的", + src="", + CreateTime=now + timedelta(minutes=5), + ), + ChatMessage( + id=3, + MsgSvrID=1003, + type_name="文本", + is_sender=0, + talker="user1", + room_name="testroom", + msg="我想问个问题", + src="", + CreateTime=now + timedelta(minutes=10), + ), + ] + + result = processor.group_consecutive_messages(messages) + assert len(result) == 3 + assert result[0].msg == "你好" + assert result[1].msg == "你好,有什么可以帮你的" + assert result[2].msg == "我想问个问题" + + +def test_skip_non_text_messages(processor): + """测试跳过非文本消息""" + now = datetime.now() + messages = [ + ChatMessage( + id=1, + MsgSvrID=1001, + type_name="文本", + is_sender=0, + talker="user1", + room_name="testroom", + msg="你好", + src="", + CreateTime=now, + ), + ChatMessage( + id=2, + MsgSvrID=1002, + type_name="图片", + is_sender=0, + talker="user1", + room_name="testroom", + msg="", + src="image.jpg", + CreateTime=now + timedelta(minutes=1), + ), + ChatMessage( + id=3, + MsgSvrID=1003, + type_name="文本", + is_sender=0, + talker="user1", + room_name="testroom", + msg="看到图片了吗", + src="", + CreateTime=now + timedelta(minutes=2), + ), + ] + + result = processor.group_consecutive_messages(messages) + assert len(result) == 1 + assert result[0].msg == "你好,看到图片了吗" + + +def test_time_window_limit(processor): + """测试时间窗口限制(超过1小时的消息不会合并)""" + now = datetime.now() + messages = [ + ChatMessage( + id=1, + MsgSvrID=1001, + type_name="文本", + is_sender=0, + talker="user1", + room_name="testroom", + msg="你好", + src="", + CreateTime=now, + ), + ChatMessage( + id=2, + MsgSvrID=1002, + type_name="文本", + is_sender=0, + talker="user1", + room_name="testroom", + msg="晚上好", + src="", + CreateTime=now + timedelta(hours=2), # 超过1小时 + ), + ] + + result = processor.group_consecutive_messages(messages) + assert len(result) == 2 + assert result[0].msg == "你好" + assert result[1].msg == "晚上好" + + +if __name__ == "__main__": + import pandas as pd + import os + from datetime import datetime + + # 创建一个测试函数,将group_consecutive_messages的结果转为DataFrame并保存为CSV + def test_save_grouped_messages_to_csv(processor): + """测试将group_consecutive_messages的结果转换为DataFrame并保存为CSV""" + now = datetime.now() + messages = [ + ChatMessage( + id=1, + MsgSvrID=1001, + type_name="文本", + is_sender=0, + talker="user1", + room_name="testroom", + msg="你好", + src="", + CreateTime=now, + ), + ChatMessage( + id=2, + MsgSvrID=1002, + type_name="文本", + is_sender=0, + talker="user1", + room_name="testroom", + msg="这是测试消息", + src="", + CreateTime=now + timedelta(minutes=10), + ), + ChatMessage( + id=3, + MsgSvrID=1003, + type_name="文本", + is_sender=1, + talker="user2", + room_name="testroom", + msg="收到了", + src="", + CreateTime=now + timedelta(minutes=20), + ), + ] + + # 获取分组后的消息 + grouped_messages = processor.group_consecutive_messages(messages) + + # 将ChatMessage对象转换为字典列表 + messages_dict = [] + for msg in grouped_messages: + messages_dict.append({ + "id": msg.id, + "MsgSvrID": msg.MsgSvrID, + "type_name": msg.type_name, + "is_sender": msg.is_sender, + "talker": msg.talker, + "room_name": msg.room_name, + "msg": msg.msg, + "src": msg.src, + "CreateTime": msg.CreateTime, + }) + + # 创建DataFrame + df = pd.DataFrame(messages_dict) + + # 确保输出目录存在 + output_dir = "./test_output" + os.makedirs(output_dir, exist_ok=True) + + # 保存为CSV文件 + output_file = os.path.join(output_dir, f"grouped_messages_{now.strftime('%Y%m%d_%H%M%S')}.csv") + df.to_csv(output_file, index=False, encoding="utf-8") + + # 验证文件已创建 + assert os.path.exists(output_file) + + # 读取CSV文件并验证内容 + loaded_df = pd.read_csv(output_file) + assert len(loaded_df) == len(grouped_messages) + assert loaded_df.iloc[0]["msg"] == "你好,这是测试消息" # 验证第一条消息已合并 + + print(f"已成功保存分组消息到: {output_file}") + return output_file