RagService.java 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343
  1. package edu.nju.software.aipaasagent.service;
  2. import edu.nju.software.aipaasagent.agent.config.AgentConfiguration;
  3. import edu.nju.software.aipaasagent.rag.client.RagClient;
  4. import edu.nju.software.aipaasagent.rag.dto.RagRetrieveRequest;
  5. import edu.nju.software.aipaasagent.memory.chat.RedisBasedChatMemory;
  6. import lombok.extern.slf4j.Slf4j;
  7. import org.springframework.ai.chat.messages.Message;
  8. import org.springframework.ai.chat.messages.UserMessage;
  9. import org.springframework.ai.chat.model.ChatModel;
  10. import org.springframework.ai.chat.model.ChatResponse;
  11. import org.springframework.ai.chat.prompt.Prompt;
  12. import org.springframework.beans.factory.annotation.Autowired;
  13. import org.springframework.scheduling.annotation.Async;
  14. import org.springframework.stereotype.Service;
  15. import reactor.core.publisher.Mono;
  16. import java.util.Comparator;
  17. import java.util.List;
  18. import java.util.stream.Collectors;
  19. /**
  20. * RAG 服务
  21. * 封装 RAG 检索逻辑
  22. */
  23. @Slf4j
  24. @Service
  25. public class RagService {
  26. @Autowired
  27. private RagClient ragClient;
  28. @Autowired
  29. private ChatModel dashscopeModel;
  30. @Autowired
  31. private RedisBasedChatMemory redisBasedChatMemory;
  32. /**
  33. * 执行 RAG 检索
  34. *
  35. * @param query 查询词
  36. * @param ragConfig RAG 配置
  37. * @return 检索结果
  38. */
  39. public Mono<RagClient.RagRetrieveResponse> retrieve(String query, AgentConfiguration.RagConfig ragConfig) {
  40. if (ragConfig == null || !Boolean.TRUE.equals(ragConfig.getEnabled())) {
  41. log.debug("RAG 未启用,跳过检索");
  42. return Mono.empty();
  43. }
  44. // 转换为 RAG 客户端请求
  45. List<RagRetrieveRequest.VectorStoreConfig> vectorStores = ragConfig.getVectorStores().stream()
  46. .map(config -> RagRetrieveRequest.VectorStoreConfig.builder()
  47. .vectorStoreId(config.getVectorStoreId())
  48. .topK(config.getTopK())
  49. .scoreThreshold(config.getScoreThreshold())
  50. .build())
  51. .collect(Collectors.toList());
  52. RagRetrieveRequest request = RagRetrieveRequest.builder()
  53. .query(query)
  54. .vectorStores(vectorStores)
  55. .build();
  56. return ragClient.retrieve(request);
  57. }
  58. /**
  59. * 生成 RAG 检索查询词(使用本地模型)
  60. *
  61. * @param userCurrentMsg 用户当前消息
  62. * @param historyDialog 历史对话
  63. * @return 精准检索关键词
  64. */
  65. public String generateRetrievalQuery(String userCurrentMsg, List<Message> historyDialog) {
  66. log.info("[RAG-Query] 开始生成检索查询词");
  67. log.info("[RAG-Query] 用户原始消息:{}", userCurrentMsg);
  68. // 构建提示词
  69. StringBuilder promptBuilder = new StringBuilder();
  70. promptBuilder.append("你是一个检索查询优化助手。请根据用户当前问题和历史对话,生成一个精准的检索查询词。\n");
  71. promptBuilder.append("要求:\n");
  72. promptBuilder.append("1. 提取核心关键词\n");
  73. promptBuilder.append("2. 去除无关词汇\n");
  74. promptBuilder.append("3. 保持简洁\n");
  75. promptBuilder.append("4. 只返回查询词,不要有其他内容\n\n");
  76. if (historyDialog != null && !historyDialog.isEmpty()) {
  77. promptBuilder.append("历史对话:\n");
  78. for (Message msg : historyDialog) {
  79. if (msg instanceof UserMessage) {
  80. promptBuilder.append("用户:").append(msg.getText()).append("\n");
  81. }
  82. }
  83. promptBuilder.append("\n");
  84. }
  85. promptBuilder.append("当前问题:").append(userCurrentMsg).append("\n");
  86. promptBuilder.append("生成的检索查询词:");
  87. String prompt = promptBuilder.toString();
  88. log.debug("[RAG-Query] 提示词:\n{}", prompt);
  89. try {
  90. // 使用本地模型生成查询词
  91. ChatResponse response = dashscopeModel.call(new Prompt(prompt));
  92. String query = response.getResult().getOutput().getText().trim();
  93. log.info("[RAG-Query] 生成的检索查询词:{}", query);
  94. return query;
  95. } catch (Exception e) {
  96. log.error("[RAG-Query] 生成检索查询词失败,使用原始问题", e);
  97. return userCurrentMsg;
  98. }
  99. }
  100. /**
  101. * 排序并限制 RAG 结果数量
  102. * RAG 接口返回的结果已经满足相似度阈值,只需排序和限制数量
  103. *
  104. * @param response RAG 响应
  105. * @param maxResults 最大结果数
  106. * @return 处理后的结果
  107. */
  108. public List<RagClient.RagRetrieveResponse.Result> sortAndLimitResults(
  109. RagClient.RagRetrieveResponse response,
  110. Integer maxResults) {
  111. if (response == null || response.getResults() == null) {
  112. log.warn("[RAG-Sort] 响应或结果为空");
  113. return List.of();
  114. }
  115. log.info("[RAG-Sort] 开始排序,原始结果数:{},最大结果数:{}", response.getResults().size(), maxResults);
  116. List<RagClient.RagRetrieveResponse.Result> sorted = response.getResults().stream()
  117. .sorted(Comparator.comparing(RagClient.RagRetrieveResponse.Result::getScore).reversed())
  118. .limit(maxResults)
  119. .collect(Collectors.toList());
  120. log.info("[RAG-Sort] 排序完成,保留结果数:{}", sorted.size());
  121. return sorted;
  122. }
  123. /**
  124. * 拼接 RAG 上下文
  125. *
  126. * @param results 检索结果
  127. * @return 拼接后的上下文
  128. */
  129. public String buildRagContext(List<RagClient.RagRetrieveResponse.Result> results) {
  130. if (results == null || results.isEmpty()) {
  131. return "";
  132. }
  133. return results.stream()
  134. .map(result -> String.format("【来源:%s】%s",
  135. result.getSource() != null ? result.getSource() : "未知",
  136. result.getText()))
  137. .collect(Collectors.joining("\n\n"));
  138. }
  139. /**
  140. * 从 RAG 结果中总结关键信息(使用本地模型)
  141. *
  142. * @param ragContext RAG 上下文
  143. * @param userCurrentMsg 用户当前消息
  144. * @return 结构化关键信息 JSON
  145. */
  146. public String extractKeyInfo(String ragContext, String userCurrentMsg) {
  147. log.info("[RAG-Extract] 开始总结 RAG 关键信息");
  148. log.info("[RAG-Extract] 用户消息:{}", userCurrentMsg);
  149. String prompt = String.format(
  150. "请从以下知识库内容中提炼与「%s」相关的关键信息,仅返回 JSON 格式(无额外文字):\n" +
  151. "知识库内容:\n%s\n\n" +
  152. "返回格式示例:{\"key_points\": [\"要点 1\", \"要点 2\"], \"summary\": \"简短总结\"}",
  153. userCurrentMsg,
  154. ragContext
  155. );
  156. log.debug("[RAG-Extract] 提示词:\n{}", prompt);
  157. try {
  158. ChatResponse response = dashscopeModel.call(new Prompt(prompt));
  159. String keyInfoJson = response.getResult().getOutput().getText().trim();
  160. log.info("[RAG-Extract] 总结的关键信息:{}", keyInfoJson);
  161. return keyInfoJson;
  162. } catch (Exception e) {
  163. log.error("[RAG-Extract] 总结关键信息失败", e);
  164. return null;
  165. }
  166. }
  167. /**
  168. * 判断关键信息是否需要存储
  169. *
  170. * @param keyInfo 结构化关键信息
  171. * @return 是否需要存储
  172. */
  173. public boolean needStore(String keyInfo) {
  174. if (keyInfo == null || keyInfo.isEmpty()) {
  175. return false;
  176. }
  177. // 简单判断:包含 key_points 或 summary 就存储
  178. return keyInfo.contains("key_points") || keyInfo.contains("summary");
  179. }
  180. /**
  181. * 执行 RAG 检索并构建最终用户消息
  182. *
  183. * @param userMessage 用户消息
  184. * @param chatId 会话 ID
  185. * @param ragConfig RAG 配置
  186. * @param history 历史对话
  187. * @return 包含 RAG 上下文的最终用户消息
  188. */
  189. public String executeRagAndBuildMessage(String userMessage, String chatId, AgentConfiguration.RagConfig ragConfig, List<Message> history) {
  190. try {
  191. log.info("[RAG] ========== 开始执行 RAG 检索 ==========");
  192. log.info("[RAG] 会话 ID: {}", chatId);
  193. log.info("[RAG] 用户原始问题: {}", userMessage);
  194. // 1. 生成检索查询词
  195. String query = generateRetrievalQuery(userMessage, history);
  196. log.info("[RAG] 生成的检索查询词:{}", query);
  197. // 2. 执行 RAG 检索
  198. log.info("[RAG] 调用 RAG 检索接口...");
  199. RagClient.RagRetrieveResponse ragResponse = retrieve(query, ragConfig).block();
  200. if (ragResponse == null || ragResponse.getResults() == null || ragResponse.getResults().isEmpty()) {
  201. log.warn("[RAG] 检索结果为空,使用原始问题");
  202. return userMessage;
  203. }
  204. log.info("[RAG] 检索成功!");
  205. log.info("[RAG] 检索 ID: {}", ragResponse.getRetrievalId());
  206. log.info("[RAG] 结果数量:{}", ragResponse.getResults().size());
  207. // 打印所有检索结果
  208. for (int i = 0; i < ragResponse.getResults().size(); i++) {
  209. RagClient.RagRetrieveResponse.Result result = ragResponse.getResults().get(i);
  210. log.info("[RAG] --- 结果 {} ---", i + 1);
  211. log.info("[RAG] 文档 ID: {}", result.getId());
  212. log.info("[RAG] 来源:{}", result.getSource());
  213. log.info("[RAG] 相似度:{}", result.getScore());
  214. log.info("[RAG] 重排序分数:{}", result.getRerankScore());
  215. log.info("[RAG] 内容:{}", result.getText().length() > 200 ?
  216. result.getText().substring(0, 200) + "..." : result.getText());
  217. }
  218. // 3. 排序并取前 N 条结果
  219. // RAG 接口返回的结果已经满足相似度阈值,只需排序
  220. // 使用 agent-config 中的 top_k 或默认 3
  221. Integer topK = 3;
  222. if (ragConfig != null && ragConfig.getVectorStores() != null
  223. && !ragConfig.getVectorStores().isEmpty()) {
  224. AgentConfiguration.RagConfig.VectorStoreConfig firstConfig = ragConfig.getVectorStores().get(0);
  225. if (firstConfig.getTopK() != null) {
  226. topK = firstConfig.getTopK();
  227. }
  228. }
  229. log.info("[RAG] 使用配置参数 - topK: {}", topK);
  230. List<RagClient.RagRetrieveResponse.Result> filteredResults = sortAndLimitResults(
  231. ragResponse,
  232. topK
  233. );
  234. if (filteredResults.isEmpty()) {
  235. log.warn("[RAG] 结果为空,使用原始问题");
  236. return userMessage;
  237. }
  238. log.info("[RAG] 保留 {} 条结果", filteredResults.size());
  239. // 4. 拼接 RAG 上下文
  240. String ragContext = buildRagContext(filteredResults);
  241. log.info("[RAG] 拼接后的 RAG 上下文:\n{}", ragContext);
  242. // 5. 触发异步总结和存储关键信息(不阻塞主流程)
  243. log.info("[RAG] 触发异步关键信息总结和存储...");
  244. summarizeAndStoreKeyInfoAsync(ragContext, userMessage, chatId);
  245. // 6. 构建最终用户消息
  246. String finalMessage = buildFinalUserMessage(userMessage, ragContext);
  247. log.info("[RAG] 构建的最终用户消息:{}", finalMessage);
  248. log.info("[RAG] ========== RAG 检索完成 ==========");
  249. return finalMessage;
  250. } catch (Exception e) {
  251. log.error("[RAG] 执行检索失败,使用原始问题", e);
  252. return userMessage;
  253. }
  254. }
  255. /**
  256. * 构建最终用户消息(包含 RAG 上下文)
  257. */
  258. public String buildFinalUserMessage(String userMessage, String ragContext) {
  259. StringBuilder sb = new StringBuilder();
  260. sb.append("请根据以下参考资料回答问题:\n\n");
  261. sb.append("【参考资料】\n");
  262. sb.append(ragContext);
  263. sb.append("\n\n【用户问题】\n");
  264. sb.append(userMessage);
  265. return sb.toString();
  266. }
  267. /**
  268. * 异步总结和存储关键信息到 Redis
  269. * 使用 @Async 在后台线程执行,不阻塞主流程
  270. * 包含:总结关键信息 + 存储到 Redis
  271. */
  272. @Async
  273. public void summarizeAndStoreKeyInfoAsync(String ragContext, String userMessage, String chatId) {
  274. try {
  275. log.info("[RAG-Async] 开始异步总结和存储关键信息 - chatId: {}", chatId);
  276. // 1. 总结关键信息(在异步线程执行,不阻塞主流程)
  277. String keyInfoJson = extractKeyInfo(ragContext, userMessage);
  278. if (keyInfoJson != null && needStore(keyInfoJson)) {
  279. log.info("[RAG-Async] 提取的关键信息 JSON: {}", keyInfoJson);
  280. // 2. 存储到 Redis
  281. redisBasedChatMemory.add(chatId, List.of(
  282. new UserMessage("[RAG 关键信息] " + keyInfoJson)
  283. ));
  284. log.info("[RAG-Async] 关键信息已存储到 Redis - chatId: {}", chatId);
  285. } else {
  286. log.info("[RAG-Async] 关键信息无需存储");
  287. }
  288. } catch (Exception e) {
  289. log.error("[RAG-Async] 总结和存储关键信息失败", e);
  290. }
  291. }
  292. }