| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208 |
- package com.agent.management.rag.kb;
- import com.agent.management.model.entity.RagKnowledgeBase;
- import com.agent.management.model.entity.RagKnowledgeSourceBinding;
- import com.agent.management.rag.api.RagRetriever;
- import com.agent.management.rag.model.*;
- import com.agent.management.repository.RagKnowledgeBaseRepository;
- import com.agent.management.repository.RagKnowledgeSourceBindingRepository;
- import com.fasterxml.jackson.core.type.TypeReference;
- import com.fasterxml.jackson.databind.ObjectMapper;
- import org.springframework.stereotype.Service;
- import org.springframework.beans.factory.annotation.Qualifier;
- import java.util.*;
- import java.util.concurrent.CompletableFuture;
- import java.util.concurrent.Executor;
- import java.util.function.Function;
- import java.util.stream.Collectors;
- @Service
- public class KnowledgeBaseRagRetriever {
- private final RagKnowledgeBaseRepository bases;
- private final RagKnowledgeSourceBindingRepository bindings;
- private final Map<RagSourceType, RagRetriever> retrievers;
- private final ObjectMapper mapper;
- private final Executor executor;
- private final RagQuestionPlanningService planner;
- public KnowledgeBaseRagRetriever(RagKnowledgeBaseRepository bases,
- RagKnowledgeSourceBindingRepository bindings,
- List<RagRetriever> retrievers,
- ObjectMapper mapper,
- @Qualifier("ragRetrievalExecutor") Executor executor,
- RagQuestionPlanningService planner) {
- this.bases = bases;
- this.bindings = bindings;
- this.retrievers = retrievers.stream()
- .collect(Collectors.toMap(RagRetriever::sourceType, Function.identity()));
- this.mapper = mapper;
- this.executor = executor;
- this.planner = planner;
- }
- public KnowledgeBaseRagResult retrieve(KnowledgeBaseRagRequest request) {
- KnowledgeBaseRagResult out = new KnowledgeBaseRagResult();
- out.setKnowledgeBaseId(request.getKnowledgeBaseId());
- Optional<RagKnowledgeBase> base = bases.findById(request.getKnowledgeBaseId());
- if (base.isEmpty() || !base.get().isEnabled()) {
- out.getDiagnostics().put("error", "knowledge base does not exist or is disabled");
- return out;
- }
- List<RagKnowledgeSourceBinding> selected = bindings
- .findByKnowledgeBaseIdAndEnabledTrueOrderByPriorityAsc(request.getKnowledgeBaseId())
- .stream().filter(binding -> requested(request, binding.getSourceType())).toList();
- selected.forEach(binding -> out.getEnabledSources().add(binding.getSourceType()));
- Map<RagSourceType,String> subQuestions = planner.plan(request.getQuery(), selected);
- List<CompletableFuture<BindingResult>> tasks = selected.stream()
- .map(binding -> CompletableFuture.supplyAsync(() -> retrieveBinding(request, binding,
- subQuestions.getOrDefault(binding.getSourceType(), request.getQuery())), executor))
- .toList();
- for (CompletableFuture<BindingResult> task : tasks) {
- merge(out, task.join());
- }
- rerank(out, request.getQuery());
- return out;
- }
- private static void merge(KnowledgeBaseRagResult out, BindingResult bindingResult) {
- out.getEvidences().addAll(bindingResult.result().getEvidences());
- if (!bindingResult.result().getDiagnostics().isEmpty()) {
- out.getDiagnostics().put(bindingResult.key(), bindingResult.result().getDiagnostics());
- }
- }
- private BindingResult retrieveBinding(KnowledgeBaseRagRequest request, RagKnowledgeSourceBinding binding, String subQuestion) {
- RagRetrievalResult result = new RagRetrievalResult();
- result.setSourceType(binding.getSourceType());
- RagRetriever retriever = retrievers.get(binding.getSourceType());
- if (retriever == null) {
- result.getDiagnostics().put("error", "retriever unavailable");
- return new BindingResult(binding, result);
- }
- try {
- Map<String, Object> filters = mergedFilters(request, binding);
- RagQuery query = new RagQuery();
- query.setQuery(subQuestion);
- query.setUserId(request.getUserId());
- query.setSessionId(request.getSessionId());
- query.setSourceIds(List.of(binding.getSourceId()));
- query.setTopK(request.getTopK() != null ? request.getTopK() : integer(filters.get("topK")));
- query.setFilters(filters);
- result = retriever.retrieve(query);
- } catch (Exception e) {
- result.getDiagnostics().put("error", e.getMessage());
- }
- return new BindingResult(binding, result);
- }
- private Map<String, Object> mergedFilters(KnowledgeBaseRagRequest request,
- RagKnowledgeSourceBinding binding) {
- Map<String, Object> result = new LinkedHashMap<>();
- if (binding.getConfigJson() != null && !binding.getConfigJson().isBlank()) {
- try {
- result.putAll(mapper.readValue(binding.getConfigJson(), new TypeReference<>() {}));
- } catch (Exception e) {
- result.put("configWarning", "invalid configJson: " + e.getMessage());
- }
- }
- if (binding.getSourceType() == RagSourceType.STRUCTURED_DATA && result.containsKey("defaultSql")) {
- result.putIfAbsent("sql", result.get("defaultSql"));
- }
- if (binding.getSourceType() == RagSourceType.GRAPH && result.containsKey("defaultCypher")) {
- result.putIfAbsent("cypher", result.get("defaultCypher"));
- }
- if (request.getFilters() != null) {
- String key = switch (binding.getSourceType()) {
- case DOCUMENT -> "document";
- case STRUCTURED_DATA -> "structured";
- case GRAPH -> "graph";
- };
- Object scoped = request.getFilters().get(key);
- if (scoped instanceof Map<?, ?> map) {
- map.forEach((filterKey, value) -> result.put(String.valueOf(filterKey), value));
- }
- }
- return result;
- }
- private static boolean requested(KnowledgeBaseRagRequest request, RagSourceType type) {
- if (request.getFilters() == null) return true;
- Object value = request.getFilters().get("sourceTypes");
- if (!(value instanceof List<?> list) || list.isEmpty()) return true;
- return list.stream().map(String::valueOf).anyMatch(type.name()::equalsIgnoreCase);
- }
- private static Integer integer(Object value) {
- if (value instanceof Number number) return number.intValue();
- try {
- return value == null ? null : Integer.valueOf(String.valueOf(value));
- } catch (NumberFormatException e) {
- return null;
- }
- }
- private static void rerank(KnowledgeBaseRagResult out, String question) {
- if (out.getEvidences().size() <= 1) return;
- out.getEvidences().forEach(evidence -> {
- double score = rerankScore(evidence, question);
- Map<String, Object> metadata = new LinkedHashMap<>(
- evidence.getMetadata() == null ? Map.of() : evidence.getMetadata());
- metadata.put("rerankScore", score);
- evidence.setMetadata(metadata);
- });
- out.getEvidences().sort(Comparator.comparingDouble(
- (RagEvidence evidence) -> number(evidence.getMetadata() == null ? null : evidence.getMetadata().get("rerankScore")))
- .reversed());
- out.getDiagnostics().put("rerank", Map.of("method", "lexical-source-weight", "evidenceCount", out.getEvidences().size()));
- }
- private static double rerankScore(RagEvidence evidence, String question) {
- double base = evidence.getScore() == null ? 0 : evidence.getScore();
- double sourceWeight = switch (evidence.getSourceType()) {
- case STRUCTURED_DATA -> 0.18;
- case GRAPH -> 0.16;
- case DOCUMENT -> 0.08;
- };
- double exactness = switch (String.valueOf(evidence.getEvidenceType())) {
- case "SQL_RESULT", "GRAPH_RESULT" -> 0.12;
- case "SQL_DIAGNOSTIC" -> -0.15;
- default -> 0.0;
- };
- return base + sourceWeight + exactness + overlap(question, evidenceText(evidence));
- }
- private static String evidenceText(RagEvidence evidence) {
- StringBuilder text = new StringBuilder();
- text.append(evidence.getTitle()).append(' ').append(evidence.getContent()).append(' ');
- if (evidence.getPayload() != null) text.append(evidence.getPayload());
- return text.toString();
- }
- private static double overlap(String question, String evidence) {
- if (question == null || question.isBlank() || evidence == null || evidence.isBlank()) return 0;
- String lowerEvidence = evidence.toLowerCase(Locale.ROOT);
- LinkedHashSet<String> terms = new LinkedHashSet<>();
- java.util.regex.Matcher ascii = java.util.regex.Pattern.compile("[A-Za-z0-9_]{3,}").matcher(question);
- while (ascii.find()) terms.add(ascii.group().toLowerCase(Locale.ROOT));
- java.util.regex.Matcher han = java.util.regex.Pattern.compile("\\p{IsHan}{2,}").matcher(question);
- while (han.find()) {
- String run = han.group();
- if (run.length() <= 4) terms.add(run);
- else for (int i = 0; i < run.length() - 1; i++) terms.add(run.substring(i, i + 2));
- }
- if (terms.isEmpty()) return 0;
- long hits = terms.stream().filter(term -> lowerEvidence.contains(term.toLowerCase(Locale.ROOT))).count();
- return Math.min(0.35, hits * 0.035);
- }
- private static double number(Object value) {
- return value instanceof Number number ? number.doubleValue() : 0;
- }
- private record BindingResult(RagKnowledgeSourceBinding binding, RagRetrievalResult result) {
- String key() { return binding.getSourceType() + ":" + binding.getSourceId(); }
- }
- }
|