feat: 支持工作流多知识库向量检索
This commit is contained in:
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -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 + '\'' +
|
||||
|
||||
@@ -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"));
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user