KnowledgeBaseRagRetriever.java 10.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208
  1. package com.agent.management.rag.kb;
  2. import com.agent.management.model.entity.RagKnowledgeBase;
  3. import com.agent.management.model.entity.RagKnowledgeSourceBinding;
  4. import com.agent.management.rag.api.RagRetriever;
  5. import com.agent.management.rag.model.*;
  6. import com.agent.management.repository.RagKnowledgeBaseRepository;
  7. import com.agent.management.repository.RagKnowledgeSourceBindingRepository;
  8. import com.fasterxml.jackson.core.type.TypeReference;
  9. import com.fasterxml.jackson.databind.ObjectMapper;
  10. import org.springframework.stereotype.Service;
  11. import org.springframework.beans.factory.annotation.Qualifier;
  12. import java.util.*;
  13. import java.util.concurrent.CompletableFuture;
  14. import java.util.concurrent.Executor;
  15. import java.util.function.Function;
  16. import java.util.stream.Collectors;
  17. @Service
  18. public class KnowledgeBaseRagRetriever {
  19. private final RagKnowledgeBaseRepository bases;
  20. private final RagKnowledgeSourceBindingRepository bindings;
  21. private final Map<RagSourceType, RagRetriever> retrievers;
  22. private final ObjectMapper mapper;
  23. private final Executor executor;
  24. private final RagQuestionPlanningService planner;
  25. public KnowledgeBaseRagRetriever(RagKnowledgeBaseRepository bases,
  26. RagKnowledgeSourceBindingRepository bindings,
  27. List<RagRetriever> retrievers,
  28. ObjectMapper mapper,
  29. @Qualifier("ragRetrievalExecutor") Executor executor,
  30. RagQuestionPlanningService planner) {
  31. this.bases = bases;
  32. this.bindings = bindings;
  33. this.retrievers = retrievers.stream()
  34. .collect(Collectors.toMap(RagRetriever::sourceType, Function.identity()));
  35. this.mapper = mapper;
  36. this.executor = executor;
  37. this.planner = planner;
  38. }
  39. public KnowledgeBaseRagResult retrieve(KnowledgeBaseRagRequest request) {
  40. KnowledgeBaseRagResult out = new KnowledgeBaseRagResult();
  41. out.setKnowledgeBaseId(request.getKnowledgeBaseId());
  42. Optional<RagKnowledgeBase> base = bases.findById(request.getKnowledgeBaseId());
  43. if (base.isEmpty() || !base.get().isEnabled()) {
  44. out.getDiagnostics().put("error", "knowledge base does not exist or is disabled");
  45. return out;
  46. }
  47. List<RagKnowledgeSourceBinding> selected = bindings
  48. .findByKnowledgeBaseIdAndEnabledTrueOrderByPriorityAsc(request.getKnowledgeBaseId())
  49. .stream().filter(binding -> requested(request, binding.getSourceType())).toList();
  50. selected.forEach(binding -> out.getEnabledSources().add(binding.getSourceType()));
  51. Map<RagSourceType,String> subQuestions = planner.plan(request.getQuery(), selected);
  52. List<CompletableFuture<BindingResult>> tasks = selected.stream()
  53. .map(binding -> CompletableFuture.supplyAsync(() -> retrieveBinding(request, binding,
  54. subQuestions.getOrDefault(binding.getSourceType(), request.getQuery())), executor))
  55. .toList();
  56. for (CompletableFuture<BindingResult> task : tasks) {
  57. merge(out, task.join());
  58. }
  59. rerank(out, request.getQuery());
  60. return out;
  61. }
  62. private static void merge(KnowledgeBaseRagResult out, BindingResult bindingResult) {
  63. out.getEvidences().addAll(bindingResult.result().getEvidences());
  64. if (!bindingResult.result().getDiagnostics().isEmpty()) {
  65. out.getDiagnostics().put(bindingResult.key(), bindingResult.result().getDiagnostics());
  66. }
  67. }
  68. private BindingResult retrieveBinding(KnowledgeBaseRagRequest request, RagKnowledgeSourceBinding binding, String subQuestion) {
  69. RagRetrievalResult result = new RagRetrievalResult();
  70. result.setSourceType(binding.getSourceType());
  71. RagRetriever retriever = retrievers.get(binding.getSourceType());
  72. if (retriever == null) {
  73. result.getDiagnostics().put("error", "retriever unavailable");
  74. return new BindingResult(binding, result);
  75. }
  76. try {
  77. Map<String, Object> filters = mergedFilters(request, binding);
  78. RagQuery query = new RagQuery();
  79. query.setQuery(subQuestion);
  80. query.setUserId(request.getUserId());
  81. query.setSessionId(request.getSessionId());
  82. query.setSourceIds(List.of(binding.getSourceId()));
  83. query.setTopK(request.getTopK() != null ? request.getTopK() : integer(filters.get("topK")));
  84. query.setFilters(filters);
  85. result = retriever.retrieve(query);
  86. } catch (Exception e) {
  87. result.getDiagnostics().put("error", e.getMessage());
  88. }
  89. return new BindingResult(binding, result);
  90. }
  91. private Map<String, Object> mergedFilters(KnowledgeBaseRagRequest request,
  92. RagKnowledgeSourceBinding binding) {
  93. Map<String, Object> result = new LinkedHashMap<>();
  94. if (binding.getConfigJson() != null && !binding.getConfigJson().isBlank()) {
  95. try {
  96. result.putAll(mapper.readValue(binding.getConfigJson(), new TypeReference<>() {}));
  97. } catch (Exception e) {
  98. result.put("configWarning", "invalid configJson: " + e.getMessage());
  99. }
  100. }
  101. if (binding.getSourceType() == RagSourceType.STRUCTURED_DATA && result.containsKey("defaultSql")) {
  102. result.putIfAbsent("sql", result.get("defaultSql"));
  103. }
  104. if (binding.getSourceType() == RagSourceType.GRAPH && result.containsKey("defaultCypher")) {
  105. result.putIfAbsent("cypher", result.get("defaultCypher"));
  106. }
  107. if (request.getFilters() != null) {
  108. String key = switch (binding.getSourceType()) {
  109. case DOCUMENT -> "document";
  110. case STRUCTURED_DATA -> "structured";
  111. case GRAPH -> "graph";
  112. };
  113. Object scoped = request.getFilters().get(key);
  114. if (scoped instanceof Map<?, ?> map) {
  115. map.forEach((filterKey, value) -> result.put(String.valueOf(filterKey), value));
  116. }
  117. }
  118. return result;
  119. }
  120. private static boolean requested(KnowledgeBaseRagRequest request, RagSourceType type) {
  121. if (request.getFilters() == null) return true;
  122. Object value = request.getFilters().get("sourceTypes");
  123. if (!(value instanceof List<?> list) || list.isEmpty()) return true;
  124. return list.stream().map(String::valueOf).anyMatch(type.name()::equalsIgnoreCase);
  125. }
  126. private static Integer integer(Object value) {
  127. if (value instanceof Number number) return number.intValue();
  128. try {
  129. return value == null ? null : Integer.valueOf(String.valueOf(value));
  130. } catch (NumberFormatException e) {
  131. return null;
  132. }
  133. }
  134. private static void rerank(KnowledgeBaseRagResult out, String question) {
  135. if (out.getEvidences().size() <= 1) return;
  136. out.getEvidences().forEach(evidence -> {
  137. double score = rerankScore(evidence, question);
  138. Map<String, Object> metadata = new LinkedHashMap<>(
  139. evidence.getMetadata() == null ? Map.of() : evidence.getMetadata());
  140. metadata.put("rerankScore", score);
  141. evidence.setMetadata(metadata);
  142. });
  143. out.getEvidences().sort(Comparator.comparingDouble(
  144. (RagEvidence evidence) -> number(evidence.getMetadata() == null ? null : evidence.getMetadata().get("rerankScore")))
  145. .reversed());
  146. out.getDiagnostics().put("rerank", Map.of("method", "lexical-source-weight", "evidenceCount", out.getEvidences().size()));
  147. }
  148. private static double rerankScore(RagEvidence evidence, String question) {
  149. double base = evidence.getScore() == null ? 0 : evidence.getScore();
  150. double sourceWeight = switch (evidence.getSourceType()) {
  151. case STRUCTURED_DATA -> 0.18;
  152. case GRAPH -> 0.16;
  153. case DOCUMENT -> 0.08;
  154. };
  155. double exactness = switch (String.valueOf(evidence.getEvidenceType())) {
  156. case "SQL_RESULT", "GRAPH_RESULT" -> 0.12;
  157. case "SQL_DIAGNOSTIC" -> -0.15;
  158. default -> 0.0;
  159. };
  160. return base + sourceWeight + exactness + overlap(question, evidenceText(evidence));
  161. }
  162. private static String evidenceText(RagEvidence evidence) {
  163. StringBuilder text = new StringBuilder();
  164. text.append(evidence.getTitle()).append(' ').append(evidence.getContent()).append(' ');
  165. if (evidence.getPayload() != null) text.append(evidence.getPayload());
  166. return text.toString();
  167. }
  168. private static double overlap(String question, String evidence) {
  169. if (question == null || question.isBlank() || evidence == null || evidence.isBlank()) return 0;
  170. String lowerEvidence = evidence.toLowerCase(Locale.ROOT);
  171. LinkedHashSet<String> terms = new LinkedHashSet<>();
  172. java.util.regex.Matcher ascii = java.util.regex.Pattern.compile("[A-Za-z0-9_]{3,}").matcher(question);
  173. while (ascii.find()) terms.add(ascii.group().toLowerCase(Locale.ROOT));
  174. java.util.regex.Matcher han = java.util.regex.Pattern.compile("\\p{IsHan}{2,}").matcher(question);
  175. while (han.find()) {
  176. String run = han.group();
  177. if (run.length() <= 4) terms.add(run);
  178. else for (int i = 0; i < run.length() - 1; i++) terms.add(run.substring(i, i + 2));
  179. }
  180. if (terms.isEmpty()) return 0;
  181. long hits = terms.stream().filter(term -> lowerEvidence.contains(term.toLowerCase(Locale.ROOT))).count();
  182. return Math.min(0.35, hits * 0.035);
  183. }
  184. private static double number(Object value) {
  185. return value instanceof Number number ? number.doubleValue() : 0;
  186. }
  187. private record BindingResult(RagKnowledgeSourceBinding binding, RagRetrievalResult result) {
  188. String key() { return binding.getSourceType() + ":" + binding.getSourceId(); }
  189. }
  190. }