AI 超级智能体第三期 扩展作业
七、扩展思路
1)自定义 Advisor,比如权限校验、违禁词校验 Advisor
违禁词检测
违禁词来源
https://gitee.com/crazypoo/badwords
违禁词存放位置
使用 ClassPathResource 来加载违禁词文件,这意味着文件必须位于类路径下。如果文件不在类路径下(例如在项目的外部目录),就会抛出 FileNotFoundException,我就放到了resources下了。
自定义违禁词Advisor
▼java复制代码package com.lihui.aiagent.advisor; import lombok.extern.slf4j.Slf4j; import org.springframework.ai.chat.client.advisor.api.*; import org.springframework.core.io.ClassPathResource; import org.springframework.util.StringUtils; import reactor.core.publisher.Flux; import java.io.BufferedReader; import java.io.InputStreamReader; import java.nio.charset.StandardCharsets; import java.util.ArrayList; import java.util.List; import java.util.stream.Collectors; /** * 违禁词校验 Advisor * 检查用户输入是否包含违禁词 * 违禁词来源 https://gitee.com/crazypoo/badwords */ @Slf4j public class ProhibitedWordAdvisor implements CallAroundAdvisor, StreamAroundAdvisor { private static final String DEFAULT_PROHIBITED_WORDS_FILE = "prohibited/prohibited-words.txt"; private final List<String> prohibitedWords; /** * 创建默认违禁词Advisor,从默认文件读取违禁词列表 */ public ProhibitedWordAdvisor() { this.prohibitedWords = loadProhibitedWordsFromFile(DEFAULT_PROHIBITED_WORDS_FILE); log.info("初始化违禁词Advisor,违禁词数量: {}", prohibitedWords.size()); } /** * 创建违禁词Advisor,从指定文件读取违禁词列表 */ public ProhibitedWordAdvisor(String prohibitedWordsFile) { this.prohibitedWords = loadProhibitedWordsFromFile(prohibitedWordsFile); log.info("初始化违禁词Advisor,违禁词数量: {}", prohibitedWords.size()); } /** * 从文件加载违禁词列表 */ private List<String> loadProhibitedWordsFromFile(String filePath) { try { var resource = new ClassPathResource(filePath); var reader = new BufferedReader( new InputStreamReader(resource.getInputStream(), StandardCharsets.UTF_8)); List<String> words = reader.lines() .filter(StringUtils::hasText) .map(String::trim) .collect(Collectors.toList()); log.info("从文件 {} 加载违禁词 {} 个", filePath, words.size()); return words; } catch (Exception e) { log.error("加载违禁词文件 {} 失败", filePath, e); return new ArrayList<>(); } } @Override public String getName() { return this.getClass().getSimpleName(); } @Override public int getOrder() { return -100; // 确保在其他Advisor之前执行 } /** * 检查请求中是否包含违禁词 */ private AdvisedRequest checkRequest(AdvisedRequest request) { String userText = request.userText(); if (containsProhibitedWord(userText)) { log.warn("检测到违禁词在用户输入中: {}", userText); throw new ProhibitedWordException("用户输入包含违禁词"); } return request; } /** * 检查文本中是否包含违禁词 */ private boolean containsProhibitedWord(String text) { if (!StringUtils.hasText(text)) { return false; } for (String word : prohibitedWords) { if (text.toLowerCase().contains(word.toLowerCase())) { return true; } } return false; } @Override public AdvisedResponse aroundCall(AdvisedRequest advisedRequest, CallAroundAdvisorChain chain) { return chain.nextAroundCall(checkRequest(advisedRequest)); } @Override public Flux<AdvisedResponse> aroundStream(AdvisedRequest advisedRequest, StreamAroundAdvisorChain chain) { return chain.nextAroundStream(checkRequest(advisedRequest)); } /** * 违禁词异常 */ public static class ProhibitedWordException extends RuntimeException { public ProhibitedWordException(String message) { super(message); } } }
实现
▼text复制代码package com.lihui.aiagent.app; import com.lihui.aiagent.advisor.ProhibitedWordAdvisor; import com.lihui.aiagent.chatmemory.FileBasedChatMemory; import lombok.extern.slf4j.Slf4j; import org.springframework.ai.chat.client.ChatClient; import org.springframework.ai.chat.client.advisor.MessageChatMemoryAdvisor; import org.springframework.ai.chat.memory.ChatMemory; import org.springframework.ai.chat.model.ChatModel; import org.springframework.ai.chat.model.ChatResponse; import org.springframework.stereotype.Component; import static org.springframework.ai.chat.client.advisor.AbstractChatMemoryAdvisor.CHAT_MEMORY_CONVERSATION_ID_KEY; import static org.springframework.ai.chat.client.advisor.AbstractChatMemoryAdvisor.CHAT_MEMORY_RETRIEVE_SIZE_KEY; @Component @Slf4j public class LoveAppProhibitedWords { private final ChatClient chatClient; private static final String SYSTEM_PROMPT = "扮演深耕恋爱心理领域的专家。开场向用户表明身份,告知用户可倾诉恋爱难题。" + "围绕单身、恋爱、已婚三种状态提问:单身状态询问社交圈拓展及追求心仪对象的困扰;" + "恋爱状态询问沟通、习惯差异引发的矛盾;已婚状态询问家庭责任与亲属关系处理的问题。" + "引导用户详述事情经过、对方反应及自身想法,以便给出专属解决方案。"; /** * 初始化 ChatClient * * @param dashscopeChatModel */ public LoveAppProhibitedWords(ChatModel dashscopeChatModel) { // // 初始化基于文件的对话记忆 String fileDir = System.getProperty("user.dir") + "/tmp/chat-memory"; ChatMemory chatMemory = new FileBasedChatMemory(fileDir); // 初始化基于内存的对话记忆 // ChatMemory chatMemory = new InMemoryChatMemory(); chatClient = ChatClient.builder(dashscopeChatModel) .defaultSystem(SYSTEM_PROMPT) .defaultAdvisors( new MessageChatMemoryAdvisor(chatMemory), // 自定义日志 Advisor,可按需开启 new ProhibitedWordAdvisor() // // 自定义推理增强 Advisor,可按需开启 // ,new ReReadingAdvisor() ) .build(); } /** * AI 基础对话(支持多轮对话记忆) * * @param message * @param chatId * @return */ public String doChat(String message, String chatId) { ChatResponse chatResponse = chatClient .prompt() .user(message) .advisors(spec -> spec.param(CHAT_MEMORY_CONVERSATION_ID_KEY, chatId) .param(CHAT_MEMORY_RETRIEVE_SIZE_KEY, 10)) .call() .chatResponse(); String content = chatResponse.getResult().getOutput().getText(); log.info("content: {}", content); return content; } }
测试
▼java复制代码package com.lihui.aiagent.app; import jakarta.annotation.Resource; import org.junit.jupiter.api.Assertions; import org.junit.jupiter.api.Test; import org.springframework.boot.test.context.SpringBootTest; import java.util.UUID; @SpringBootTest class LoveAppProhibitedWordsTest { @Resource private LoveAppProhibitedWords loveAppProhibitedWords; @Test void testProhibitedwordst() { String chatId = UUID.randomUUID().toString(); // 正常输入测试 String normalMessage = "你好,我是lihui"; String answer = loveAppProhibitedWords.doChat(normalMessage, chatId); Assertions.assertNotNull(answer); // 包含违禁词输入测试,验证是否抛出异常 String prohibitedMessage = "成人偷拍"; Assertions.assertThrows( com.lihui.aiagent.advisor.ProhibitedWordAdvisor.ProhibitedWordException.class, () -> loveAppProhibitedWords.doChat(prohibitedMessage, chatId), "Expected ProhibitedWordException for prohibited message" ); // 对话记忆相关测试 String recallMessage = "我之前说的内容,你能再帮我想想吗?"; answer = loveAppProhibitedWords.doChat(recallMessage, chatId); Assertions.assertNotNull(answer); } }

2)自定义对话记忆,比如持久化对话到 MySQL 或 Redis 存储中
持久化对话到 MySQL
▼java复制代码package com.lihui.aiagent.chatmemory; import cn.hutool.json.JSONConfig; import cn.hutool.json.JSONUtil; import lombok.extern.slf4j.Slf4j; import org.springframework.ai.chat.memory.ChatMemory; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.SystemMessage; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.jdbc.core.JdbcTemplate; import org.springframework.stereotype.Component; import org.springframework.transaction.annotation.Transactional; import javax.sql.DataSource; import java.sql.Timestamp; import java.time.LocalDateTime; import java.util.*; import java.util.stream.Collectors; /** * MySQL实现的对话记忆 * 将对话内容持久化到MySQL数据库 */ @Component @Slf4j public class MySQLChatMemory implements ChatMemory { private final JdbcTemplate jdbcTemplate; private final JSONConfig jsonConfig; public MySQLChatMemory(DataSource dataSource) { this.jdbcTemplate = new JdbcTemplate(dataSource); this.jsonConfig = new JSONConfig().setIgnoreNullValue(true); log.info("初始化MySQL对话记忆"); } @Override @Transactional public void add(String conversationId, Message message) { if (message != null && conversationId != null) { List<Message> messages = Collections.singletonList(message); add(conversationId, messages); } } @Override @Transactional public void add(String conversationId, List<Message> messages) { if (messages == null || messages.isEmpty() || conversationId == null) { return; } // 获取当前最大序号 Integer maxOrder = getMaxOrder(conversationId).orElse(0); int nextOrder = maxOrder + 1; // 使用批处理提高效率 String insertSql = "INSERT INTO chatmemory (conversation_id, message_order, message_type, content, message_json, create_time, update_time, is_delete) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"; log.info("添加消息到会话 {}, 消息数量: {}", conversationId, messages.size()); jdbcTemplate.batchUpdate(insertSql, messages, messages.size(), (ps, message) -> { int order = nextOrder + messages.indexOf(message); String messageJson = serializeMessage(message); String content = message.getText(); Timestamp now = Timestamp.valueOf(LocalDateTime.now()); ps.setString(1, conversationId); ps.setInt(2, order); ps.setString(3, message.getMessageType().toString()); ps.setString(4, content); ps.setString(5, messageJson); ps.setTimestamp(6, now); // create_time ps.setTimestamp(7, now); // update_time ps.setBoolean(8, false); // is_delete = 0 }); } @Override public List<Message> get(String conversationId, int lastN) { String sql; Object[] params; // 修改查询逻辑:lastN > 0 时获取前N条消息,而不是最后N条 if (lastN > 0) { sql = "SELECT message_json, message_type, content FROM chatmemory " + "WHERE conversation_id = ? AND is_delete = 0 ORDER BY message_order DESC LIMIT ?"; params = new Object[] { conversationId, lastN }; } else { sql = "SELECT message_json, message_type, content FROM chatmemory " + "WHERE conversation_id = ? AND is_delete = 0 ORDER BY message_order DESC"; params = new Object[] { conversationId }; } List<Message> messages = executeMessageQuery(sql, params); log.info("从会话 {} 中检索到 {} 条消息", conversationId, messages.size()); return messages; } @Override @Transactional public void clear(String conversationId) { // 将物理删除改为逻辑删除 String sql = "UPDATE chatmemory SET is_delete = 1, update_time = ? WHERE conversation_id = ? AND is_delete = 0"; Timestamp now = Timestamp.valueOf(LocalDateTime.now()); Object[] params = new Object[] { now, conversationId }; int count = jdbcTemplate.update(sql, params); log.info("从会话 {} 中逻辑删除 {} 条消息", conversationId, count); } /** * 获取会话中最大的消息序号 */ private Optional<Integer> getMaxOrder(String conversationId) { String sql = "SELECT MAX(message_order) FROM chatmemory WHERE conversation_id = ? AND is_delete = 0"; Integer result = jdbcTemplate.queryForObject(sql, Integer.class, conversationId); return Optional.ofNullable(result); } /** * 将消息序列化为JSON字符串 */ private String serializeMessage(Message message) { Map<String, Object> map = new HashMap<>(); map.put("type", message.getMessageType().toString()); map.put("text", message.getText()); // 添加消息类名,便于反序列化 if (message instanceof UserMessage) { map.put("messageClass", "UserMessage"); } else if (message instanceof AssistantMessage) { map.put("messageClass", "AssistantMessage"); } else if (message instanceof SystemMessage) { map.put("messageClass", "SystemMessage"); } else { map.put("messageClass", "OtherMessage"); } return JSONUtil.toJsonStr(map, jsonConfig); } /** * 从JSON字符串反序列化消息 */ private Message deserializeMessage(String messageJson, String messageType, String content) { switch (messageType) { case "USER": return new UserMessage(content); case "ASSISTANT": return new AssistantMessage(content); case "SYSTEM": return new SystemMessage(content); default: log.warn("未知的消息类型: {}", messageType); return new AssistantMessage("未知消息类型: " + content); } } /** * 执行消息查询并返回结果列表 */ private List<Message> executeMessageQuery(String sql, Object[] params) { log.info("SQL: {}, 参数: {}", sql, Arrays.toString(params)); return jdbcTemplate.query(sql, params, (rs, rowNum) -> { String messageJson = rs.getString("message_json"); String messageType = rs.getString("message_type"); String content = rs.getString("content"); return deserializeMessage(messageJson, messageType, content); }).stream() .filter(Objects::nonNull) .collect(Collectors.toList()); } }
▼text复制代码package com.lihui.aiagent.app; import com.lihui.aiagent.advisor.MyLoggerAdvisor; import com.lihui.aiagent.chatmemory.MySQLChatMemory; import com.lihui.aiagent.chatmemory.MybatisPlusChatMemory; import lombok.extern.slf4j.Slf4j; import org.springframework.ai.chat.client.ChatClient; import org.springframework.ai.chat.client.advisor.MessageChatMemoryAdvisor; import org.springframework.ai.chat.model.ChatModel; import org.springframework.ai.chat.model.ChatResponse; import org.springframework.stereotype.Component; import static org.springframework.ai.chat.client.advisor.AbstractChatMemoryAdvisor.CHAT_MEMORY_CONVERSATION_ID_KEY; import static org.springframework.ai.chat.client.advisor.AbstractChatMemoryAdvisor.CHAT_MEMORY_RETRIEVE_SIZE_KEY; @Component @Slf4j public class LoveAppMysqlPersistence { private final ChatClient chatClient; private static final String SYSTEM_PROMPT = "扮演深耕恋爱心理领域的专家。开场向用户表明身份,告知用户可倾诉恋爱难题。" + "围绕单身、恋爱、已婚三种状态提问:单身状态询问社交圈拓展及追求心仪对象的困扰;" + "恋爱状态询问沟通、习惯差异引发的矛盾;已婚状态询问家庭责任与亲属关系处理的问题。" + "引导用户详述事情经过、对方反应及自身想法,以便给出专属解决方案。"; /** * 初始化 ChatClient * * @param dashscopeChatModel */ public LoveAppMysqlPersistence(ChatModel dashscopeChatModel, MySQLChatMemory chatMemory ) { // ChatMemory chatMemory = new MySQLChatMemory; chatClient = ChatClient.builder(dashscopeChatModel) .defaultSystem(SYSTEM_PROMPT) .defaultAdvisors( new MessageChatMemoryAdvisor(chatMemory), // 自定义日志 Advisor,可按需开启 new MyLoggerAdvisor() // // 自定义推理增强 Advisor,可按需开启 // ,new ReReadingAdvisor() ) .build(); } /** * AI 基础对话(支持多轮对话记忆) * * @param message * @param chatId * @return */ public String doChat(String message, String chatId) { ChatResponse chatResponse = chatClient .prompt() .user(message) .advisors(spec -> spec.param(CHAT_MEMORY_CONVERSATION_ID_KEY, chatId) .param(CHAT_MEMORY_RETRIEVE_SIZE_KEY, 10)) .call() .chatResponse(); String content = chatResponse.getResult().getOutput().getText(); log.info("content: {}", content); return content; } }
评论
问答助学
相关内容
0个评论
全部评论
点击登录,快来和大家讨论吧~
表情
图片
暂无评论
