feat: 支持工作流知识库多库检索
This commit is contained in:
@@ -52,11 +52,17 @@ public record WorkflowDesignerOptionsView(
|
||||
* @param id 知识库 ID
|
||||
* @param title 知识库标题
|
||||
* @param description 知识库描述
|
||||
* @param vectorEmbedModelId Embedding 模型 ID
|
||||
* @param dimensionOfVectorModel 向量维度
|
||||
* @param vectorStoreEnabled 是否可用于向量检索
|
||||
*/
|
||||
public record KnowledgeOption(
|
||||
@JsonSerialize(using = ToStringSerializer.class) BigInteger id,
|
||||
String title,
|
||||
String description
|
||||
String description,
|
||||
@JsonSerialize(using = ToStringSerializer.class) BigInteger vectorEmbedModelId,
|
||||
Integer dimensionOfVectorModel,
|
||||
Boolean vectorStoreEnabled
|
||||
) {
|
||||
}
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ import com.mybatisflex.core.query.QueryWrapper;
|
||||
import org.springframework.stereotype.Service;
|
||||
import tech.easyflow.admin.model.ai.WorkflowDesignerOptionsView;
|
||||
import tech.easyflow.ai.easyagentsflow.service.WorkflowDatacenterContentService;
|
||||
import tech.easyflow.ai.easyagentsflow.knowledge.WorkflowKnowledgeContractService;
|
||||
import tech.easyflow.ai.entity.DocumentCollection;
|
||||
import tech.easyflow.ai.entity.Model;
|
||||
import tech.easyflow.ai.entity.ModelProvider;
|
||||
@@ -75,6 +76,7 @@ public class WorkflowDesignerOptionService {
|
||||
private final DatacenterSourceService datacenterSourceService;
|
||||
private final DatacenterDatasetRegistryService datacenterDatasetRegistryService;
|
||||
private final DatacenterDatasetQueryService datacenterDatasetQueryService;
|
||||
private final WorkflowKnowledgeContractService workflowKnowledgeContractService;
|
||||
|
||||
/**
|
||||
* 创建工作流设计器选项服务。
|
||||
@@ -93,6 +95,7 @@ public class WorkflowDesignerOptionService {
|
||||
* @param datacenterSourceService 数据源服务
|
||||
* @param datacenterDatasetRegistryService 数据集注册服务
|
||||
* @param datacenterDatasetQueryService 数据集查询服务
|
||||
* @param workflowKnowledgeContractService 工作流知识库契约服务
|
||||
*/
|
||||
public WorkflowDesignerOptionService(
|
||||
ModelService modelService,
|
||||
@@ -108,7 +111,8 @@ public class WorkflowDesignerOptionService {
|
||||
ResourceAccessService resourceAccessService,
|
||||
DatacenterSourceService datacenterSourceService,
|
||||
DatacenterDatasetRegistryService datacenterDatasetRegistryService,
|
||||
DatacenterDatasetQueryService datacenterDatasetQueryService) {
|
||||
DatacenterDatasetQueryService datacenterDatasetQueryService,
|
||||
WorkflowKnowledgeContractService workflowKnowledgeContractService) {
|
||||
this.modelService = modelService;
|
||||
this.documentCollectionService = documentCollectionService;
|
||||
this.pluginService = pluginService;
|
||||
@@ -123,6 +127,7 @@ public class WorkflowDesignerOptionService {
|
||||
this.datacenterSourceService = datacenterSourceService;
|
||||
this.datacenterDatasetRegistryService = datacenterDatasetRegistryService;
|
||||
this.datacenterDatasetQueryService = datacenterDatasetQueryService;
|
||||
this.workflowKnowledgeContractService = workflowKnowledgeContractService;
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -163,6 +168,7 @@ public class WorkflowDesignerOptionService {
|
||||
LoginAccount account = requireAccount();
|
||||
Set<BigInteger> modelIds = new HashSet<>();
|
||||
Set<BigInteger> knowledgeIds = new HashSet<>();
|
||||
List<List<BigInteger>> knowledgeGroups = new ArrayList<>();
|
||||
Set<BigInteger> checkedPluginItemIds = new HashSet<>();
|
||||
Set<BigInteger> checkedWorkflowIds = new HashSet<>();
|
||||
Set<BigInteger> checkedSourceIds = new HashSet<>();
|
||||
@@ -177,14 +183,22 @@ public class WorkflowDesignerOptionService {
|
||||
if (data == null) {
|
||||
continue;
|
||||
}
|
||||
String nodeType = data.getString("type");
|
||||
String nodeType = node.getString("type");
|
||||
String dataType = data.getString("type");
|
||||
if (nodeType != null && !nodeType.isBlank()
|
||||
&& dataType != null && !dataType.isBlank()
|
||||
&& !Objects.equals(nodeType, dataType)) {
|
||||
throw new BusinessException("工作流节点类型与节点数据类型不一致");
|
||||
}
|
||||
if (nodeType == null || nodeType.isBlank()) {
|
||||
nodeType = node.getString("type");
|
||||
nodeType = dataType;
|
||||
}
|
||||
if ("llmNode".equals(nodeType)) {
|
||||
addReferenceId(modelIds, readReferenceId(data, "llmId", "模型"));
|
||||
} else if ("knowledgeNode".equals(nodeType)) {
|
||||
addReferenceId(knowledgeIds, readReferenceId(data, "knowledgeId", "知识库"));
|
||||
List<BigInteger> nodeKnowledgeIds = readKnowledgeReferenceIds(data);
|
||||
knowledgeIds.addAll(nodeKnowledgeIds);
|
||||
knowledgeGroups.add(nodeKnowledgeIds);
|
||||
} else if ("plugin-node".equals(nodeType)) {
|
||||
BigInteger pluginItemId = readReferenceId(data, "pluginId", "插件工具");
|
||||
if (pluginItemId != null && checkedPluginItemIds.add(pluginItemId)) {
|
||||
@@ -198,6 +212,8 @@ public class WorkflowDesignerOptionService {
|
||||
}
|
||||
assertModelReferences(modelIds, account);
|
||||
assertKnowledgeReferences(knowledgeIds, account);
|
||||
workflowKnowledgeContractService.assertMultiKnowledgeContracts(
|
||||
knowledgeGroups, account.getTenantId());
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -460,17 +476,55 @@ public class WorkflowDesignerOptionService {
|
||||
}
|
||||
|
||||
private List<WorkflowDesignerOptionsView.KnowledgeOption> listKnowledgeOptions(LoginAccount account) {
|
||||
return documentCollectionService.list(QueryWrapper.create()
|
||||
List<DocumentCollection> collections = documentCollectionService.list(QueryWrapper.create()
|
||||
.eq(DocumentCollection::getTenantId, account.getTenantId())
|
||||
.orderBy(DocumentCollection::getModified, false))
|
||||
.stream()
|
||||
.orderBy(DocumentCollection::getModified, false));
|
||||
Set<BigInteger> vectorReadyIds = workflowKnowledgeContractService
|
||||
.findVectorReadyKnowledgeIds(collections, account.getTenantId());
|
||||
return collections.stream()
|
||||
.filter(item -> resourceAccessService.canAccess(
|
||||
account, CategoryResourceType.KNOWLEDGE, item, ResourceAction.USE))
|
||||
.map(item -> new WorkflowDesignerOptionsView.KnowledgeOption(
|
||||
item.getId(), item.getTitle(), item.getDescription()))
|
||||
item.getId(),
|
||||
item.getTitle(),
|
||||
item.getDescription(),
|
||||
item.getVectorEmbedModelId(),
|
||||
item.getDimensionOfVectorModel(),
|
||||
vectorReadyIds.contains(item.getId())))
|
||||
.toList();
|
||||
}
|
||||
|
||||
private List<BigInteger> readKnowledgeReferenceIds(JSONObject data) {
|
||||
if (data.containsKey("knowledgeIds")) {
|
||||
Object rawIds = data.get("knowledgeIds");
|
||||
if (!(rawIds instanceof JSONArray ids) || ids.isEmpty()) {
|
||||
throw new BusinessException("知识库节点至少需要选择一个知识库");
|
||||
}
|
||||
List<BigInteger> result = new ArrayList<>();
|
||||
for (Object id : ids) {
|
||||
BigInteger parsed = parseReferenceId(id, "知识库");
|
||||
if (result.contains(parsed)) {
|
||||
throw new BusinessException("知识库节点不能重复选择同一知识库");
|
||||
}
|
||||
result.add(parsed);
|
||||
}
|
||||
return result;
|
||||
}
|
||||
BigInteger legacyId = readReferenceId(data, "knowledgeId", "知识库");
|
||||
return legacyId == null ? List.of() : List.of(legacyId);
|
||||
}
|
||||
|
||||
private BigInteger parseReferenceId(Object value, String resourceName) {
|
||||
if (value == null || String.valueOf(value).isBlank()) {
|
||||
throw new BusinessException(resourceName + "ID不能为空");
|
||||
}
|
||||
try {
|
||||
return new BigInteger(String.valueOf(value));
|
||||
} catch (NumberFormatException exception) {
|
||||
throw new BusinessException(resourceName + "ID无效");
|
||||
}
|
||||
}
|
||||
|
||||
private void addReferenceId(Set<BigInteger> resourceIds, BigInteger resourceId) {
|
||||
if (resourceId != null) {
|
||||
resourceIds.add(resourceId);
|
||||
|
||||
@@ -6,7 +6,10 @@ import org.mockito.MockedStatic;
|
||||
import org.testng.Assert;
|
||||
import org.testng.annotations.Test;
|
||||
import tech.easyflow.ai.easyagentsflow.service.WorkflowDatacenterContentService;
|
||||
import tech.easyflow.ai.easyagentsflow.knowledge.WorkflowKnowledgeContractService;
|
||||
import tech.easyflow.ai.entity.Model;
|
||||
import tech.easyflow.ai.entity.ModelProvider;
|
||||
import tech.easyflow.ai.entity.DocumentCollection;
|
||||
import tech.easyflow.ai.entity.Workflow;
|
||||
import tech.easyflow.ai.plugin.workflow.snapshot.WorkflowPluginSnapshotResolver;
|
||||
import tech.easyflow.ai.service.DocumentCollectionService;
|
||||
@@ -129,6 +132,96 @@ public class WorkflowDesignerOptionServiceTest {
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldAcceptCompatibleMultiKnowledgeReferences() {
|
||||
ModelService modelService = mock(ModelService.class);
|
||||
DocumentCollectionService knowledgeService =
|
||||
mock(DocumentCollectionService.class);
|
||||
ResourceAccessService accessService = mock(ResourceAccessService.class);
|
||||
when(accessService.canAccess(any(), any(), any(), any()))
|
||||
.thenReturn(true);
|
||||
when(knowledgeService.listByIds(any()))
|
||||
.thenReturn(List.of(
|
||||
knowledge(1, 7, 3),
|
||||
knowledge(2, 7, 3)));
|
||||
when(modelService.listModelInstances(any()))
|
||||
.thenReturn(List.of(embeddingModel(7)));
|
||||
WorkflowDesignerOptionService service = createService(
|
||||
modelService,
|
||||
knowledgeService,
|
||||
mock(DatacenterSourceService.class),
|
||||
mock(WorkflowService.class),
|
||||
accessService);
|
||||
|
||||
try (MockedStatic<SaTokenUtil> saToken = mockStatic(SaTokenUtil.class)) {
|
||||
saToken.when(SaTokenUtil::getLoginAccount).thenReturn(loginAccount());
|
||||
|
||||
service.assertContentReferences("""
|
||||
{"nodes":[{"type":"knowledgeNode","data":{
|
||||
"knowledgeIds":["1","2"],"retrievalMode":"VECTOR"
|
||||
}}]}
|
||||
""");
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectIncompatibleMultiKnowledgeReferences() {
|
||||
ModelService modelService = mock(ModelService.class);
|
||||
DocumentCollectionService knowledgeService =
|
||||
mock(DocumentCollectionService.class);
|
||||
ResourceAccessService accessService = mock(ResourceAccessService.class);
|
||||
when(accessService.canAccess(any(), any(), any(), any()))
|
||||
.thenReturn(true);
|
||||
when(knowledgeService.listByIds(any()))
|
||||
.thenReturn(List.of(
|
||||
knowledge(1, 7, 3),
|
||||
knowledge(2, 8, 3)));
|
||||
when(modelService.listModelInstances(any()))
|
||||
.thenReturn(List.of(embeddingModel(7), embeddingModel(8)));
|
||||
WorkflowDesignerOptionService service = createService(
|
||||
modelService,
|
||||
knowledgeService,
|
||||
mock(DatacenterSourceService.class),
|
||||
mock(WorkflowService.class),
|
||||
accessService);
|
||||
|
||||
try (MockedStatic<SaTokenUtil> saToken = mockStatic(SaTokenUtil.class)) {
|
||||
saToken.when(SaTokenUtil::getLoginAccount).thenReturn(loginAccount());
|
||||
|
||||
BusinessException exception = Assert.expectThrows(
|
||||
BusinessException.class,
|
||||
() -> service.assertContentReferences("""
|
||||
{"nodes":[{"type":"knowledgeNode","data":{
|
||||
"knowledgeIds":["1","2"],"retrievalMode":"VECTOR"
|
||||
}}]}
|
||||
"""));
|
||||
|
||||
Assert.assertTrue(exception.getMessage().contains("Embedding"));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectConflictingRootAndDataNodeTypes() {
|
||||
WorkflowDesignerOptionService service = createService(
|
||||
mock(ModelService.class),
|
||||
mock(DocumentCollectionService.class),
|
||||
mock(DatacenterSourceService.class));
|
||||
|
||||
try (MockedStatic<SaTokenUtil> saToken = mockStatic(SaTokenUtil.class)) {
|
||||
saToken.when(SaTokenUtil::getLoginAccount).thenReturn(loginAccount());
|
||||
|
||||
BusinessException exception = Assert.expectThrows(
|
||||
BusinessException.class,
|
||||
() -> service.assertContentReferences("""
|
||||
{"nodes":[{"type":"knowledgeNode","data":{
|
||||
"type":"llmNode","knowledgeId":"1"
|
||||
}}]}
|
||||
"""));
|
||||
|
||||
Assert.assertTrue(exception.getMessage().contains("类型"));
|
||||
}
|
||||
}
|
||||
|
||||
private WorkflowDesignerOptionService createService(
|
||||
ModelService modelService,
|
||||
DocumentCollectionService knowledgeService,
|
||||
@@ -142,6 +235,20 @@ public class WorkflowDesignerOptionServiceTest {
|
||||
DatacenterSourceService sourceService,
|
||||
WorkflowService workflowService) {
|
||||
ResourceAccessService resourceAccessService = mock(ResourceAccessService.class);
|
||||
return createService(
|
||||
modelService,
|
||||
knowledgeService,
|
||||
sourceService,
|
||||
workflowService,
|
||||
resourceAccessService);
|
||||
}
|
||||
|
||||
private WorkflowDesignerOptionService createService(
|
||||
ModelService modelService,
|
||||
DocumentCollectionService knowledgeService,
|
||||
DatacenterSourceService sourceService,
|
||||
WorkflowService workflowService,
|
||||
ResourceAccessService resourceAccessService) {
|
||||
return new WorkflowDesignerOptionService(
|
||||
modelService,
|
||||
knowledgeService,
|
||||
@@ -156,10 +263,40 @@ public class WorkflowDesignerOptionServiceTest {
|
||||
resourceAccessService,
|
||||
sourceService,
|
||||
mock(DatacenterDatasetRegistryService.class),
|
||||
mock(DatacenterDatasetQueryService.class)
|
||||
mock(DatacenterDatasetQueryService.class),
|
||||
new WorkflowKnowledgeContractService(
|
||||
knowledgeService, modelService)
|
||||
);
|
||||
}
|
||||
|
||||
private DocumentCollection knowledge(
|
||||
long id, long embeddingModelId, int dimension) {
|
||||
DocumentCollection collection = new DocumentCollection();
|
||||
collection.setId(BigInteger.valueOf(id));
|
||||
collection.setTenantId(BigInteger.valueOf(100));
|
||||
collection.setVectorEmbedModelId(BigInteger.valueOf(embeddingModelId));
|
||||
collection.setDimensionOfVectorModel(dimension);
|
||||
collection.setVectorStoreEnable(true);
|
||||
collection.setVectorStoreCollection("collection_" + id);
|
||||
return collection;
|
||||
}
|
||||
|
||||
private Model embeddingModel(long id) {
|
||||
Model model = new Model();
|
||||
model.setId(BigInteger.valueOf(id));
|
||||
model.setTenantId(BigInteger.valueOf(100));
|
||||
model.setModelType(Model.MODEL_TYPES[1]);
|
||||
model.setProviderId(BigInteger.ONE);
|
||||
model.setModelName("embedding-" + id);
|
||||
model.setEndpoint("https://embedding.example");
|
||||
model.setRequestPath("/v1/embeddings");
|
||||
ModelProvider provider = new ModelProvider();
|
||||
provider.setId(BigInteger.ONE);
|
||||
provider.setProviderType("openai");
|
||||
model.setModelProvider(provider);
|
||||
return model;
|
||||
}
|
||||
|
||||
private LoginAccount loginAccount() {
|
||||
LoginAccount account = new LoginAccount();
|
||||
account.setId(BigInteger.ONE);
|
||||
|
||||
Reference in New Issue
Block a user