From 6f6d0893ecd96da167aa81816374fbc085cbb69c Mon Sep 17 00:00:00 2001 From: evo <446796145@qq.com> Date: Tue, 21 Jul 2026 09:34:56 +0800 Subject: [PATCH] =?UTF-8?q?fix(rag):=20=E7=BB=9F=E4=B8=80=E6=A3=80?= =?UTF-8?q?=E7=B4=A2=E9=98=88=E5=80=BC=E9=87=8D=E6=8E=92=E9=99=8D=E7=BA=A7?= =?UTF-8?q?=E4=B8=8E=E7=BC=93=E5=AD=98=E8=AF=AD=E4=B9=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../retriever/CustomVectorRetriever.java | 11 ++- .../retrieval/KnowledgeRetrievalService.java | 2 + .../impl/KnowledgeRetrievalServiceImpl.java | 83 ++++++++++++++++--- ...owledgeRetrievalServiceRegressionTest.java | 66 +++++++++++++++ 4 files changed, 148 insertions(+), 14 deletions(-) create mode 100644 ruoyi-modules/ruoyi-chat/src/test/java/org/ruoyi/service/retrieval/impl/KnowledgeRetrievalServiceRegressionTest.java diff --git a/ruoyi-modules/ruoyi-chat/src/main/java/org/ruoyi/service/knowledge/retriever/CustomVectorRetriever.java b/ruoyi-modules/ruoyi-chat/src/main/java/org/ruoyi/service/knowledge/retriever/CustomVectorRetriever.java index f79206bc..a160b2a4 100644 --- a/ruoyi-modules/ruoyi-chat/src/main/java/org/ruoyi/service/knowledge/retriever/CustomVectorRetriever.java +++ b/ruoyi-modules/ruoyi-chat/src/main/java/org/ruoyi/service/knowledge/retriever/CustomVectorRetriever.java @@ -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 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()); } } diff --git a/ruoyi-modules/ruoyi-chat/src/main/java/org/ruoyi/service/retrieval/KnowledgeRetrievalService.java b/ruoyi-modules/ruoyi-chat/src/main/java/org/ruoyi/service/retrieval/KnowledgeRetrievalService.java index 3e0a6cab..398d3601 100644 --- a/ruoyi-modules/ruoyi-chat/src/main/java/org/ruoyi/service/retrieval/KnowledgeRetrievalService.java +++ b/ruoyi-modules/ruoyi-chat/src/main/java/org/ruoyi/service/retrieval/KnowledgeRetrievalService.java @@ -31,4 +31,6 @@ public interface KnowledgeRetrievalService { * @return 检索结果列表 */ List retrieve(QueryVectorBo queryVectorBo); + + void invalidateKnowledge(String kid); } diff --git a/ruoyi-modules/ruoyi-chat/src/main/java/org/ruoyi/service/retrieval/impl/KnowledgeRetrievalServiceImpl.java b/ruoyi-modules/ruoyi-chat/src/main/java/org/ruoyi/service/retrieval/impl/KnowledgeRetrievalServiceImpl.java index b6841ef1..3d9f5850 100644 --- a/ruoyi-modules/ruoyi-chat/src/main/java/org/ruoyi/service/retrieval/impl/KnowledgeRetrievalServiceImpl.java +++ b/ruoyi-modules/ruoyi-chat/src/main/java/org/ruoyi/service/retrieval/impl/KnowledgeRetrievalServiceImpl.java @@ -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 retrievalCache = new ConcurrentHashMap<>(); @Override public List retrieveTexts(QueryVectorBo queryVectorBo) { @@ -51,6 +57,11 @@ public class KnowledgeRetrievalServiceImpl implements KnowledgeRetrievalService @Override public List 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 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 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 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 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 copyResults(List 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 results) { } } diff --git a/ruoyi-modules/ruoyi-chat/src/test/java/org/ruoyi/service/retrieval/impl/KnowledgeRetrievalServiceRegressionTest.java b/ruoyi-modules/ruoyi-chat/src/test/java/org/ruoyi/service/retrieval/impl/KnowledgeRetrievalServiceRegressionTest.java new file mode 100644 index 00000000..0841af78 --- /dev/null +++ b/ruoyi-modules/ruoyi-chat/src/test/java/org/ruoyi/service/retrieval/impl/KnowledgeRetrievalServiceRegressionTest.java @@ -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 result = (List) 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 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(); + } +}