feat: 支持工作流多知识库向量检索

This commit is contained in:
2026-09-04 17:40:15 +08:00
parent e959a772b5
commit 45c708a212
15 changed files with 2092 additions and 129 deletions

View File

@@ -18,6 +18,7 @@ package com.easyagents.flow.core.knowledge;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;
public class KnowledgeManager {
@@ -51,4 +52,20 @@ public class KnowledgeManager {
}
return null;
}
/**
* 将完整知识检索请求交给首个能够处理它的 Provider。
*
* @param request 检索请求
* @return 节点输出;没有 Provider 能处理时返回 null
*/
public Map<String, Object> search(KnowledgeSearchRequest request) {
for (KnowledgeProvider provider : providers) {
Map<String, Object> result = provider.search(request);
if (result != null) {
return result;
}
}
return null;
}
}

View File

@@ -15,6 +15,37 @@
*/
package com.easyagents.flow.core.knowledge;
import com.easyagents.flow.core.util.Maps;
import java.util.List;
import java.util.Map;
public interface KnowledgeProvider {
Knowledge getKnowledge(Object id);
/**
* 执行完整的知识库节点检索请求。
*
* <p>默认实现保留单知识库兼容。需要跨知识库汇总的业务 Provider
* 应覆盖本方法并返回完整节点输出。</p>
*
* @param request 检索请求
* @return 节点输出;当前 Provider 不支持该请求时返回 null
*/
default Map<String, Object> search(KnowledgeSearchRequest request) {
if (request == null || request.getKnowledgeIds().size() != 1) {
return null;
}
Object knowledgeId = request.getKnowledgeIds().get(0);
Knowledge knowledge = getKnowledge(knowledgeId);
if (knowledge == null) {
return null;
}
List<Map<String, Object>> documents = knowledge.search(
request.getKeyword(),
request.getLimit(),
request.getKnowledgeNode(),
request.getChain());
return Maps.of("documents", documents);
}
}

View File

@@ -0,0 +1,67 @@
/**
* Copyright (c) 2025-2026, Michael Yang 杨福海 (fuhai999@gmail.com).
* <p>
* Licensed under the GNU Lesser General Public License (LGPL) ,Version 3.0.
*/
package com.easyagents.flow.core.knowledge;
import com.easyagents.flow.core.chain.Chain;
import com.easyagents.flow.core.node.KnowledgeNode;
import java.util.Collections;
import java.util.ArrayList;
import java.util.List;
/**
* 工作流知识库节点的完整检索请求。
*/
public class KnowledgeSearchRequest {
private final List<Object> knowledgeIds;
private final String keyword;
private final int limit;
private final String retrievalMode;
private final KnowledgeNode knowledgeNode;
private final Chain chain;
public KnowledgeSearchRequest(
List<Object> knowledgeIds,
String keyword,
int limit,
String retrievalMode,
KnowledgeNode knowledgeNode,
Chain chain) {
this.knowledgeIds = knowledgeIds == null
? Collections.emptyList()
: Collections.unmodifiableList(new ArrayList<>(knowledgeIds));
this.keyword = keyword;
this.limit = limit;
this.retrievalMode = retrievalMode;
this.knowledgeNode = knowledgeNode;
this.chain = chain;
}
public List<Object> getKnowledgeIds() {
return knowledgeIds;
}
public String getKeyword() {
return keyword;
}
public int getLimit() {
return limit;
}
public String getRetrievalMode() {
return retrievalMode;
}
public KnowledgeNode getKnowledgeNode() {
return knowledgeNode;
}
public Chain getChain() {
return chain;
}
}

View File

@@ -17,15 +17,15 @@ package com.easyagents.flow.core.node;
import com.easyagents.flow.core.chain.Chain;
import com.easyagents.flow.core.chain.ChainState;
import com.easyagents.flow.core.knowledge.Knowledge;
import com.easyagents.flow.core.knowledge.KnowledgeManager;
import com.easyagents.flow.core.util.Maps;
import com.easyagents.flow.core.knowledge.KnowledgeSearchRequest;
import com.easyagents.flow.core.util.StringUtil;
import com.easyagents.flow.core.util.TextTemplate;
import org.slf4j.Logger;
import java.util.Arrays;
import java.util.ArrayList;
import java.util.Collections;
import java.util.LinkedHashSet;
import java.util.List;
import java.util.Map;
@@ -36,6 +36,7 @@ public class KnowledgeNode extends BaseNode {
private static final Logger logger = org.slf4j.LoggerFactory.getLogger(KnowledgeNode.class);
private Object knowledgeId;
private List<Object> knowledgeIds = new ArrayList<>();
private String keyword;
private String limit;
private String retrievalMode = "HYBRID";
@@ -48,6 +49,32 @@ public class KnowledgeNode extends BaseNode {
this.knowledgeId = knowledgeId;
}
/**
* 获取规范化的知识库集合,兼容历史单值字段。
*
* @return 去重后的知识库 ID
*/
public List<Object> getKnowledgeIds() {
if (knowledgeIds != null && !knowledgeIds.isEmpty()) {
return Collections.unmodifiableList(knowledgeIds);
}
return knowledgeId == null
? Collections.emptyList()
: Collections.singletonList(knowledgeId);
}
public void setKnowledgeIds(List<?> knowledgeIds) {
LinkedHashSet<Object> normalized = new LinkedHashSet<>();
if (knowledgeIds != null) {
for (Object id : knowledgeIds) {
if (id != null && StringUtil.hasText(String.valueOf(id))) {
normalized.add(id);
}
}
}
this.knowledgeIds = new ArrayList<>(normalized);
}
public String getKeyword() {
return keyword;
}
@@ -88,25 +115,44 @@ public class KnowledgeNode extends BaseNode {
if (StringUtil.hasText(realLimitString)) {
try {
realLimit = Integer.parseInt(realLimitString);
} catch (Exception e) {
logger.error(e.toString(), e);
} catch (NumberFormatException exception) {
throw new IllegalArgumentException(
"知识库节点最终返回条数必须为正整数", exception);
}
}
Knowledge knowledge = KnowledgeManager.getInstance().getKnowledge(knowledgeId);
if (knowledge == null) {
return Collections.emptyMap();
if (realLimit <= 0) {
throw new IllegalArgumentException(
"知识库节点最终返回条数必须为正整数");
}
List<Map<String, Object>> result = knowledge.search(realKeyword, realLimit, this, chain);
return Maps.of("documents", result);
List<Object> resolvedKnowledgeIds = getKnowledgeIds();
if (resolvedKnowledgeIds.isEmpty()) {
throw new IllegalArgumentException("知识库节点至少需要选择一个知识库");
}
if (resolvedKnowledgeIds.size() > 1
&& !"VECTOR".equalsIgnoreCase(retrievalMode)) {
throw new IllegalArgumentException("多知识库检索仅支持 VECTOR 模式");
}
Map<String, Object> result = KnowledgeManager.getInstance().search(
new KnowledgeSearchRequest(
resolvedKnowledgeIds,
realKeyword,
realLimit,
retrievalMode,
this,
chain));
if (result == null) {
throw new IllegalStateException("没有可用的知识库 Provider");
}
return result;
}
@Override
public String toString() {
return "KnowledgeNode{" +
"knowledgeId=" + knowledgeId +
", knowledgeIds=" + knowledgeIds +
", keyword='" + keyword + '\'' +
", limit='" + limit + '\'' +
", retrievalMode='" + retrievalMode + '\'' +

View File

@@ -16,15 +16,41 @@
package com.easyagents.flow.core.parser.impl;
import com.alibaba.fastjson.JSONObject;
import com.alibaba.fastjson.JSONArray;
import com.easyagents.flow.core.node.KnowledgeNode;
import com.easyagents.flow.core.parser.BaseNodeParser;
import java.util.ArrayList;
public class KnowledgeNodeParser extends BaseNodeParser<KnowledgeNode> {
@Override
public KnowledgeNode doParse(JSONObject root, JSONObject data, JSONObject chainJSONObject) {
KnowledgeNode knowledgeNode = new KnowledgeNode();
knowledgeNode.setKnowledgeId(data.get("knowledgeId"));
if (data.containsKey("knowledgeIds")) {
Object rawIds = data.get("knowledgeIds");
if (!(rawIds instanceof JSONArray)) {
throw new IllegalArgumentException("knowledgeIds 必须为数组");
}
JSONArray ids = (JSONArray) rawIds;
if (ids.isEmpty()) {
throw new IllegalArgumentException("knowledgeIds 不能为空");
}
java.util.LinkedHashSet<String> normalized =
new java.util.LinkedHashSet<>();
for (Object id : ids) {
String value = id == null ? null : String.valueOf(id).trim();
if (!com.easyagents.flow.core.util.StringUtil.hasText(value)) {
throw new IllegalArgumentException("knowledgeIds 不能包含空值");
}
if (!normalized.add(value)) {
throw new IllegalArgumentException("knowledgeIds 不能包含重复值");
}
}
knowledgeNode.setKnowledgeIds(new ArrayList<>(normalized));
} else {
knowledgeNode.setKnowledgeId(data.get("knowledgeId"));
}
knowledgeNode.setLimit(data.getString("limit"));
knowledgeNode.setKeyword(data.getString("keyword"));
knowledgeNode.setRetrievalMode(data.getString("retrievalMode"));

View File

@@ -0,0 +1,217 @@
/**
* Copyright (c) 2025-2026, Michael Yang 杨福海 (fuhai999@gmail.com).
* <p>
* Licensed under the GNU Lesser General Public License (LGPL) ,Version 3.0.
*/
package com.easyagents.flow.core.test;
import com.alibaba.fastjson.JSONArray;
import com.alibaba.fastjson.JSONObject;
import com.easyagents.flow.core.chain.Chain;
import com.easyagents.flow.core.chain.ChainDefinition;
import com.easyagents.flow.core.chain.ChainState;
import com.easyagents.flow.core.knowledge.Knowledge;
import com.easyagents.flow.core.knowledge.KnowledgeManager;
import com.easyagents.flow.core.knowledge.KnowledgeProvider;
import com.easyagents.flow.core.knowledge.KnowledgeSearchRequest;
import com.easyagents.flow.core.node.KnowledgeNode;
import com.easyagents.flow.core.parser.impl.KnowledgeNodeParser;
import org.junit.Assert;
import org.junit.Test;
import java.util.List;
import java.util.Map;
import java.util.concurrent.atomic.AtomicInteger;
/**
* 知识库节点多来源契约测试。
*/
public class KnowledgeNodeTest {
@Test
public void shouldParseLegacyKnowledgeId() {
JSONObject data = baseData();
data.put("knowledgeId", "101");
KnowledgeNode node = parse(data);
Assert.assertEquals(List.of("101"), node.getKnowledgeIds());
}
@Test
public void shouldPreferKnowledgeIdsAndKeepOrder() {
JSONObject data = baseData();
data.put("knowledgeId", "legacy");
JSONArray ids = new JSONArray();
ids.addAll(List.of("201", "202"));
data.put("knowledgeIds", ids);
KnowledgeNode node = parse(data);
Assert.assertEquals(List.of("201", "202"), node.getKnowledgeIds());
}
@Test
public void shouldNormalizeKnowledgeIdsBeforeStoringThem() {
JSONObject data = baseData();
JSONArray ids = new JSONArray();
ids.addAll(List.of(" 201 ", "202"));
data.put("knowledgeIds", ids);
KnowledgeNode node = parse(data);
Assert.assertEquals(List.of("201", "202"), node.getKnowledgeIds());
}
@Test
public void shouldRejectInvalidKnowledgeIds() {
JSONObject data = baseData();
data.put("knowledgeIds", "[201,202]");
assertParseFailure(data, "必须为数组");
JSONArray duplicateIds = new JSONArray();
duplicateIds.addAll(List.of("201", "201"));
data.put("knowledgeIds", duplicateIds);
assertParseFailure(data, "重复值");
}
@Test
public void defaultProviderShouldKeepSingleKnowledgeCompatibility() {
KnowledgeProvider provider = id ->
(keyword, limit, node, chain) -> List.of(Map.of(
"knowledgeId", id,
"content", keyword));
KnowledgeNode node = new KnowledgeNode();
node.setKnowledgeId("301");
Map<String, Object> output = provider.search(
new KnowledgeSearchRequest(
node.getKnowledgeIds(),
"问题",
3,
"HYBRID",
node,
null));
Assert.assertNotNull(output);
Assert.assertEquals(1, ((List<?>) output.get("documents")).size());
}
@Test
public void defaultProviderShouldDeclineMultiKnowledgeRequest() {
KnowledgeProvider provider = id -> null;
KnowledgeNode node = new KnowledgeNode();
node.setKnowledgeIds(List.of("401", "402"));
Assert.assertNull(provider.search(new KnowledgeSearchRequest(
node.getKnowledgeIds(),
"问题",
3,
"VECTOR",
node,
null)));
}
@Test
public void shouldResolveVariableLimitAndDefaultBlankValueAtRuntime() {
Assert.assertEquals(7, executeAndCaptureLimit("{{start.limit}}", "7"));
Assert.assertEquals(10, executeAndCaptureLimit("{{start.limit}}", " "));
Assert.assertEquals(10, executeAndCaptureLimit("{{start.limit ?? }}", null));
}
@Test
public void shouldRejectInvalidResolvedVariableLimitAtRuntime() {
assertRuntimeLimitFailure("abc");
assertRuntimeLimitFailure("0");
assertRuntimeLimitFailure("-2");
}
private static KnowledgeNode parse(JSONObject data) {
return new KnowledgeNodeParser().doParse(
new JSONObject(), data, new JSONObject());
}
private static JSONObject baseData() {
JSONObject data = new JSONObject();
data.put("keyword", "问题");
data.put("limit", "5");
data.put("retrievalMode", "VECTOR");
return data;
}
private static void assertParseFailure(
JSONObject data, String expectedMessage) {
try {
parse(data);
Assert.fail("invalid knowledgeIds must be rejected");
} catch (IllegalArgumentException expected) {
Assert.assertTrue(expected.getMessage().contains(expectedMessage));
}
}
private static int executeAndCaptureLimit(
String limitTemplate,
String runtimeValue) {
AtomicInteger capturedLimit = new AtomicInteger(-1);
KnowledgeProvider provider = new KnowledgeProvider() {
@Override
public Knowledge getKnowledge(Object id) {
return null;
}
@Override
public Map<String, Object> search(KnowledgeSearchRequest request) {
capturedLimit.set(request.getLimit());
return Map.of("documents", List.of());
}
};
KnowledgeManager.getInstance().registerProvider(provider);
try {
KnowledgeNode node = runtimeNode(limitTemplate);
ChainState state = new ChainState();
if (runtimeValue != null) {
state.getMemory().put("start.limit", runtimeValue);
}
node.execute(new FixedStateChain(state));
return capturedLimit.get();
} finally {
KnowledgeManager.getInstance().removeProvider(provider);
}
}
private static void assertRuntimeLimitFailure(String runtimeValue) {
KnowledgeNode node = runtimeNode("{{start.limit}}");
ChainState state = new ChainState();
state.getMemory().put("start.limit", runtimeValue);
IllegalArgumentException exception = Assert.assertThrows(
IllegalArgumentException.class,
() -> node.execute(new FixedStateChain(state)));
Assert.assertTrue(exception.getMessage().contains("必须为正整数"));
}
private static KnowledgeNode runtimeNode(String limitTemplate) {
KnowledgeNode node = new KnowledgeNode();
node.setKnowledgeIds(List.of("501", "502"));
node.setKeyword("问题");
node.setLimit(limitTemplate);
node.setRetrievalMode("VECTOR");
return node;
}
private static final class FixedStateChain extends Chain {
private final ChainState state;
private FixedStateChain(ChainState state) {
super(new ChainDefinition(), "knowledge-node-limit-test");
this.state = state;
}
@Override
public ChainState getExecutionState() {
return state;
}
}
}