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 retrievers; private final ObjectMapper mapper; private final Executor executor; private final RagQuestionPlanningService planner; public KnowledgeBaseRagRetriever(RagKnowledgeBaseRepository bases, RagKnowledgeSourceBindingRepository bindings, List 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 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 selected = bindings .findByKnowledgeBaseIdAndEnabledTrueOrderByPriorityAsc(request.getKnowledgeBaseId()) .stream().filter(binding -> requested(request, binding.getSourceType())).toList(); selected.forEach(binding -> out.getEnabledSources().add(binding.getSourceType())); Map subQuestions = planner.plan(request.getQuery(), selected); List> tasks = selected.stream() .map(binding -> CompletableFuture.supplyAsync(() -> retrieveBinding(request, binding, subQuestions.getOrDefault(binding.getSourceType(), request.getQuery())), executor)) .toList(); for (CompletableFuture 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 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 mergedFilters(KnowledgeBaseRagRequest request, RagKnowledgeSourceBinding binding) { Map 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 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 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(); } } }