feat(rag): hybrid multi-path search with RRF fusion

Add configurable hybrid mode on KnowledgeSearchPort that fuses dense
unfiltered, dense filtered, and lexical ranks via RRF while preserving
dense-compatible threshold scores. Archives Delivery 2 OpenSpec change.
This commit is contained in:
zhuyongxin
2026-07-27 18:33:31 +08:00
parent ac1f831903
commit 376ad0c241
19 changed files with 661 additions and 8 deletions
@@ -3,12 +3,15 @@ package com.superbiz.agent.service;
import com.superbiz.agent.dto.RetrievalTrace;
import com.superbiz.agent.dto.RetrievedEvidenceCandidate;
import com.superbiz.agent.service.retrieval.KnowledgeSearchHit;
import com.superbiz.agent.service.retrieval.KnowledgeSearchMode;
import com.superbiz.agent.service.retrieval.KnowledgeSearchPort;
import com.superbiz.agent.service.retrieval.KnowledgeSearchRequest;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.stereotype.Service;
import java.util.ArrayList;
import java.util.List;
import java.util.Locale;
/**
* L1 向量检索适配器。
@@ -23,6 +26,9 @@ public class KnowledgeDocumentRetriever {
private final KnowledgeSearchPort knowledgeSearchPort;
@Value("${retrieval.search.mode:dense}")
private String searchMode = "dense";
public KnowledgeDocumentRetriever(KnowledgeSearchPort knowledgeSearchPort) {
this.knowledgeSearchPort = knowledgeSearchPort;
}
@@ -39,8 +45,11 @@ public class KnowledgeDocumentRetriever {
public RetrievalAttemptResult retrieve(String attemptName, String query, String categoryFilter, int topK) {
long start = System.currentTimeMillis();
try {
KnowledgeSearchMode mode = "hybrid".equalsIgnoreCase(trim(searchMode))
? KnowledgeSearchMode.HYBRID
: KnowledgeSearchMode.DENSE;
List<KnowledgeSearchHit> hits = knowledgeSearchPort.search(
KnowledgeSearchRequest.dense(query, topK, categoryFilter));
new KnowledgeSearchRequest(query, topK, categoryFilter, mode));
List<RetrievedEvidenceCandidate> candidates = toCandidates(attemptName, hits);
return new RetrievalAttemptResult(
attempt(attemptName, query, categoryFilter, candidates.size(), null,
@@ -110,6 +119,10 @@ public class KnowledgeDocumentRetriever {
return hits.get(0).score();
}
private static String trim(String value) {
return value == null ? "" : value.trim().toLowerCase(Locale.ROOT);
}
/**
* 单次检索 attempt 的结果包。
*/
@@ -0,0 +1,73 @@
package com.superbiz.agent.service.retrieval;
import java.util.ArrayList;
import java.util.Comparator;
import java.util.LinkedHashSet;
import java.util.List;
import java.util.Locale;
import java.util.Set;
/**
* Sparse-lite lexical ranking over already recalled candidates.
* Not a substitute for inverted-index BM25; expands ordering signal only.
*/
public final class LexicalRanker {
private LexicalRanker() {
}
public static List<KnowledgeSearchHit> rank(String query, List<KnowledgeSearchHit> candidates) {
if (candidates == null || candidates.isEmpty()) {
return List.of();
}
Set<String> terms = tokenize(query);
if (terms.isEmpty()) {
return List.copyOf(candidates);
}
List<ScoredHit> scored = new ArrayList<>(candidates.size());
for (KnowledgeSearchHit hit : candidates) {
String haystack = (nullToEmpty(hit.title()) + " "
+ nullToEmpty(hit.breadcrumb()) + " "
+ nullToEmpty(hit.content())).toLowerCase(Locale.ROOT);
int hits = 0;
for (String term : terms) {
if (haystack.contains(term)) {
hits++;
}
}
double coverage = hits / (double) terms.size();
scored.add(new ScoredHit(hit, coverage, hits));
}
scored.sort(Comparator
.comparingDouble((ScoredHit s) -> s.coverage).reversed()
.thenComparingInt((ScoredHit s) -> s.hits).reversed()
.thenComparingInt(s -> s.hit.originalRank()));
return scored.stream().map(s -> s.hit).toList();
}
static Set<String> tokenize(String query) {
if (query == null || query.isBlank()) {
return Set.of();
}
String normalized = query.toLowerCase(Locale.ROOT);
String[] parts = normalized.split("[^\\p{IsAlphabetic}\\p{IsDigit}]+");
Set<String> terms = new LinkedHashSet<>();
for (String part : parts) {
if (part == null) {
continue;
}
String term = part.trim();
if (term.length() >= 2) {
terms.add(term);
}
}
return terms;
}
private static String nullToEmpty(String value) {
return value == null ? "" : value;
}
private record ScoredHit(KnowledgeSearchHit hit, double coverage, int hits) {
}
}
@@ -0,0 +1,85 @@
package com.superbiz.agent.service.retrieval;
import java.util.ArrayList;
import java.util.Comparator;
import java.util.HashMap;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import java.util.function.Function;
/**
* Reciprocal Rank Fusion helpers.
*
* <pre>
* RRF_w(d) = Σ w_i / (k + rank_i(d))
* </pre>
*/
public final class RrfFusion {
private RrfFusion() {
}
public static <T> List<Scored<T>> fuse(List<RankedPath<T>> paths,
int rrfK,
Function<T, String> identityFn) {
if (paths == null || paths.isEmpty()) {
return List.of();
}
int k = Math.max(1, rrfK);
Map<String, Acc<T>> acc = new LinkedHashMap<>();
for (RankedPath<T> path : paths) {
if (path == null || path.items() == null || path.items().isEmpty()) {
continue;
}
double weight = path.weight() <= 0 ? 1.0 : path.weight();
List<T> items = path.items();
for (int i = 0; i < items.size(); i++) {
T item = items.get(i);
if (item == null) {
continue;
}
String id = identityFn.apply(item);
if (id == null || id.isBlank()) {
continue;
}
int rank = i + 1;
double contrib = weight / (k + rank);
Acc<T> bucket = acc.computeIfAbsent(id, ignored -> new Acc<>(item));
bucket.score += contrib;
bucket.ranks.put(path.name(), rank);
// Prefer first-seen item payload; callers should put preferred path first if needed.
}
}
List<Scored<T>> scored = new ArrayList<>(acc.size());
for (Map.Entry<String, Acc<T>> entry : acc.entrySet()) {
Acc<T> value = entry.getValue();
scored.add(new Scored<>(entry.getKey(), value.item, value.score, Map.copyOf(value.ranks)));
}
scored.sort(Comparator
.comparingDouble((Scored<T> s) -> s.rrfScore()).reversed()
.thenComparing(Scored::identity));
return scored;
}
public record RankedPath<T>(String name, List<T> items, double weight) {
public RankedPath {
Objects.requireNonNull(name, "name");
items = items == null ? List.of() : List.copyOf(items);
}
}
public record Scored<T>(String identity, T item, double rrfScore, Map<String, Integer> ranks) {
}
private static final class Acc<T> {
private final T item;
private double score;
private final Map<String, Integer> ranks = new HashMap<>();
private Acc(T item) {
this.item = item;
}
}
}
@@ -2,15 +2,22 @@ package com.superbiz.agent.service.retrieval;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.service.VectorSearchService;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.stereotype.Component;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Locale;
import java.util.Map;
/**
* Dense-only {@link KnowledgeSearchPort} backed by the existing vector search facade.
* {@link KnowledgeSearchPort} backed by the existing vector search facade.
*
* <ul>
* <li>{@code DENSE}: single dense path (legacy behavior)</li>
* <li>{@code HYBRID}: dense unfiltered + optional dense filtered + lexical rank, fused by RRF</li>
* </ul>
*/
@Component
public class VectorKnowledgeSearchAdapter implements KnowledgeSearchPort {
@@ -18,6 +25,21 @@ public class VectorKnowledgeSearchAdapter implements KnowledgeSearchPort {
private final VectorSearchService vectorSearchService;
private final ObjectMapper objectMapper;
@Value("${retrieval.search.mode:dense}")
private String configuredMode = "dense";
@Value("${retrieval.hybrid.rrf-k:60}")
private int rrfK = 60;
@Value("${retrieval.hybrid.weight.dense-unfiltered:1.0}")
private double weightDenseUnfiltered = 1.0;
@Value("${retrieval.hybrid.weight.dense-filtered:1.0}")
private double weightDenseFiltered = 1.0;
@Value("${retrieval.hybrid.weight.lexical:1.0}")
private double weightLexical = 1.0;
public VectorKnowledgeSearchAdapter(VectorSearchService vectorSearchService, ObjectMapper objectMapper) {
this.vectorSearchService = vectorSearchService;
this.objectMapper = objectMapper;
@@ -25,13 +47,101 @@ public class VectorKnowledgeSearchAdapter implements KnowledgeSearchPort {
@Override
public List<KnowledgeSearchHit> search(KnowledgeSearchRequest request) {
if (request.mode() == KnowledgeSearchMode.HYBRID) {
// Delivery 2 will implement true hybrid. Until then, fall back to dense.
KnowledgeSearchMode mode = resolveMode(request.mode());
if (mode == KnowledgeSearchMode.HYBRID) {
return searchHybrid(request);
}
return searchDense(request.query(), request.topK(), request.categoryFilter());
}
private List<KnowledgeSearchHit> searchDense(String query, int topK, String categoryFilter) {
List<VectorSearchService.SearchResult> results = vectorSearchService.searchSimilarDocuments(
request.query(),
request.topK(),
request.categoryFilter());
query, topK, categoryFilter);
return toHits(results);
}
private List<KnowledgeSearchHit> searchHybrid(KnowledgeSearchRequest request) {
String query = request.query();
int topK = request.topK();
String category = trimToNull(request.categoryFilter());
List<KnowledgeSearchHit> unfiltered = searchDense(query, topK, null);
List<KnowledgeSearchHit> filtered = category == null
? List.of()
: searchDense(query, topK, category);
Map<String, KnowledgeSearchHit> unionByKey = new LinkedHashMap<>();
for (KnowledgeSearchHit hit : unfiltered) {
unionByKey.putIfAbsent(hit.evidenceKey(), hit);
}
for (KnowledgeSearchHit hit : filtered) {
unionByKey.putIfAbsent(hit.evidenceKey(), hit);
}
List<KnowledgeSearchHit> union = new ArrayList<>(unionByKey.values());
List<KnowledgeSearchHit> lexical = LexicalRanker.rank(query, union);
List<RrfFusion.RankedPath<KnowledgeSearchHit>> paths = new ArrayList<>();
paths.add(new RrfFusion.RankedPath<>("dense_unfiltered", unfiltered, weightDenseUnfiltered));
if (!filtered.isEmpty()) {
paths.add(new RrfFusion.RankedPath<>("dense_filtered", filtered, weightDenseFiltered));
}
if (!lexical.isEmpty()) {
paths.add(new RrfFusion.RankedPath<>("lexical", lexical, weightLexical));
}
List<RrfFusion.Scored<KnowledgeSearchHit>> fused = RrfFusion.fuse(
paths, rrfK, KnowledgeSearchHit::evidenceKey);
List<KnowledgeSearchHit> ordered = new ArrayList<>();
int rank = 1;
for (RrfFusion.Scored<KnowledgeSearchHit> scored : fused) {
if (ordered.size() >= topK) {
break;
}
KnowledgeSearchHit base = scored.item();
Map<String, String> metadata = new LinkedHashMap<>(
base.metadata() == null ? Map.of() : base.metadata());
metadata.put("fusedScore", Double.toString(scored.rrfScore()));
metadata.put("fusionRanks", scored.ranks().toString());
metadata.put("fusionRank", Integer.toString(rank));
ordered.add(new KnowledgeSearchHit(
base.id(),
base.content(),
base.score(),
base.rawScore(),
base.scoreLabel(),
base.metadataJson(),
metadata,
base.docId(),
base.chunkIndex(),
base.evidenceKey(),
base.source(),
base.title(),
base.breadcrumb(),
rank
));
rank++;
}
return ordered;
}
private KnowledgeSearchMode resolveMode(KnowledgeSearchMode requestMode) {
if (requestMode == KnowledgeSearchMode.HYBRID) {
return KnowledgeSearchMode.HYBRID;
}
if (requestMode == KnowledgeSearchMode.DENSE) {
// Allow global config to force hybrid even if caller passes DENSE default.
String configured = configuredMode == null ? "dense" : configuredMode.trim().toLowerCase(Locale.ROOT);
if ("hybrid".equals(configured)) {
return KnowledgeSearchMode.HYBRID;
}
return KnowledgeSearchMode.DENSE;
}
String configured = configuredMode == null ? "dense" : configuredMode.trim().toLowerCase(Locale.ROOT);
return "hybrid".equals(configured) ? KnowledgeSearchMode.HYBRID : KnowledgeSearchMode.DENSE;
}
private List<KnowledgeSearchHit> toHits(List<VectorSearchService.SearchResult> results) {
if (results == null || results.isEmpty()) {
return List.of();
}
@@ -91,4 +201,11 @@ public class VectorKnowledgeSearchAdapter implements KnowledgeSearchPort {
return Map.of();
}
}
private static String trimToNull(String value) {
if (value == null || value.isBlank()) {
return null;
}
return value.trim();
}
}
@@ -0,0 +1,50 @@
package com.superbiz.agent.service.retrieval;
import org.junit.jupiter.api.Test;
import java.util.List;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertTrue;
class RrfFusionTest {
@Test
void multiPathAgreementOutranksSinglePathHead() {
List<String> dense = List.of("a", "b", "c");
List<String> lexical = List.of("c", "b", "d");
List<RrfFusion.Scored<String>> fused = RrfFusion.fuse(
List.of(
new RrfFusion.RankedPath<>("dense", dense, 1.0),
new RrfFusion.RankedPath<>("lexical", lexical, 1.0)
),
60,
s -> s
);
// c: dense#3 + lexical#1 ; b: dense#2 + lexical#2 ; a: dense#1 only
// With k=60, c edges b slightly, and both beat single-path a.
assertEquals("c", fused.get(0).identity());
assertEquals("b", fused.get(1).identity());
assertEquals("a", fused.get(2).identity());
assertTrue(fused.get(0).rrfScore() > fused.get(2).rrfScore());
}
@Test
void pathWeightCanElevateSecondaryPath() {
List<String> dense = List.of("a", "b");
List<String> lexical = List.of("b", "a");
List<RrfFusion.Scored<String>> fused = RrfFusion.fuse(
List.of(
new RrfFusion.RankedPath<>("dense", dense, 1.0),
new RrfFusion.RankedPath<>("lexical", lexical, 2.0)
),
60,
s -> s
);
assertEquals("b", fused.get(0).identity());
}
}
@@ -0,0 +1,53 @@
package com.superbiz.agent.service.retrieval;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.service.VectorSearchService;
import org.junit.jupiter.api.Test;
import org.springframework.test.util.ReflectionTestUtils;
import java.util.List;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.when;
class VectorKnowledgeSearchAdapterHybridTest {
@Test
void hybridFusesFilteredAndUnfilteredDensePaths() {
VectorSearchService vectorSearchService = mock(VectorSearchService.class);
when(vectorSearchService.searchSimilarDocuments("pool timeout", 3, null)).thenReturn(List.of(
result("u1", "{\"_source\":\"a.md\",\"docId\":\"a\",\"chunkIndex\":0,\"title\":\"generic\"}", "generic pool", 0.4f),
result("u2", "{\"_source\":\"b.md\",\"docId\":\"b\",\"chunkIndex\":0,\"title\":\"other\"}", "other", 0.5f)
));
when(vectorSearchService.searchSimilarDocuments("pool timeout", 3, "mysql")).thenReturn(List.of(
result("f1", "{\"_source\":\"c.md\",\"docId\":\"c\",\"chunkIndex\":0,\"title\":\"mysql pool timeout\"}", "mysql pool timeout runbook", 0.35f)
));
VectorKnowledgeSearchAdapter adapter = new VectorKnowledgeSearchAdapter(vectorSearchService, new ObjectMapper());
ReflectionTestUtils.setField(adapter, "configuredMode", "hybrid");
ReflectionTestUtils.setField(adapter, "rrfK", 60);
ReflectionTestUtils.setField(adapter, "weightDenseUnfiltered", 1.0);
ReflectionTestUtils.setField(adapter, "weightDenseFiltered", 1.0);
ReflectionTestUtils.setField(adapter, "weightLexical", 1.0);
List<KnowledgeSearchHit> hits = adapter.search(
new KnowledgeSearchRequest("pool timeout", 3, "mysql", KnowledgeSearchMode.HYBRID));
assertEquals(3, hits.size());
assertTrue(hits.stream().anyMatch(hit -> "c#chunk-0".equals(hit.evidenceKey())));
assertTrue(hits.get(0).metadata().containsKey("fusedScore"));
}
private static VectorSearchService.SearchResult result(String id, String metadata, String content, float score) {
VectorSearchService.SearchResult result = new VectorSearchService.SearchResult();
result.setId(id);
result.setMetadata(metadata);
result.setContent(content);
result.setScore(score);
result.setRawScore((double) score);
result.setScoreLabel("l2_distance");
return result;
}
}