基于MySQL的AI对话持久化
此处先@笨蛋皮卡丘,对着这位大佬的基本实现做了一个小扩展,多存了一点信息 (目前媒体信息只支持URL类型) 原始链接
依赖
▼xml复制代码<!--JSON序列化,兼容cos使用低版本--> <dependency> <groupId>com.fasterxml.jackson.core</groupId> <artifactId>jackson-databind</artifactId> <version>2.18.3</version> </dependency>
数据表chat_message
service和mapper就对着表一键生成吧,这里就不给了
▼sql复制代码DROP TABLE IF EXISTS `chat_message`; CREATE TABLE `chat_message` ( `id` bigint UNSIGNED NOT NULL AUTO_INCREMENT COMMENT '主键ID', `conversationId` varchar(64) CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci NOT NULL COMMENT '会话ID', `messageType` varchar(20) CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci NOT NULL COMMENT '消息类型', `content` text CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci NOT NULL COMMENT '消息内容', `metadata` text CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci NOT NULL COMMENT '元数据', `createTime` datetime NOT NULL DEFAULT CURRENT_TIMESTAMP COMMENT '创建时间', `responses` longtext CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci NULL COMMENT '工具响应', `mediaJson` json CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci NULL COMMENT '媒体信息',, `toolCalls` longtext CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci NULL COMMENT '工具调用', PRIMARY KEY (`id`) USING BTREE, INDEX `idx_type_content_prefix`(`messageType` ASC, `content`(1) ASC) USING BTREE ) ENGINE = InnoDB CHARACTER SET = utf8mb4 COLLATE = utf8mb4_unicode_ci COMMENT = '消息' ROW_FORMAT = DYNAMIC;
序列化器MediaSerializerDeserializer
▼java复制代码public class MediaSerializerDeserializer { // 自定义序列化器 public static class MediaSerializer extends JsonSerializer<Media> { @Override public void serialize(Media media, JsonGenerator gen, SerializerProvider provider) throws IOException { gen.writeStartObject(); if (media.getId() != null) { gen.writeStringField("id", media.getId()); } gen.writeStringField("mimeType", media.getMimeType().toString()); gen.writeStringField("data", media.getData().toString()); // 假设 data 是 URL.toString() gen.writeStringField("name", media.getName()); gen.writeEndObject(); } } // 自定义反序列化器(从 JSON 构建 Media) public static class MediaDeserializer extends JsonDeserializer<Media> { @Override public Media deserialize(JsonParser p, DeserializationContext ctxt) throws IOException { JsonNode node = p.getCodec().readTree(p); // 1. 读取 mimeType String mimeTypeStr = node.get("mimeType").asText(); MimeType mimeType = MimeType.valueOf(mimeTypeStr); // 假设 MimeType 有 valueOf 方法 // 2. 读取 data(假设 data 是 URL 的字符串形式) String dataStr = node.get("data").asText(); URL dataUrl = new URL(dataStr); // 3. 读取 id(可选) String id = null; if (node.has("id")) { id = node.get("id").asText(); } // 4. 读取 name String name = node.get("name").asText(); // 5. 使用 Builder 构建 Media 对象 return Media.builder() .id(id) .mimeType(mimeType) .data(dataUrl) .name(name) .build(); } } // 注册序列化器和反序列化器到 ObjectMapper public static ObjectMapper createCustomObjectMapper() { ObjectMapper mapper = new ObjectMapper(); SimpleModule module = new SimpleModule(); module.addSerializer(Media.class, new MediaSerializer()); module.addDeserializer(Media.class, new MediaDeserializer()); mapper.registerModule(module); return mapper; } }
实体类ChatMessage
▼java复制代码@TableName(value ="chat_message") @Data public class ChatMessage implements Serializable { /** * 主键ID */ @TableId(value = "id", type = IdType.AUTO) private Long id; /** * 会话ID */ @TableField(value = "conversationId") private String conversationId; /** * 消息类型 */ @TableField(value = "messageType") private MessageType messageType; /** * 消息内容 */ @TableField(value = "content") private String content; /** * 元数据 */ @TableField(value = "metadata", typeHandler = JacksonTypeHandler.class) private Map<String, Object> metadata; /** * 创建时间 */ @TableField(value = "createTime") private Date createTime; /** * 工具调用结果,仅工具消息有 */ @TableField(value = "responses", typeHandler = JacksonTypeHandler.class) private List<ToolResponseMessage.ToolResponse> responses; /** * 媒体文件,仅用户消息和助手消息有 */ @TableField(value = "mediaJson") protected String mediaJson; /** * 工具调用列表,仅助手消息有 */ @TableField(value = "toolCalls", typeHandler = JacksonTypeHandler.class) private List<AssistantMessage.ToolCall> toolCalls; /** * 临时字段(List<Media> 类型,不映射到数据库) */ @TableField(exist = false) private List<Media> media; // 用于业务逻辑 @TableField(exist = false) private static final long serialVersionUID = 1L; // 序列化为json public void setMediaFromList(List<Media> mediaList) { ObjectMapper objectMapper = MediaSerializerDeserializer.createCustomObjectMapper(); try { this.mediaJson = objectMapper.writeValueAsString(mediaList); } catch (JsonProcessingException e) { throw new RuntimeException("Failed to serialize Media list to JSON", e); } } //反序列化为List<Media> public List<Media> getMediaAsList() { if (mediaJson == null) { return null; } ObjectMapper objectMapper = MediaSerializerDeserializer.createCustomObjectMapper(); try { List<Media> media1 = objectMapper.readValue(mediaJson, new TypeReference<List<Media>>() { }); return media1; } catch (JsonProcessingException e) { throw new RuntimeException("Failed to deserialize Media list from JSON", e); } } }
ChatMessage <=> Message 转换 MessageConverter
▼java复制代码public class MessageConverter { /** * 将 Message 转换为 ChatMessage */ public static ChatMessage toChatMessage(Message message, String conversationId) { ChatMessage chatMessage = new ChatMessage(); chatMessage.setConversationId(conversationId); chatMessage.setMessageType(message.getMessageType()); chatMessage.setMetadata(message.getMetadata() == null ? new HashMap<>() : message.getMetadata()); if (message instanceof UserMessage userMessage) { chatMessage.setContent(userMessage.getText()); //设置业务逻辑字段 chatMessage.setMedia(userMessage.getMedia()); //设置数据库字段 chatMessage.setMediaFromList(userMessage.getMedia() != null ? userMessage.getMedia() : null); } else if (message instanceof AssistantMessage assistantMessage) { chatMessage.setContent(assistantMessage.getText()); chatMessage.setToolCalls(assistantMessage.getToolCalls()); //设置业务逻辑字段 chatMessage.setMedia(assistantMessage.getMedia()); //设置数据库字段 chatMessage.setMediaFromList(assistantMessage.getMedia() != null ? assistantMessage.getMedia() : null); } else if (message instanceof SystemMessage systemMessage) { // 只存 content chatMessage.setContent(systemMessage.getText()); } else if (message instanceof ToolResponseMessage toolMessage) { chatMessage.setResponses(toolMessage.getResponses()); chatMessage.setContent(toolMessage.getText()); // 可选,工具消息一般无content // 不设置 toolCalls/media } return chatMessage; } /** * 将 ChatMessage 转换为 Message */ public static Message toMessage(ChatMessage chatMessage) { MessageType messageType = chatMessage.getMessageType(); String text = chatMessage.getContent(); Map<String, Object> metadata = chatMessage.getMetadata() != null ? chatMessage.getMetadata() : new HashMap<>(); return switch (messageType) { case USER -> new UserMessage(text , chatMessage.getMediaAsList(), metadata); case ASSISTANT -> new AssistantMessage(text, metadata , chatMessage.getToolCalls()!=null?chatMessage.getToolCalls(): List.of() , chatMessage.getMediaAsList()); case SYSTEM -> new SystemMessage(text); case TOOL -> new ToolResponseMessage( chatMessage.getResponses() != null ? chatMessage.getResponses() : List.of(), metadata ); default -> throw new IllegalArgumentException("Unknown message type: " + messageType); }; } }
ChatMemory 实现和使用
和大佬的一模一样(开闭原则)
▼java复制代码@Component @RequiredArgsConstructor public class DatabaseChatMemory implements ChatMemory { @Resource private ChatMessageService chatMessageService; /** * 批量添加消息到数据库 * * @param conversationId 会话ID,用于关联消息和会话 * @param messages 消息列表,包含多个消息对象 */ @Override public void add(String conversationId, List<Message> messages) { // 将Message对象转换为ChatMessage对象,并收集到列表中 List<ChatMessage> chatMessages = messages.stream() .map(message -> MessageConverter.toChatMessage(message, conversationId)) .collect(Collectors.toList()); // 批量保存消息到数据库 chatMessageService.saveBatch(chatMessages, chatMessages.size()); } /** * 获取指定会话的最近N条消息 * * @param conversationId 会话ID,用于查询特定会话的消息 * @param lastN 获取最近的消息数量,必须大于0 * @return 返回按照时间顺序排列的消息列表 */ @Override public List<Message> get(String conversationId, int lastN) { LambdaQueryWrapper<ChatMessage> queryWrapper = new LambdaQueryWrapper<>(); // 查询最近的 lastN 条消息 queryWrapper.eq(ChatMessage::getConversationId, conversationId) .orderByDesc(ChatMessage::getCreateTime) .last(lastN > 0, "LIMIT " + lastN); List<ChatMessage> chatMessages = chatMessageService.list(queryWrapper); // 按照时间顺序返回 if (!chatMessages.isEmpty()) { Collections.reverse(chatMessages); } // 将ChatMessage对象转换为Message对象,并收集到列表中 return chatMessages .stream() .map(MessageConverter::toMessage) .collect(Collectors.toList()); } /** * 清除指定会话的所有消息 * * @param conversationId 会话ID,用于删除特定会话的所有消息 */ @Override public void clear(String conversationId) { LambdaQueryWrapper<ChatMessage> queryWrapper = new LambdaQueryWrapper<>(); queryWrapper.eq(ChatMessage::getConversationId, conversationId); // 删除特定会话的所有消息 chatMessageService.remove(queryWrapper); } }
基本使用也是一样的
▼java复制代码public LoveApp(@Qualifier("dashscopeChatModel") ChatModel dashscopeChatModel, ChatMemory databaseChatMemory) { // 初始化基于内存的对话记忆 // ChatMemory chatMemory = new InMemoryChatMemory(); // 初始化基于文件的对话记忆 String fileDir = System.getProperty("user.dir") + "/chat-memory"; ChatMemory chatMemory = new FileBasedChatMemory(fileDir); // 构建ChatClient实例,配置默认的系统提示和对话记忆顾问 chatClient = ChatClient.builder(dashscopeChatModel) // 设置默认的对话记忆 .defaultAdvisors(new MessageChatMemoryAdvisor(databaseChatMemory)) .build(); }
评论
问答助学
相关内容
0个评论
全部评论
点击登录,快来和大家讨论吧~
表情
图片
暂无评论
