diff --git a/muse-cloud/muse-module-knowledge/muse-module-knowledge-server/src/main/java/cn/iocoder/muse/module/knowledge/api/MuseKnowledgeRetrievalApiImpl.java b/muse-cloud/muse-module-knowledge/muse-module-knowledge-server/src/main/java/cn/iocoder/muse/module/knowledge/api/MuseKnowledgeRetrievalApiImpl.java index 8586029e..bfd3c020 100644 --- a/muse-cloud/muse-module-knowledge/muse-module-knowledge-server/src/main/java/cn/iocoder/muse/module/knowledge/api/MuseKnowledgeRetrievalApiImpl.java +++ b/muse-cloud/muse-module-knowledge/muse-module-knowledge-server/src/main/java/cn/iocoder/muse/module/knowledge/api/MuseKnowledgeRetrievalApiImpl.java @@ -2,6 +2,7 @@ package cn.iocoder.muse.module.knowledge.api; import cn.iocoder.muse.framework.common.util.json.JsonUtils; import cn.iocoder.muse.module.knowledge.application.muse.facade.KnowledgeRuntimeClient; +import cn.iocoder.muse.module.knowledge.application.muse.facade.KnowledgeRuntimeClient.FailureClass; import cn.iocoder.muse.module.knowledge.application.muse.facade.KnowledgeRuntimeClient.RetrieveChunksCommand; import cn.iocoder.muse.module.knowledge.application.muse.facade.KnowledgeRuntimeClient.RuntimeResult; import cn.iocoder.muse.module.knowledge.application.muse.facade.KnowledgeRuntimeClient.Status; @@ -13,6 +14,7 @@ import cn.iocoder.muse.module.knowledge.dal.mysql.muse.MuseKnowledgeRagflowBindi import cn.iocoder.muse.module.knowledge.dal.mysql.muse.MuseKnowledgeSourceBindingProjectionMapper; import com.fasterxml.jackson.databind.JsonNode; import jakarta.annotation.Resource; +import lombok.extern.slf4j.Slf4j; import org.springframework.stereotype.Service; import org.springframework.util.StringUtils; @@ -36,6 +38,7 @@ import java.util.Set; * {@code empty(omittedReason)},不抛异常阻断主生成链

*/ @Service +@Slf4j public class MuseKnowledgeRetrievalApiImpl implements MuseKnowledgeRetrievalApi { /** 单 chunk 正文摘要最大长度(token 预算意识,防 prompt 撑爆;专题-03 §4)。 */ @@ -135,41 +138,74 @@ public class MuseKnowledgeRetrievalApiImpl implements MuseKnowledgeRetrievalApi ? "not_authorized" : "no_dataset"; return RetrievalResult.empty(reason, omitted); } - // 4. 一次跨所有授权 dataset 检索(RAGFlow 支持多 dataset_ids);chunk 再按 dataset 反查授权来源 - List datasetIds = new ArrayList<>(); - Map datasetToSource = new LinkedHashMap<>(); - for (AuthorizedSource source : authorized) { - if (datasetToSource.putIfAbsent(source.datasetId(), source) == null) { - datasetIds.add(source.datasetId()); - } - } - AuthorizedSource primary = authorized.get(0); + // 4. 逐授权 dataset 检索并按相似度合并(provider 无关): + // WHY 逐库而非一次多库——Dify 的 retrieve 是单 dataset 端点(见 DifyKnowledgeRuntimeClient), + // 一次只能检索一个库;逐库调用后每次 chunk 直接归属【本次检索的授权来源】(Dify 单库响应不含 + // dataset_id、无法事后反查),再按相似度降序合并、截断到全局 topK。RAGFlow 亦兼容单库调用。 + // 顺序执行(非并行):检索处于租户 ThreadLocal 上下文内,同线程保留上下文,避免并行池线程丢 + // 租户上下文的隔离风险;内测期单作品绑定 KB 数少、顺序延迟可接受,并行留作后续优化。 + // 同一 dataset 多来源只检一次、归属首个(与旧多库去重 putIfAbsent 语义一致)。 int topK = request.topK() != null && request.topK() > 0 ? request.topK() : DEFAULT_TOP_K; - RetrieveChunksCommand command = new RetrieveChunksCommand( - request.tenantId(), request.ownerUserId(), primary.kbId(), datasetIds, null, - request.question(), topK, DEFAULT_THRESHOLD, null, - request.correlationId(), null, 1); - RuntimeResult result = ragFlowClient.retrieveChunks(command); - if (result == null || result.status() != Status.SUCCEEDED) { - String reason = result == null || result.failureClass() == null - ? "retrieval_failed" : "retrieval_failed:" + result.failureClass().name(); - return RetrievalResult.empty(reason, omitted); + List merged = new ArrayList<>(); + Set retrievedDataset = new HashSet<>(); + FailureClass firstFailure = null; + boolean anySuccess = false; + for (AuthorizedSource source : authorized) { + if (!retrievedDataset.add(source.datasetId())) { + continue; + } + RetrieveChunksCommand command = new RetrieveChunksCommand( + request.tenantId(), request.ownerUserId(), source.kbId(), + List.of(source.datasetId()), null, request.question(), topK, + DEFAULT_THRESHOLD, null, request.correlationId(), null, 1); + RuntimeResult result; + try { + result = ragFlowClient.retrieveChunks(command); + } catch (RuntimeException ex) { + // 单库检索异常不阻断其它库:记失败日志、降级跳过(fail-closed,不抛断主生成链) + if (firstFailure == null) { + firstFailure = FailureClass.UNKNOWN_FAILURE; + } + log.warn("[retrieveForWork][单库检索异常 correlationId={}, kbId={}, datasetId={}]", + request.correlationId(), source.kbId(), source.datasetId(), ex); + continue; + } + if (result == null || result.status() != Status.SUCCEEDED) { + FailureClass fc = result == null ? null : result.failureClass(); + if (firstFailure == null && fc != null) { + firstFailure = fc; + } + log.info("[retrieveForWork][单库检索失败 correlationId={}, kbId={}, datasetId={}, failureClass={}]", + request.correlationId(), source.kbId(), source.datasetId(), fc); + continue; + } + anySuccess = true; + merged.addAll(parseChunks(result, source)); } - // 5. 解析 data.chunks[] → 补 §5.3 合同字段(按 dataset 反查授权来源) - List chunks = parseChunks(result, datasetToSource, primary); - if (chunks.isEmpty()) { + if (merged.isEmpty()) { + if (!anySuccess) { + // 全部授权库检索失败:回传代表性失败类(单库场景等价旧 retrieval_failed:CLASS) + String reason = firstFailure == null ? "retrieval_failed" : "retrieval_failed:" + firstFailure.name(); + return RetrievalResult.empty(reason, omitted); + } + // 有库检索成功但都无命中 return RetrievalResult.empty("no_chunk", omitted); } - return RetrievalResult.ok(chunks, omitted); + // 5. 合并排序:相似度降序(相似度缺失排最后),截断到全局 topK + merged.sort((a, b) -> Double.compare( + b.similarity() == null ? Double.NEGATIVE_INFINITY : b.similarity(), + a.similarity() == null ? Double.NEGATIVE_INFINITY : a.similarity())); + List topChunks = merged.size() > topK + ? new ArrayList<>(merged.subList(0, topK)) : merged; + return RetrievalResult.ok(topChunks, omitted); } catch (RuntimeException ex) { // fail-closed 降级:检索任何异常都不阻断主生成链 return RetrievalResult.empty("retrieval_error"); } } - /** 从 RAGFlow 检索响应树 data.chunks[] 解析脱敏片段;按 dataset 反查授权来源补齐 §5.3 字段,反查不到归 primary。 */ - private List parseChunks(RuntimeResult result, Map datasetToSource, - AuthorizedSource primary) { + /** 解析单次检索(单库)响应 data.chunks[] → 直接归属【本次检索的授权来源】补齐 §5.3 合同字段。 */ + private List parseChunks(RuntimeResult result, AuthorizedSource source) { List chunks = new ArrayList<>(); Object responseBody = result.summary() == null ? null : result.summary().get("responseBody"); JsonNode tree = toJsonNode(responseBody); @@ -185,10 +221,12 @@ public class MuseKnowledgeRetrievalApiImpl implements MuseKnowledgeRetrievalApi if (!StringUtils.hasText(content)) { continue; } + // datasetId:优先 chunk 自带字段(RAGFlow 多库遗留),缺失(Dify 单库响应无 dataset_id)则用本次检索来源的 datasetId String datasetId = firstText(chunk, "dataset_id", "kb_id"); + if (!StringUtils.hasText(datasetId)) { + datasetId = source.datasetId(); + } String documentId = firstText(chunk, "document_id", "doc_id"); - AuthorizedSource source = datasetId != null && datasetToSource.containsKey(datasetId) - ? datasetToSource.get(datasetId) : primary; Double similarity = chunk.path("similarity").isNumber() ? chunk.path("similarity").asDouble() : null; chunks.add(new RetrievedChunk(source.kbId(), datasetId, documentId, summarize(content), similarity, source.sourceOwner(), source.sourceObjectVersion(), source.authorizationSnapshotId(), diff --git a/muse-cloud/muse-module-knowledge/muse-module-knowledge-server/src/test/java/cn/iocoder/muse/module/knowledge/api/MuseKnowledgeRetrievalApiImplTest.java b/muse-cloud/muse-module-knowledge/muse-module-knowledge-server/src/test/java/cn/iocoder/muse/module/knowledge/api/MuseKnowledgeRetrievalApiImplTest.java index b235c8d7..a099fa67 100644 --- a/muse-cloud/muse-module-knowledge/muse-module-knowledge-server/src/test/java/cn/iocoder/muse/module/knowledge/api/MuseKnowledgeRetrievalApiImplTest.java +++ b/muse-cloud/muse-module-knowledge/muse-module-knowledge-server/src/test/java/cn/iocoder/muse/module/knowledge/api/MuseKnowledgeRetrievalApiImplTest.java @@ -30,6 +30,7 @@ import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertNull; import static org.junit.jupiter.api.Assertions.assertTrue; import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.times; import static org.mockito.Mockito.verify; import static org.mockito.Mockito.verifyNoInteractions; import static org.mockito.Mockito.when; @@ -230,6 +231,88 @@ class MuseKnowledgeRetrievalApiImplTest extends BaseMockitoUnitTest { assertEquals("no_chunk", result.omittedReason()); } + // ======================================================================================== + // S6c 多授权 dataset 逐库扇出 + 按相似度合并(Dify retrieve 单库端点,一次只检一个库)。 + // ======================================================================================== + + /** 两授权来源(各自 dataset)逐库检索,按相似度降序合并,每 chunk 归属各自来源的 §5.3 字段。 */ + @Test + void should_mergeAndSortChunksAcrossDatasets_withPerSourceContractFields() { + when(bindingMapper.selectActiveByWorkId(4001L)) + .thenReturn(List.of(authorizedBinding(5001L), authorizedBinding(5002L))); + when(projectionMapper.selectActiveByWorkId(4001L)) + .thenReturn(List.of(projection(5001L, "active"), projection(5002L, "active"))); + when(ragflowBindingMapper.selectActiveDatasetByKbId(5001L)).thenReturn(dataset("ds-A")); + when(ragflowBindingMapper.selectActiveDatasetByKbId(5002L)).thenReturn(dataset("ds-B")); + // 每库单独响应:ds-A 相似度 0.7、ds-B 相似度 0.9(合并后 B 应排前) + when(ragFlowClient.retrieveChunks(any(RetrieveChunksCommand.class))).thenAnswer(inv -> { + RetrieveChunksCommand cmd = inv.getArgument(0); + if ("ds-A".equals(cmd.ragflowDatasetIds().get(0))) { + return retrievalSuccess("{\"data\":{\"chunks\":[{\"content\":\"来自A\",\"document_id\":\"a1\",\"kb_id\":\"ds-A\",\"similarity\":0.7}]}}"); + } + return retrievalSuccess("{\"data\":{\"chunks\":[{\"content\":\"来自B\",\"document_id\":\"b1\",\"kb_id\":\"ds-B\",\"similarity\":0.9}]}}"); + }); + + RetrievalResult result = retrievalApi.retrieveForWork(req()); + + assertTrue(result.hasChunks()); + assertEquals(2, result.chunks().size()); + // 相似度降序:B(0.9)在前、A(0.7)在后,各自归属正确来源 + assertEquals("来自B", result.chunks().get(0).contentSummary()); + assertEquals(5002L, result.chunks().get(0).sourceKbId()); + assertEquals("auth-5002", result.chunks().get(0).authorizationSnapshotId()); + assertEquals("来自A", result.chunks().get(1).contentSummary()); + assertEquals(5001L, result.chunks().get(1).sourceKbId()); + assertEquals("auth-5001", result.chunks().get(1).authorizationSnapshotId()); + // 逐库扇出:两库各一次检索 + verify(ragFlowClient, times(2)).retrieveChunks(any(RetrieveChunksCommand.class)); + } + + /** 单库检索失败降级:一库失败、另一库成功 → 返回成功库命中,不整体阻断(fail-closed 降级)。 */ + @Test + void should_degradeToSuccessfulDatasets_whenOneDatasetFails() { + when(bindingMapper.selectActiveByWorkId(4001L)) + .thenReturn(List.of(authorizedBinding(5001L), authorizedBinding(5002L))); + when(projectionMapper.selectActiveByWorkId(4001L)) + .thenReturn(List.of(projection(5001L, "active"), projection(5002L, "active"))); + when(ragflowBindingMapper.selectActiveDatasetByKbId(5001L)).thenReturn(dataset("ds-A")); + when(ragflowBindingMapper.selectActiveDatasetByKbId(5002L)).thenReturn(dataset("ds-B")); + when(ragFlowClient.retrieveChunks(any(RetrieveChunksCommand.class))).thenAnswer(inv -> { + RetrieveChunksCommand cmd = inv.getArgument(0); + if ("ds-A".equals(cmd.ragflowDatasetIds().get(0))) { + return new RuntimeResult(null, Operation.RETRIEVE_CHUNKS, "c", null, 1, 5L, + Status.FAILED, FailureClass.TIMEOUT, null, "r", List.of(), Map.of()); + } + return retrievalSuccess("{\"data\":{\"chunks\":[{\"content\":\"来自B\",\"document_id\":\"b1\",\"kb_id\":\"ds-B\",\"similarity\":0.9}]}}"); + }); + + RetrievalResult result = retrievalApi.retrieveForWork(req()); + + assertTrue(result.hasChunks(), "单库失败应降级、不整体阻断"); + assertEquals(1, result.chunks().size()); + assertEquals("来自B", result.chunks().get(0).contentSummary()); + assertEquals("ok", result.status()); + } + + /** 全部授权库检索失败 → retrieval_failed(代表性失败类),不误报无结果。 */ + @Test + void should_returnRetrievalFailed_whenAllDatasetsFail() { + when(bindingMapper.selectActiveByWorkId(4001L)) + .thenReturn(List.of(authorizedBinding(5001L), authorizedBinding(5002L))); + when(projectionMapper.selectActiveByWorkId(4001L)) + .thenReturn(List.of(projection(5001L, "active"), projection(5002L, "active"))); + when(ragflowBindingMapper.selectActiveDatasetByKbId(5001L)).thenReturn(dataset("ds-A")); + when(ragflowBindingMapper.selectActiveDatasetByKbId(5002L)).thenReturn(dataset("ds-B")); + RuntimeResult failed = new RuntimeResult(null, Operation.RETRIEVE_CHUNKS, "c", null, 1, 5L, + Status.FAILED, FailureClass.RAGFLOW_UNAVAILABLE, null, "r", List.of(), Map.of()); + when(ragFlowClient.retrieveChunks(any(RetrieveChunksCommand.class))).thenReturn(failed); + + RetrievalResult result = retrievalApi.retrieveForWork(req()); + + assertFalse(result.hasChunks()); + assertEquals("retrieval_failed:RAGFLOW_UNAVAILABLE", result.omittedReason()); + } + // ======================================================================================== // U0 SPIKE(临时-03 §U0 / 决策门 D1):把"多租户共享同一 RAGFlow dataset 必越权"钉成可复现证据。 // 这是测试代码(spike),不是业务代码;不改任何检索/绑定主路径。