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),不是业务代码;不改任何检索/绑定主路径。