mirror of
https://gitcode.com/ageerle/ruoyi-ai.git
synced 2026-09-13 00:14:59 +00:00
fix(rag): 统一检索阈值重排降级与缓存语义
This commit is contained in:
@@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -31,4 +31,6 @@ public interface KnowledgeRetrievalService {
|
||||
* @return 检索结果列表
|
||||
*/
|
||||
List<KnowledgeRetrievalVo> retrieve(QueryVectorBo queryVectorBo);
|
||||
|
||||
void invalidateKnowledge(String kid);
|
||||
}
|
||||
|
||||
@@ -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,18 +79,21 @@ 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)
|
||||
.collect(Collectors.toList());
|
||||
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);
|
||||
return vectorStoreService.search(copyOf(queryVectorBo, targetMaxResults));
|
||||
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) { }
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user