| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343 |
- package edu.nju.software.aipaasagent.service;
- import edu.nju.software.aipaasagent.agent.config.AgentConfiguration;
- import edu.nju.software.aipaasagent.rag.client.RagClient;
- import edu.nju.software.aipaasagent.rag.dto.RagRetrieveRequest;
- import edu.nju.software.aipaasagent.memory.chat.RedisBasedChatMemory;
- import lombok.extern.slf4j.Slf4j;
- import org.springframework.ai.chat.messages.Message;
- import org.springframework.ai.chat.messages.UserMessage;
- import org.springframework.ai.chat.model.ChatModel;
- import org.springframework.ai.chat.model.ChatResponse;
- import org.springframework.ai.chat.prompt.Prompt;
- import org.springframework.beans.factory.annotation.Autowired;
- import org.springframework.scheduling.annotation.Async;
- import org.springframework.stereotype.Service;
- import reactor.core.publisher.Mono;
- import java.util.Comparator;
- import java.util.List;
- import java.util.stream.Collectors;
- /**
- * RAG 服务
- * 封装 RAG 检索逻辑
- */
- @Slf4j
- @Service
- public class RagService {
- @Autowired
- private RagClient ragClient;
- @Autowired
- private ChatModel dashscopeModel;
- @Autowired
- private RedisBasedChatMemory redisBasedChatMemory;
- /**
- * 执行 RAG 检索
- *
- * @param query 查询词
- * @param ragConfig RAG 配置
- * @return 检索结果
- */
- public Mono<RagClient.RagRetrieveResponse> retrieve(String query, AgentConfiguration.RagConfig ragConfig) {
- if (ragConfig == null || !Boolean.TRUE.equals(ragConfig.getEnabled())) {
- log.debug("RAG 未启用,跳过检索");
- return Mono.empty();
- }
- // 转换为 RAG 客户端请求
- List<RagRetrieveRequest.VectorStoreConfig> vectorStores = ragConfig.getVectorStores().stream()
- .map(config -> RagRetrieveRequest.VectorStoreConfig.builder()
- .vectorStoreId(config.getVectorStoreId())
- .topK(config.getTopK())
- .scoreThreshold(config.getScoreThreshold())
- .build())
- .collect(Collectors.toList());
- RagRetrieveRequest request = RagRetrieveRequest.builder()
- .query(query)
- .vectorStores(vectorStores)
- .build();
- return ragClient.retrieve(request);
- }
- /**
- * 生成 RAG 检索查询词(使用本地模型)
- *
- * @param userCurrentMsg 用户当前消息
- * @param historyDialog 历史对话
- * @return 精准检索关键词
- */
- public String generateRetrievalQuery(String userCurrentMsg, List<Message> historyDialog) {
- log.info("[RAG-Query] 开始生成检索查询词");
- log.info("[RAG-Query] 用户原始消息:{}", userCurrentMsg);
- // 构建提示词
- StringBuilder promptBuilder = new StringBuilder();
- promptBuilder.append("你是一个检索查询优化助手。请根据用户当前问题和历史对话,生成一个精准的检索查询词。\n");
- promptBuilder.append("要求:\n");
- promptBuilder.append("1. 提取核心关键词\n");
- promptBuilder.append("2. 去除无关词汇\n");
- promptBuilder.append("3. 保持简洁\n");
- promptBuilder.append("4. 只返回查询词,不要有其他内容\n\n");
- if (historyDialog != null && !historyDialog.isEmpty()) {
- promptBuilder.append("历史对话:\n");
- for (Message msg : historyDialog) {
- if (msg instanceof UserMessage) {
- promptBuilder.append("用户:").append(msg.getText()).append("\n");
- }
- }
- promptBuilder.append("\n");
- }
- promptBuilder.append("当前问题:").append(userCurrentMsg).append("\n");
- promptBuilder.append("生成的检索查询词:");
- String prompt = promptBuilder.toString();
- log.debug("[RAG-Query] 提示词:\n{}", prompt);
- try {
- // 使用本地模型生成查询词
- ChatResponse response = dashscopeModel.call(new Prompt(prompt));
- String query = response.getResult().getOutput().getText().trim();
-
- log.info("[RAG-Query] 生成的检索查询词:{}", query);
- return query;
- } catch (Exception e) {
- log.error("[RAG-Query] 生成检索查询词失败,使用原始问题", e);
- return userCurrentMsg;
- }
- }
- /**
- * 排序并限制 RAG 结果数量
- * RAG 接口返回的结果已经满足相似度阈值,只需排序和限制数量
- *
- * @param response RAG 响应
- * @param maxResults 最大结果数
- * @return 处理后的结果
- */
- public List<RagClient.RagRetrieveResponse.Result> sortAndLimitResults(
- RagClient.RagRetrieveResponse response,
- Integer maxResults) {
-
- if (response == null || response.getResults() == null) {
- log.warn("[RAG-Sort] 响应或结果为空");
- return List.of();
- }
- log.info("[RAG-Sort] 开始排序,原始结果数:{},最大结果数:{}", response.getResults().size(), maxResults);
-
- List<RagClient.RagRetrieveResponse.Result> sorted = response.getResults().stream()
- .sorted(Comparator.comparing(RagClient.RagRetrieveResponse.Result::getScore).reversed())
- .limit(maxResults)
- .collect(Collectors.toList());
-
- log.info("[RAG-Sort] 排序完成,保留结果数:{}", sorted.size());
- return sorted;
- }
- /**
- * 拼接 RAG 上下文
- *
- * @param results 检索结果
- * @return 拼接后的上下文
- */
- public String buildRagContext(List<RagClient.RagRetrieveResponse.Result> results) {
- if (results == null || results.isEmpty()) {
- return "";
- }
- return results.stream()
- .map(result -> String.format("【来源:%s】%s",
- result.getSource() != null ? result.getSource() : "未知",
- result.getText()))
- .collect(Collectors.joining("\n\n"));
- }
- /**
- * 从 RAG 结果中总结关键信息(使用本地模型)
- *
- * @param ragContext RAG 上下文
- * @param userCurrentMsg 用户当前消息
- * @return 结构化关键信息 JSON
- */
- public String extractKeyInfo(String ragContext, String userCurrentMsg) {
- log.info("[RAG-Extract] 开始总结 RAG 关键信息");
- log.info("[RAG-Extract] 用户消息:{}", userCurrentMsg);
- String prompt = String.format(
- "请从以下知识库内容中提炼与「%s」相关的关键信息,仅返回 JSON 格式(无额外文字):\n" +
- "知识库内容:\n%s\n\n" +
- "返回格式示例:{\"key_points\": [\"要点 1\", \"要点 2\"], \"summary\": \"简短总结\"}",
- userCurrentMsg,
- ragContext
- );
- log.debug("[RAG-Extract] 提示词:\n{}", prompt);
- try {
- ChatResponse response = dashscopeModel.call(new Prompt(prompt));
- String keyInfoJson = response.getResult().getOutput().getText().trim();
-
- log.info("[RAG-Extract] 总结的关键信息:{}", keyInfoJson);
- return keyInfoJson;
- } catch (Exception e) {
- log.error("[RAG-Extract] 总结关键信息失败", e);
- return null;
- }
- }
- /**
- * 判断关键信息是否需要存储
- *
- * @param keyInfo 结构化关键信息
- * @return 是否需要存储
- */
- public boolean needStore(String keyInfo) {
- if (keyInfo == null || keyInfo.isEmpty()) {
- return false;
- }
- // 简单判断:包含 key_points 或 summary 就存储
- return keyInfo.contains("key_points") || keyInfo.contains("summary");
- }
- /**
- * 执行 RAG 检索并构建最终用户消息
- *
- * @param userMessage 用户消息
- * @param chatId 会话 ID
- * @param ragConfig RAG 配置
- * @param history 历史对话
- * @return 包含 RAG 上下文的最终用户消息
- */
- public String executeRagAndBuildMessage(String userMessage, String chatId, AgentConfiguration.RagConfig ragConfig, List<Message> history) {
- try {
- log.info("[RAG] ========== 开始执行 RAG 检索 ==========");
- log.info("[RAG] 会话 ID: {}", chatId);
- log.info("[RAG] 用户原始问题: {}", userMessage);
- // 1. 生成检索查询词
- String query = generateRetrievalQuery(userMessage, history);
- log.info("[RAG] 生成的检索查询词:{}", query);
- // 2. 执行 RAG 检索
- log.info("[RAG] 调用 RAG 检索接口...");
- RagClient.RagRetrieveResponse ragResponse = retrieve(query, ragConfig).block();
-
- if (ragResponse == null || ragResponse.getResults() == null || ragResponse.getResults().isEmpty()) {
- log.warn("[RAG] 检索结果为空,使用原始问题");
- return userMessage;
- }
- log.info("[RAG] 检索成功!");
- log.info("[RAG] 检索 ID: {}", ragResponse.getRetrievalId());
- log.info("[RAG] 结果数量:{}", ragResponse.getResults().size());
-
- // 打印所有检索结果
- for (int i = 0; i < ragResponse.getResults().size(); i++) {
- RagClient.RagRetrieveResponse.Result result = ragResponse.getResults().get(i);
- log.info("[RAG] --- 结果 {} ---", i + 1);
- log.info("[RAG] 文档 ID: {}", result.getId());
- log.info("[RAG] 来源:{}", result.getSource());
- log.info("[RAG] 相似度:{}", result.getScore());
- log.info("[RAG] 重排序分数:{}", result.getRerankScore());
- log.info("[RAG] 内容:{}", result.getText().length() > 200 ?
- result.getText().substring(0, 200) + "..." : result.getText());
- }
- // 3. 排序并取前 N 条结果
- // RAG 接口返回的结果已经满足相似度阈值,只需排序
- // 使用 agent-config 中的 top_k 或默认 3
- Integer topK = 3;
- if (ragConfig != null && ragConfig.getVectorStores() != null
- && !ragConfig.getVectorStores().isEmpty()) {
- AgentConfiguration.RagConfig.VectorStoreConfig firstConfig = ragConfig.getVectorStores().get(0);
- if (firstConfig.getTopK() != null) {
- topK = firstConfig.getTopK();
- }
- }
- log.info("[RAG] 使用配置参数 - topK: {}", topK);
-
- List<RagClient.RagRetrieveResponse.Result> filteredResults = sortAndLimitResults(
- ragResponse,
- topK
- );
- if (filteredResults.isEmpty()) {
- log.warn("[RAG] 结果为空,使用原始问题");
- return userMessage;
- }
- log.info("[RAG] 保留 {} 条结果", filteredResults.size());
- // 4. 拼接 RAG 上下文
- String ragContext = buildRagContext(filteredResults);
- log.info("[RAG] 拼接后的 RAG 上下文:\n{}", ragContext);
- // 5. 触发异步总结和存储关键信息(不阻塞主流程)
- log.info("[RAG] 触发异步关键信息总结和存储...");
- summarizeAndStoreKeyInfoAsync(ragContext, userMessage, chatId);
- // 6. 构建最终用户消息
- String finalMessage = buildFinalUserMessage(userMessage, ragContext);
- log.info("[RAG] 构建的最终用户消息:{}", finalMessage);
- log.info("[RAG] ========== RAG 检索完成 ==========");
- return finalMessage;
- } catch (Exception e) {
- log.error("[RAG] 执行检索失败,使用原始问题", e);
- return userMessage;
- }
- }
- /**
- * 构建最终用户消息(包含 RAG 上下文)
- */
- public String buildFinalUserMessage(String userMessage, String ragContext) {
- StringBuilder sb = new StringBuilder();
- sb.append("请根据以下参考资料回答问题:\n\n");
- sb.append("【参考资料】\n");
- sb.append(ragContext);
- sb.append("\n\n【用户问题】\n");
- sb.append(userMessage);
- return sb.toString();
- }
- /**
- * 异步总结和存储关键信息到 Redis
- * 使用 @Async 在后台线程执行,不阻塞主流程
- * 包含:总结关键信息 + 存储到 Redis
- */
- @Async
- public void summarizeAndStoreKeyInfoAsync(String ragContext, String userMessage, String chatId) {
- try {
- log.info("[RAG-Async] 开始异步总结和存储关键信息 - chatId: {}", chatId);
-
- // 1. 总结关键信息(在异步线程执行,不阻塞主流程)
- String keyInfoJson = extractKeyInfo(ragContext, userMessage);
-
- if (keyInfoJson != null && needStore(keyInfoJson)) {
- log.info("[RAG-Async] 提取的关键信息 JSON: {}", keyInfoJson);
-
- // 2. 存储到 Redis
- redisBasedChatMemory.add(chatId, List.of(
- new UserMessage("[RAG 关键信息] " + keyInfoJson)
- ));
-
- log.info("[RAG-Async] 关键信息已存储到 Redis - chatId: {}", chatId);
- } else {
- log.info("[RAG-Async] 关键信息无需存储");
- }
- } catch (Exception e) {
- log.error("[RAG-Async] 总结和存储关键信息失败", e);
- }
- }
- }
|