feat(knowledge): S6c 检索多 dataset 逐库扇出+按相似度合并(Dify 单库端点适配)

Dify 的 retrieve 是单 dataset 端点、一次只检索一个库;把 MuseKnowledgeRetrievalApiImpl 从一次
多库调用(RAGFlow 原生多 dataset_ids)改为逐授权库检索:每次 chunk 直接归属本次检索的授权来源
(Dify 单库响应无 dataset_id、无法事后反查),再按相似度降序合并、截断到全局 topK。单库失败降级
跳过并记可追溯日志、全部失败回传代表性失败类。顺序执行(非并行)保留租户 ThreadLocal 上下文,
避免并行池线程丢上下文的隔离风险;内测期单作品 KB 数少延迟可接受,并行留作后续优化。授权双门
/§5.3 合同字段/同库多来源去重语义不变。

测试 18/0/0/0:15 个原测试(含 SPIKE 越权 + uRetrieve)证明单来源行为逐字不变,新增多库合并排序
/单库失败降级/全库失败三例。

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
lili 2026-07-08 03:12:55 -07:00
parent 09ff354e2f
commit a8f5e912f5
2 changed files with 148 additions and 27 deletions

View File

@ -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)},<b>不抛异常阻断主生成链</b></p>
*/
@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<String> datasetIds = new ArrayList<>();
Map<String, AuthorizedSource> 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无法事后反查),再按相似度降序合并截断到全局 topKRAGFlow 亦兼容单库调用
// 顺序执行(非并行):检索处于租户 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<RetrievedChunk> merged = new ArrayList<>();
Set<String> 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<RetrievedChunk> 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<RetrievedChunk> 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<RetrievedChunk> parseChunks(RuntimeResult result, Map<String, AuthorizedSource> datasetToSource,
AuthorizedSource primary) {
/** 解析单次检索(单库)响应 data.chunks[] → 直接归属【本次检索的授权来源】补齐 §5.3 合同字段。 */
private List<RetrievedChunk> parseChunks(RuntimeResult result, AuthorizedSource source) {
List<RetrievedChunk> 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(),

View File

@ -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.7ds-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不是业务代码不改任何检索/绑定主路径