fix(rag): 统一检索阈值重排降级与缓存语义

This commit is contained in:
evo
2026-07-21 09:34:56 +08:00
parent a6a55202a3
commit 6f6d0893ec
4 changed files with 148 additions and 14 deletions

View File

@@ -1,6 +1,7 @@
package org.ruoyi.service.knowledge.retriever;
import dev.langchain4j.data.segment.TextSegment;
import dev.langchain4j.data.document.Metadata;
import dev.langchain4j.rag.content.Content;
import dev.langchain4j.rag.content.retriever.ContentRetriever;
import dev.langchain4j.rag.query.Query;
@@ -55,11 +56,17 @@ public class CustomVectorRetriever implements ContentRetriever {
queryVectorBo.setRerankScoreThreshold(knowledgeInfoVo.getRerankScoreThreshold());
// 通过统一服务执行检索
List<String> nearestList = knowledgeRetrievalService.retrieveTexts(queryVectorBo);
var nearestList = knowledgeRetrievalService.retrieve(queryVectorBo);
// 将结果包装为标准的 Content 返回
return nearestList.stream()
.map(text -> Content.from(TextSegment.from(text)))
.map(vo -> {
Metadata metadata = new Metadata();
metadata.put("kid", String.valueOf(knowledgeInfoVo.getId()));
metadata.put("docId", Objects.toString(vo.getDocId(), ""));
metadata.put("fid", Objects.toString(vo.getId(), ""));
return Content.from(TextSegment.from(vo.getContent(), metadata));
})
.collect(Collectors.toList());
}
}

View File

@@ -31,4 +31,6 @@ public interface KnowledgeRetrievalService {
* @return 检索结果列表
*/
List<KnowledgeRetrievalVo> retrieve(QueryVectorBo queryVectorBo);
void invalidateKnowledge(String kid);
}

View File

@@ -3,6 +3,7 @@ package org.ruoyi.service.retrieval.impl;
import lombok.RequiredArgsConstructor;
import lombok.extern.slf4j.Slf4j;
import org.ruoyi.common.core.utils.StringUtils;
import org.ruoyi.common.core.exception.ServiceException;
import org.ruoyi.domain.bo.rerank.RerankRequest;
import org.ruoyi.domain.bo.rerank.RerankResult;
import org.ruoyi.domain.bo.vector.QueryVectorBo;
@@ -17,6 +18,8 @@ import org.springframework.stereotype.Service;
import java.util.*;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.TimeUnit;
import java.util.stream.Collectors;
/**
@@ -40,6 +43,9 @@ public class KnowledgeRetrievalServiceImpl implements KnowledgeRetrievalService
* 如果启用重排序,粗召回会获取更多结果供重排序筛选
*/
private static final int RERANK_EXPANSION_FACTOR = 3;
private static final long CACHE_TTL_MILLIS = TimeUnit.MINUTES.toMillis(5);
private static final int CACHE_MAX_ENTRIES = 1000;
private final Map<String, CacheEntry> retrievalCache = new ConcurrentHashMap<>();
@Override
public List<String> retrieveTexts(QueryVectorBo queryVectorBo) {
@@ -51,6 +57,11 @@ public class KnowledgeRetrievalServiceImpl implements KnowledgeRetrievalService
@Override
public List<KnowledgeRetrievalVo> retrieve(QueryVectorBo queryVectorBo) {
String cacheKey = cacheKey(queryVectorBo);
CacheEntry cached = retrievalCache.get(cacheKey);
if (cached != null && System.currentTimeMillis() - cached.createdAt < CACHE_TTL_MILLIS) {
return copyResults(cached.results);
}
log.info("开始知识库检索, kid={}, query={}", queryVectorBo.getKid(), queryVectorBo.getQuery());
// 1. 粗召回阶段 (向量检索 + 关键词搜索)
@@ -68,19 +79,22 @@ public class KnowledgeRetrievalServiceImpl implements KnowledgeRetrievalService
// 3. 重排序阶段 (可选)
List<KnowledgeRetrievalVo> finalResults = coarseResults;
if (Boolean.TRUE.equals(queryVectorBo.getEnableRerank()) &&
StringUtils.isNotBlank(queryVectorBo.getRerankModelName())) {
boolean rerankApplied = Boolean.TRUE.equals(queryVectorBo.getEnableRerank()) &&
StringUtils.isNotBlank(queryVectorBo.getRerankModelName());
if (rerankApplied) {
finalResults = performRerank(queryVectorBo, coarseResults);
}
// 4. 应用分值阈值过滤 (重排分值或 RRF 分值)
double threshold = queryVectorBo.getRerankScoreThreshold() != null ?
queryVectorBo.getRerankScoreThreshold() : 0.0;
return finalResults.stream()
.filter(res -> res.getScore() >= threshold)
if (rerankApplied && queryVectorBo.getRerankScoreThreshold() != null) {
double threshold = queryVectorBo.getRerankScoreThreshold();
finalResults = finalResults.stream()
.filter(res -> res.getScore() != null && res.getScore() >= threshold)
.collect(Collectors.toList());
}
cache(cacheKey, finalResults);
return copyResults(finalResults);
}
/**
* 粗召回阶段:根据配置执行向量搜索或混合搜索
@@ -132,7 +146,8 @@ public class KnowledgeRetrievalServiceImpl implements KnowledgeRetrievalService
List<KnowledgeFragmentVo> fragments = fragmentMapper.searchByKeyword(kid, queryVectorBo.getQuery(), finalTargetMaxResults);
return fragments.stream().map(f -> {
KnowledgeRetrievalVo vo = new KnowledgeRetrievalVo();
vo.setId(f.getId().toString());
// 优先使用 fid 作为融合标识(与向量侧一致),历史数据无 fid 时回退主键
vo.setId(StringUtils.isNotBlank(f.getFid()) ? f.getFid() : f.getId().toString());
vo.setContent(f.getContent());
vo.setDocId(f.getDocId());
vo.setIdx(f.getIdx());
@@ -155,7 +170,11 @@ public class KnowledgeRetrievalServiceImpl implements KnowledgeRetrievalService
} catch (Exception e) {
log.error("混合检索执行失败,回退到纯向量检索: {}", e.getMessage(), e);
try {
return vectorStoreService.search(copyOf(queryVectorBo, targetMaxResults));
} catch (Exception vectorError) {
throw new ServiceException("知识库检索不可用:向量与混合检索均失败");
}
}
}
@@ -182,19 +201,21 @@ public class KnowledgeRetrievalServiceImpl implements KnowledgeRetrievalService
RerankResult rerankResult = rerankModel.rerank(rerankRequest);
// 写回分数并记录原始分
List<KnowledgeRetrievalVo> reranked = new ArrayList<>();
for (RerankResult.RerankDocument doc : rerankResult.getDocuments()) {
if (doc.getIndex() != null && doc.getIndex() < coarseResults.size()) {
KnowledgeRetrievalVo vo = coarseResults.get(doc.getIndex());
vo.setRawScore(vo.getScore());
vo.setScore(doc.getRelevanceScore());
reranked.add(vo);
}
}
// 按新分排序
coarseResults.sort((a, b) -> b.getScore().compareTo(a.getScore()));
reranked.sort((a, b) -> b.getScore().compareTo(a.getScore()));
// 截断到 topN
return coarseResults.subList(0, Math.min(topN, coarseResults.size()));
return reranked.subList(0, Math.min(topN, reranked.size()));
} catch (Exception e) {
log.error("重排序流程失败: {}", e.getMessage());
@@ -253,4 +274,42 @@ public class KnowledgeRetrievalServiceImpl implements KnowledgeRetrievalService
copy.setBaseUrl(original.getBaseUrl());
return copy;
}
@Override
public void invalidateKnowledge(String kid) {
if (StringUtils.isBlank(kid)) {
retrievalCache.clear();
} else {
retrievalCache.keySet().removeIf(key -> key.startsWith(kid + "|"));
}
}
private void cache(String key, List<KnowledgeRetrievalVo> results) {
if (retrievalCache.size() >= CACHE_MAX_ENTRIES) {
long now = System.currentTimeMillis();
retrievalCache.entrySet().removeIf(e -> now - e.getValue().createdAt >= CACHE_TTL_MILLIS);
if (retrievalCache.size() >= CACHE_MAX_ENTRIES) {
retrievalCache.clear();
}
}
retrievalCache.put(key, new CacheEntry(System.currentTimeMillis(), copyResults(results)));
}
private String cacheKey(QueryVectorBo bo) {
return String.join("|", Objects.toString(bo.getKid(), ""), Objects.toString(bo.getQuery(), ""),
Objects.toString(bo.getMaxResults(), ""), Objects.toString(bo.getVectorModelName(), ""),
Objects.toString(bo.getEmbeddingModelName(), ""), Objects.toString(bo.getSimilarityThreshold(), ""),
Objects.toString(bo.getEnableHybrid(), ""), Objects.toString(bo.getHybridAlpha(), ""),
Objects.toString(bo.getEnableRerank(), ""), Objects.toString(bo.getRerankModelName(), ""),
Objects.toString(bo.getRerankTopN(), ""), Objects.toString(bo.getRerankScoreThreshold(), ""));
}
private List<KnowledgeRetrievalVo> copyResults(List<KnowledgeRetrievalVo> source) {
return source.stream().map(vo -> KnowledgeRetrievalVo.builder()
.id(vo.getId()).docId(vo.getDocId()).knowledgeId(vo.getKnowledgeId()).idx(vo.getIdx())
.content(vo.getContent()).score(vo.getScore()).originalIndex(vo.getOriginalIndex())
.rawScore(vo.getRawScore()).sourceName(vo.getSourceName()).build()).collect(Collectors.toList());
}
private record CacheEntry(long createdAt, List<KnowledgeRetrievalVo> results) { }
}

View File

@@ -0,0 +1,66 @@
package org.ruoyi.service.retrieval.impl;
import org.junit.jupiter.api.Tag;
import org.junit.jupiter.api.Test;
import org.ruoyi.domain.bo.vector.QueryVectorBo;
import org.ruoyi.domain.vo.knowledge.KnowledgeRetrievalVo;
import org.ruoyi.factory.RerankModelFactory;
import org.ruoyi.mapper.knowledge.KnowledgeFragmentMapper;
import org.ruoyi.service.vector.VectorStoreService;
import java.lang.reflect.Method;
import java.util.List;
import static org.junit.jupiter.api.Assertions.*;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.when;
@Tag("dev")
class KnowledgeRetrievalServiceRegressionTest {
@Test
void rrfMergesSameStableIdAndKeepsDistinctItems() throws Exception {
KnowledgeRetrievalServiceImpl service = new KnowledgeRetrievalServiceImpl(
mock(VectorStoreService.class), mock(RerankModelFactory.class), mock(KnowledgeFragmentMapper.class));
KnowledgeRetrievalVo vectorA = item("fid-a", "A", 0.9);
KnowledgeRetrievalVo vectorB = item("fid-b", "B", 0.8);
KnowledgeRetrievalVo keywordA = item("fid-a", "A", 10.0);
KnowledgeRetrievalVo keywordC = item("fid-c", "C", 10.0);
Method method = KnowledgeRetrievalServiceImpl.class.getDeclaredMethod(
"calculateRRF", List.class, List.class, double.class);
method.setAccessible(true);
@SuppressWarnings("unchecked")
List<KnowledgeRetrievalVo> result = (List<KnowledgeRetrievalVo>) method.invoke(
service, List.of(vectorA, vectorB), List.of(keywordA, keywordC), 0.5);
assertEquals(3, result.size());
assertEquals("fid-a", result.get(0).getId(), "an item present in both rankings must win");
assertEquals(3, result.stream().map(KnowledgeRetrievalVo::getId).distinct().count());
}
@Test
void rerankThresholdDoesNotLeakIntoNonRerankRetrieval() {
VectorStoreService vectorStore = mock(VectorStoreService.class);
when(vectorStore.search(org.mockito.ArgumentMatchers.any())).thenReturn(List.of(item("fid-a", "A", 0.75)));
KnowledgeRetrievalServiceImpl service = new KnowledgeRetrievalServiceImpl(
vectorStore, mock(RerankModelFactory.class), mock(KnowledgeFragmentMapper.class));
QueryVectorBo query = new QueryVectorBo();
query.setKid("1");
query.setQuery("test");
query.setMaxResults(5);
query.setEnableHybrid(false);
query.setEnableRerank(false);
query.setSimilarityThreshold(0.5);
query.setRerankScoreThreshold(0.9);
List<KnowledgeRetrievalVo> result = service.retrieve(query);
assertEquals(1, result.size());
assertEquals("fid-a", result.get(0).getId());
}
private static KnowledgeRetrievalVo item(String id, String content, double score) {
return KnowledgeRetrievalVo.builder().id(id).content(content).score(score).build();
}
}