Compare commits
12 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 45c708a212 | |||
| e959a772b5 | |||
| f311822a3d | |||
| 2cdbc57174 | |||
| 4bc68ec7f7 | |||
| 93e0ff2204 | |||
| ea67d0519f | |||
| 2d26494775 | |||
| 651b292d83 | |||
| 67b69669c4 | |||
| 6192df4bd1 | |||
| bb37d9d708 |
@@ -1,14 +1,17 @@
|
||||
package com.easyagents.agent.runtime;
|
||||
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeRetriever;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeRegistration;
|
||||
import com.easyagents.agent.runtime.media.AgentMediaResolver;
|
||||
import com.easyagents.agent.runtime.memory.AgentMemorySnapshot;
|
||||
import com.easyagents.agent.runtime.persistence.conversation.AgentConversationRecorder;
|
||||
import com.easyagents.agent.runtime.persistence.conversation.noop.NoopAgentConversationRecorder;
|
||||
import com.easyagents.agent.runtime.persistence.session.AgentSessionStore;
|
||||
import com.easyagents.agent.runtime.persistence.session.noop.NoopAgentSessionStore;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolInvoker;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/**
|
||||
@@ -44,7 +47,12 @@ public class AgentInitRequest {
|
||||
/**
|
||||
* 知识库集合,实现AgentKnowledgeRetriever接口以进行知识检索动作。
|
||||
*/
|
||||
private Map<String, AgentKnowledgeRetriever> knowledgeRetrievers = new LinkedHashMap<>();
|
||||
private List<AgentKnowledgeRegistration> knowledgeRegistrations = new ArrayList<>();
|
||||
|
||||
/**
|
||||
* 首次构建 Agent 时装载的对话历史快照。
|
||||
*/
|
||||
private AgentMemorySnapshot memorySnapshot = new AgentMemorySnapshot();
|
||||
|
||||
/**
|
||||
* 对话事件记录器,用于记录运行时事件流。
|
||||
@@ -156,17 +164,37 @@ public class AgentInitRequest {
|
||||
*
|
||||
* @return 知识库检索器
|
||||
*/
|
||||
public Map<String, AgentKnowledgeRetriever> getKnowledgeRetrievers() {
|
||||
return knowledgeRetrievers;
|
||||
public List<AgentKnowledgeRegistration> getKnowledgeRegistrations() {
|
||||
return knowledgeRegistrations;
|
||||
}
|
||||
|
||||
/**
|
||||
* 设置知识库检索器。
|
||||
*
|
||||
* @param knowledgeRetrievers 知识库检索器
|
||||
* @param knowledgeRegistrations 知识库运行时绑定
|
||||
*/
|
||||
public void setKnowledgeRetrievers(Map<String, AgentKnowledgeRetriever> knowledgeRetrievers) {
|
||||
this.knowledgeRetrievers = knowledgeRetrievers == null ? new LinkedHashMap<>() : knowledgeRetrievers;
|
||||
public void setKnowledgeRegistrations(List<AgentKnowledgeRegistration> knowledgeRegistrations) {
|
||||
this.knowledgeRegistrations = knowledgeRegistrations == null
|
||||
? new ArrayList<>()
|
||||
: new ArrayList<>(knowledgeRegistrations);
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取首次构建 Agent 时的对话历史快照。
|
||||
*
|
||||
* @return 对话历史快照
|
||||
*/
|
||||
public AgentMemorySnapshot getMemorySnapshot() {
|
||||
return memorySnapshot;
|
||||
}
|
||||
|
||||
/**
|
||||
* 设置首次构建 Agent 时的对话历史快照。
|
||||
*
|
||||
* @param memorySnapshot 对话历史快照
|
||||
*/
|
||||
public void setMemorySnapshot(AgentMemorySnapshot memorySnapshot) {
|
||||
this.memorySnapshot = memorySnapshot == null ? new AgentMemorySnapshot() : memorySnapshot;
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package com.easyagents.agent.runtime;
|
||||
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeRetriever;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeRegistration;
|
||||
import com.easyagents.agent.runtime.memory.AgentMemorySnapshot;
|
||||
import com.easyagents.agent.runtime.message.AgentMessage;
|
||||
import com.easyagents.agent.runtime.persistence.conversation.AgentConversationRecorder;
|
||||
@@ -9,7 +9,9 @@ import com.easyagents.agent.runtime.persistence.session.AgentSessionStore;
|
||||
import com.easyagents.agent.runtime.persistence.session.noop.NoopAgentSessionStore;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolInvoker;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/**
|
||||
@@ -60,7 +62,7 @@ public class AgentRuntimeExecutionContext {
|
||||
/**
|
||||
* 按知识库ID索引的检索器。
|
||||
*/
|
||||
private Map<String, AgentKnowledgeRetriever> knowledgeRetrievers = new LinkedHashMap<>();
|
||||
private List<AgentKnowledgeRegistration> knowledgeRegistrations = new ArrayList<>();
|
||||
|
||||
/**
|
||||
* 会话状态存储。
|
||||
@@ -231,17 +233,19 @@ public class AgentRuntimeExecutionContext {
|
||||
*
|
||||
* @return 知识库检索器
|
||||
*/
|
||||
public Map<String, AgentKnowledgeRetriever> getKnowledgeRetrievers() {
|
||||
return knowledgeRetrievers;
|
||||
public List<AgentKnowledgeRegistration> getKnowledgeRegistrations() {
|
||||
return knowledgeRegistrations;
|
||||
}
|
||||
|
||||
/**
|
||||
* 设置知识库检索器。
|
||||
*
|
||||
* @param knowledgeRetrievers 知识库检索器
|
||||
* @param knowledgeRegistrations 知识库运行时绑定
|
||||
*/
|
||||
public void setKnowledgeRetrievers(Map<String, AgentKnowledgeRetriever> knowledgeRetrievers) {
|
||||
this.knowledgeRetrievers = knowledgeRetrievers == null ? new LinkedHashMap<>() : knowledgeRetrievers;
|
||||
public void setKnowledgeRegistrations(List<AgentKnowledgeRegistration> knowledgeRegistrations) {
|
||||
this.knowledgeRegistrations = knowledgeRegistrations == null
|
||||
? new ArrayList<>()
|
||||
: new ArrayList<>(knowledgeRegistrations);
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -2,309 +2,452 @@ package com.easyagents.agent.runtime.agentscope;
|
||||
|
||||
import com.easyagents.agent.runtime.AgentRuntimeException;
|
||||
import com.easyagents.agent.runtime.AgentRuntimeExecutionContext;
|
||||
import com.easyagents.agent.runtime.event.*;
|
||||
import com.easyagents.agent.runtime.knowledge.*;
|
||||
import io.agentscope.core.message.TextBlock;
|
||||
import io.agentscope.core.rag.Knowledge;
|
||||
import io.agentscope.core.rag.model.Document;
|
||||
import io.agentscope.core.rag.model.DocumentMetadata;
|
||||
import io.agentscope.core.rag.model.RetrieveConfig;
|
||||
import reactor.core.publisher.Mono;
|
||||
import reactor.core.publisher.Sinks;
|
||||
import com.easyagents.agent.runtime.event.AgentRuntimeEvent;
|
||||
import com.easyagents.agent.runtime.event.AgentRuntimeEventType;
|
||||
import com.easyagents.agent.runtime.event.AgentRuntimeTurnContextHolder;
|
||||
import com.easyagents.agent.runtime.hitl.AgentToolApprovalCoordinator;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeDocument;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeRegistration;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeRetrievalRequest;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeRetrievalResult;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeSpec;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeToolNames;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolCategory;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolContext;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolResult;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolSpec;
|
||||
import io.agentscope.core.tool.Toolkit;
|
||||
|
||||
import java.util.*;
|
||||
import java.util.ArrayList;
|
||||
import java.util.Comparator;
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.LinkedHashSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Objects;
|
||||
import java.util.Set;
|
||||
|
||||
/**
|
||||
* 将运行时知识库检索器适配为一个聚合 AgentScope Knowledge。
|
||||
* 将中立知识库绑定适配为 AgentScope 一库一工具。
|
||||
*/
|
||||
public class AgentScopeKnowledgeAdapter {
|
||||
|
||||
/**
|
||||
* 创建聚合 Knowledge。
|
||||
* 根据 Agent 知识库声明创建模型可见工具定义。
|
||||
*
|
||||
* @param request 运行请求
|
||||
* @return 聚合 Knowledge;未配置知识库时返回 null
|
||||
* @param context 运行时上下文
|
||||
* @return 知识库工具定义
|
||||
* @throws AgentRuntimeException 声明、运行名或 Retriever 绑定不合法时抛出
|
||||
*/
|
||||
public Knowledge createAggregateKnowledge(AgentRuntimeExecutionContext request) {
|
||||
return createAggregateKnowledge(request, (Sinks.Many<AgentRuntimeEvent>) null);
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建带事件 sink 的聚合 Knowledge。
|
||||
*
|
||||
* @param request 运行请求
|
||||
* @param eventSink 事件 sink
|
||||
* @return 聚合 Knowledge;未配置知识库时返回 null
|
||||
*/
|
||||
public Knowledge createAggregateKnowledge(AgentRuntimeExecutionContext request, Sinks.Many<AgentRuntimeEvent> eventSink) {
|
||||
return createAggregateKnowledge(request, fixedHolder(request, eventSink));
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建可读取当前运行轮次事件出口的聚合 Knowledge。
|
||||
*
|
||||
* @param request 运行时级上下文
|
||||
* @param turnContextHolder 当前运行轮次上下文持有器
|
||||
* @return 聚合 Knowledge;未配置知识库时返回 null
|
||||
*/
|
||||
public Knowledge createAggregateKnowledge(AgentRuntimeExecutionContext request,
|
||||
AgentRuntimeTurnContextHolder turnContextHolder) {
|
||||
if (request.getAgentDefinition().getKnowledgeSpecs().isEmpty()) {
|
||||
return null;
|
||||
public List<AgentToolSpec> createToolSpecs(AgentRuntimeExecutionContext context) {
|
||||
if (context == null || context.getAgentDefinition() == null) {
|
||||
throw new AgentRuntimeException("Agent runtime context and definition are required for knowledge tools.");
|
||||
}
|
||||
return new AggregateKnowledge(request, turnContextHolder);
|
||||
}
|
||||
|
||||
private AgentRuntimeTurnContextHolder fixedHolder(AgentRuntimeExecutionContext request,
|
||||
Sinks.Many<AgentRuntimeEvent> eventSink) {
|
||||
AgentRuntimeTurnContextHolder holder = new AgentRuntimeTurnContextHolder();
|
||||
AgentRuntimeEventBridge bridge = new AgentRuntimeEventBridge(request, holder);
|
||||
holder.set(new AgentRuntimeTurnContext(null, eventSink, bridge));
|
||||
return holder;
|
||||
List<AgentKnowledgeSpec> knowledgeSpecs = context.getAgentDefinition().getKnowledgeSpecs();
|
||||
List<AgentKnowledgeRegistration> registrations = context.getKnowledgeRegistrations();
|
||||
if (knowledgeSpecs == null || knowledgeSpecs.isEmpty()) {
|
||||
if (registrations != null && !registrations.isEmpty()) {
|
||||
throw new AgentRuntimeException("Knowledge registrations require matching knowledge specs.");
|
||||
}
|
||||
return List.of();
|
||||
}
|
||||
Map<String, AgentKnowledgeRegistration> registrationIndex = registrationIndex(registrations);
|
||||
if (registrationIndex.size() != knowledgeSpecs.size()) {
|
||||
throw new AgentRuntimeException("Knowledge specs and registrations must match one-to-one.");
|
||||
}
|
||||
List<AgentToolSpec> toolSpecs = new ArrayList<>(knowledgeSpecs.size());
|
||||
Set<String> toolNames = new LinkedHashSet<>();
|
||||
for (AgentKnowledgeSpec knowledgeSpec : knowledgeSpecs) {
|
||||
validateKnowledgeSpec(knowledgeSpec);
|
||||
AgentKnowledgeRegistration registration = registrationIndex.get(knowledgeSpec.getKnowledgeId());
|
||||
if (registration == null) {
|
||||
throw new AgentRuntimeException(
|
||||
"Knowledge retriever is required: " + knowledgeSpec.getKnowledgeId());
|
||||
}
|
||||
String toolName = AgentKnowledgeToolNames.build(knowledgeSpec.getRuntimeName());
|
||||
if (!toolNames.add(toolName)) {
|
||||
throw new AgentRuntimeException("Duplicate knowledge tool name: " + toolName);
|
||||
}
|
||||
toolSpecs.add(toolSpec(knowledgeSpec, toolName));
|
||||
}
|
||||
return toolSpecs;
|
||||
}
|
||||
|
||||
/**
|
||||
* 将运行时文档转换为 AgentScope 文档。
|
||||
* 将知识库工具注册到现有 AgentScope Toolkit。
|
||||
*
|
||||
* @param documents 运行时文档
|
||||
* @return AgentScope 文档
|
||||
* @param context 运行时上下文
|
||||
* @param toolSpecs 知识库工具定义
|
||||
* @param toolkit AgentScope Toolkit
|
||||
* @param toolAdapter 中立工具适配器
|
||||
* @param approvalCoordinator 工具审批协调器
|
||||
* @param turnContextHolder 当前运行轮次上下文持有器
|
||||
*/
|
||||
public List<Document> toDocuments(List<AgentKnowledgeDocument> documents) {
|
||||
List<Document> converted = new ArrayList<>();
|
||||
public void registerTools(AgentRuntimeExecutionContext context,
|
||||
List<AgentToolSpec> toolSpecs,
|
||||
Toolkit toolkit,
|
||||
AgentScopeToolAdapter toolAdapter,
|
||||
AgentToolApprovalCoordinator approvalCoordinator,
|
||||
AgentRuntimeTurnContextHolder turnContextHolder) {
|
||||
Objects.requireNonNull(toolkit, "toolkit");
|
||||
Objects.requireNonNull(toolAdapter, "toolAdapter");
|
||||
if (toolSpecs == null || toolSpecs.isEmpty()) {
|
||||
return;
|
||||
}
|
||||
Map<String, AgentKnowledgeRegistration> registrationIndex =
|
||||
registrationIndex(context.getKnowledgeRegistrations());
|
||||
for (AgentToolSpec toolSpec : toolSpecs) {
|
||||
String knowledgeId = stringValue(toolSpec.getMetadata().get("knowledgeId"));
|
||||
AgentKnowledgeRegistration registration = registrationIndex.get(knowledgeId);
|
||||
if (registration == null) {
|
||||
throw new AgentRuntimeException("Knowledge registration is required: " + knowledgeId);
|
||||
}
|
||||
toolkit.registerAgentTool(toolAdapter.adapt(
|
||||
toolSpec,
|
||||
(input, toolContext) -> retrieve(registration, input, toolContext),
|
||||
context,
|
||||
approvalCoordinator,
|
||||
turnContextHolder,
|
||||
null,
|
||||
null,
|
||||
false,
|
||||
false,
|
||||
false));
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 按知识库 ID 建立运行时绑定索引并拒绝重复绑定。
|
||||
*
|
||||
* @param registrations 知识库运行时绑定
|
||||
* @return 以知识库 ID 为键的绑定索引
|
||||
*/
|
||||
private Map<String, AgentKnowledgeRegistration> registrationIndex(
|
||||
List<AgentKnowledgeRegistration> registrations) {
|
||||
Map<String, AgentKnowledgeRegistration> index = new LinkedHashMap<>();
|
||||
if (registrations == null) {
|
||||
return index;
|
||||
}
|
||||
for (AgentKnowledgeRegistration registration : registrations) {
|
||||
if (registration == null || registration.getKnowledgeSpec() == null) {
|
||||
throw new AgentRuntimeException("Knowledge registration and spec are required.");
|
||||
}
|
||||
String knowledgeId = registration.getKnowledgeSpec().getKnowledgeId();
|
||||
if (knowledgeId == null || knowledgeId.isBlank()) {
|
||||
throw new AgentRuntimeException("Knowledge id is required.");
|
||||
}
|
||||
if (index.putIfAbsent(knowledgeId, registration) != null) {
|
||||
throw new AgentRuntimeException("Duplicate knowledge registration: " + knowledgeId);
|
||||
}
|
||||
}
|
||||
return index;
|
||||
}
|
||||
|
||||
/**
|
||||
* 校验知识库声明的标识和英文运行名。
|
||||
*
|
||||
* @param knowledgeSpec 知识库声明
|
||||
* @throws AgentRuntimeException 声明不合法时抛出
|
||||
*/
|
||||
private void validateKnowledgeSpec(AgentKnowledgeSpec knowledgeSpec) {
|
||||
if (knowledgeSpec == null) {
|
||||
throw new AgentRuntimeException("Knowledge spec is required.");
|
||||
}
|
||||
if (knowledgeSpec.getKnowledgeId() == null || knowledgeSpec.getKnowledgeId().isBlank()) {
|
||||
throw new AgentRuntimeException("Knowledge id is required.");
|
||||
}
|
||||
AgentKnowledgeToolNames.build(knowledgeSpec.getRuntimeName());
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建单个知识库对应的模型工具声明。
|
||||
*
|
||||
* @param knowledgeSpec 知识库声明
|
||||
* @param toolName 已规范化的工具名
|
||||
* @return 模型可见工具声明
|
||||
*/
|
||||
private AgentToolSpec toolSpec(AgentKnowledgeSpec knowledgeSpec, String toolName) {
|
||||
AgentToolSpec toolSpec = new AgentToolSpec();
|
||||
toolSpec.setName(toolName);
|
||||
toolSpec.setDescription(toolDescription(knowledgeSpec));
|
||||
toolSpec.setCategory(AgentToolCategory.KNOWLEDGE);
|
||||
toolSpec.setParametersSchema(Map.of(
|
||||
"type", "object",
|
||||
"properties", Map.of(
|
||||
"query", Map.of(
|
||||
"type", "string",
|
||||
"description", "A standalone search query containing all context needed to retrieve relevant knowledge."
|
||||
)
|
||||
),
|
||||
"required", List.of("query"),
|
||||
"additionalProperties", false
|
||||
));
|
||||
toolSpec.getMetadata().putAll(knowledgeSpec.getMetadata());
|
||||
toolSpec.getMetadata().put("knowledgeId", knowledgeSpec.getKnowledgeId());
|
||||
toolSpec.getMetadata().put("knowledgeName", knowledgeSpec.getName());
|
||||
toolSpec.getMetadata().put("knowledgeRuntimeName", knowledgeSpec.getRuntimeName());
|
||||
toolSpec.getMetadata().put("toolDisplayName", knowledgeSpec.getName());
|
||||
return toolSpec;
|
||||
}
|
||||
|
||||
/**
|
||||
* 生成包含知识库名称、范围和查询要求的工具描述。
|
||||
*
|
||||
* @param knowledgeSpec 知识库声明
|
||||
* @return 模型可见工具描述
|
||||
*/
|
||||
private String toolDescription(AgentKnowledgeSpec knowledgeSpec) {
|
||||
String knowledgeName = hasText(knowledgeSpec.getName())
|
||||
? knowledgeSpec.getName().trim()
|
||||
: knowledgeSpec.getRuntimeName().trim();
|
||||
StringBuilder description = new StringBuilder()
|
||||
.append("Search the knowledge base \"")
|
||||
.append(knowledgeName)
|
||||
.append("\" for relevant and grounded information.");
|
||||
if (hasText(knowledgeSpec.getDescription())) {
|
||||
description.append(" Knowledge scope: ")
|
||||
.append(knowledgeSpec.getDescription().trim())
|
||||
.append('.');
|
||||
}
|
||||
description.append(" Build a standalone query from the current question and necessary conversation context.");
|
||||
return description.toString();
|
||||
}
|
||||
|
||||
/**
|
||||
* 执行单库检索并生成模型证据及旁路检索事件。
|
||||
*
|
||||
* @param registration 知识库运行时绑定
|
||||
* @param input Function Call 输入
|
||||
* @param toolContext 工具执行上下文
|
||||
* @return 模型可见工具结果
|
||||
*/
|
||||
private AgentToolResult retrieve(AgentKnowledgeRegistration registration,
|
||||
Map<String, Object> input,
|
||||
AgentToolContext toolContext) {
|
||||
AgentKnowledgeSpec knowledgeSpec = registration.getKnowledgeSpec();
|
||||
String query = input == null ? null : stringValue(input.get("query"));
|
||||
if (query == null || query.isBlank()) {
|
||||
throw new AgentRuntimeException("Knowledge tool query is required: "
|
||||
+ AgentKnowledgeToolNames.build(knowledgeSpec.getRuntimeName()));
|
||||
}
|
||||
String normalizedQuery = query.trim();
|
||||
AgentKnowledgeRetrievalRequest retrievalRequest = new AgentKnowledgeRetrievalRequest();
|
||||
retrievalRequest.setQuery(normalizedQuery);
|
||||
retrievalRequest.setLimit(knowledgeSpec.getLimit());
|
||||
retrievalRequest.setScoreThreshold(knowledgeSpec.getScoreThreshold());
|
||||
retrievalRequest.setKnowledgeSpec(knowledgeSpec);
|
||||
retrievalRequest.setRuntimeContext(toolContext.getRuntimeContext());
|
||||
retrievalRequest.getMetadata().put("requestId", toolContext.getRequestId());
|
||||
retrievalRequest.getMetadata().put("traceId", toolContext.getTraceId());
|
||||
retrievalRequest.getMetadata().put("sessionId", toolContext.getSessionId());
|
||||
retrievalRequest.getMetadata().put("toolCallId", toolContext.getToolCallId());
|
||||
AgentKnowledgeRetrievalResult retrievalResult = registration.getRetriever().retrieve(retrievalRequest);
|
||||
if (retrievalResult == null) {
|
||||
throw new AgentRuntimeException("Knowledge retriever returned null result: "
|
||||
+ knowledgeSpec.getKnowledgeId());
|
||||
}
|
||||
List<AgentKnowledgeDocument> documents = finalDocuments(knowledgeSpec, retrievalResult.getDocuments());
|
||||
toolContext.emitEvent(retrievalEvent(toolContext, knowledgeSpec, retrievalRequest, documents));
|
||||
|
||||
List<Map<String, Object>> summaries = documentSummaries(documents);
|
||||
AgentToolResult toolResult = AgentToolResult.success(modelContent(knowledgeSpec, normalizedQuery, documents));
|
||||
toolResult.setDisplayContent(summaries);
|
||||
toolResult.getMetadata().putAll(retrievalResult.getMetadata());
|
||||
toolResult.getMetadata().put("knowledgeId", knowledgeSpec.getKnowledgeId());
|
||||
toolResult.getMetadata().put("knowledgeName", knowledgeSpec.getName());
|
||||
toolResult.getMetadata().put("query", normalizedQuery);
|
||||
toolResult.getMetadata().put("documentCount", documents.size());
|
||||
toolResult.getMetadata().put("documents", summaries);
|
||||
return toolResult;
|
||||
}
|
||||
|
||||
/**
|
||||
* 对检索结果执行统一阈值过滤、降序排序和绑定数量截断。
|
||||
*
|
||||
* @param knowledgeSpec 知识库声明
|
||||
* @param documents 原始检索文档
|
||||
* @return 最终进入模型上下文的文档
|
||||
*/
|
||||
private List<AgentKnowledgeDocument> finalDocuments(AgentKnowledgeSpec knowledgeSpec,
|
||||
List<AgentKnowledgeDocument> documents) {
|
||||
List<AgentKnowledgeDocument> finalDocuments = new ArrayList<>();
|
||||
if (documents == null) {
|
||||
return converted;
|
||||
return finalDocuments;
|
||||
}
|
||||
for (AgentKnowledgeDocument document : documents) {
|
||||
converted.add(toDocument(document));
|
||||
if (document == null || !passesThreshold(document, knowledgeSpec.getScoreThreshold())) {
|
||||
continue;
|
||||
}
|
||||
preserveKnowledgeMetadata(knowledgeSpec, document);
|
||||
finalDocuments.add(document);
|
||||
}
|
||||
return converted;
|
||||
}
|
||||
|
||||
private Document toDocument(AgentKnowledgeDocument document) {
|
||||
Map<String, Object> payload = new LinkedHashMap<>();
|
||||
payload.put("documentId", document.getDocumentId());
|
||||
payload.put("documentName", document.getDocumentName());
|
||||
payload.put("chunkId", document.getChunkId());
|
||||
payload.put("sourceUri", document.getSourceUri());
|
||||
payload.put("knowledgeMetadata", document.getKnowledgeMetadata());
|
||||
payload.put("documentMetadata", document.getMetadata());
|
||||
payload.putAll(document.getMetadata());
|
||||
DocumentMetadata metadata = DocumentMetadata.builder()
|
||||
.content(TextBlock.builder().text(safeContent(document)).build())
|
||||
.docId(safeDocumentId(document))
|
||||
.chunkId(safeChunkId(document))
|
||||
.payload(payload)
|
||||
.build();
|
||||
Document converted = new Document(metadata);
|
||||
converted.setScore(document.getScore());
|
||||
return converted;
|
||||
finalDocuments.sort(Comparator.comparing(
|
||||
AgentKnowledgeDocument::getScore,
|
||||
Comparator.nullsLast(Comparator.reverseOrder())));
|
||||
int limit = Math.max(knowledgeSpec.getLimit(), 1);
|
||||
if (finalDocuments.size() > limit) {
|
||||
return new ArrayList<>(finalDocuments.subList(0, limit));
|
||||
}
|
||||
return finalDocuments;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取 AgentScope 要求的非空文档 ID。
|
||||
* 判断文档最终分数是否达到绑定阈值。
|
||||
*
|
||||
* @param document 知识文档
|
||||
* @return 非空文档 ID
|
||||
* @param document 检索文档
|
||||
* @param scoreThreshold 分数阈值
|
||||
* @return 达到阈值时为 true
|
||||
*/
|
||||
private String safeDocumentId(AgentKnowledgeDocument document) {
|
||||
if (document.getDocumentId() != null && !document.getDocumentId().isBlank()) {
|
||||
return document.getDocumentId();
|
||||
private boolean passesThreshold(AgentKnowledgeDocument document, double scoreThreshold) {
|
||||
if (scoreThreshold <= 0D) {
|
||||
return true;
|
||||
}
|
||||
if (document.getChunkId() != null && !document.getChunkId().isBlank()) {
|
||||
return document.getChunkId();
|
||||
}
|
||||
return "knowledge-document";
|
||||
return document.getScore() != null && document.getScore() >= scoreThreshold;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取 AgentScope 要求的非空分片 ID。
|
||||
* 将知识库归属信息合并到文档元数据中。
|
||||
*
|
||||
* @param document 知识文档
|
||||
* @return 非空分片 ID
|
||||
* @param knowledgeSpec 知识库声明
|
||||
* @param document 检索文档
|
||||
*/
|
||||
private String safeChunkId(AgentKnowledgeDocument document) {
|
||||
if (document.getChunkId() != null && !document.getChunkId().isBlank()) {
|
||||
return document.getChunkId();
|
||||
}
|
||||
if (document.getDocumentId() != null && !document.getDocumentId().isBlank()) {
|
||||
return document.getDocumentId();
|
||||
}
|
||||
return "0";
|
||||
private void preserveKnowledgeMetadata(AgentKnowledgeSpec knowledgeSpec,
|
||||
AgentKnowledgeDocument document) {
|
||||
Map<String, Object> knowledgeMetadata = new LinkedHashMap<>(knowledgeSpec.getMetadata());
|
||||
knowledgeMetadata.put("knowledgeId", knowledgeSpec.getKnowledgeId());
|
||||
knowledgeMetadata.put("knowledgeName", knowledgeSpec.getName());
|
||||
knowledgeMetadata.put("knowledgeRuntimeName", knowledgeSpec.getRuntimeName());
|
||||
knowledgeMetadata.putAll(document.getKnowledgeMetadata());
|
||||
document.setKnowledgeMetadata(knowledgeMetadata);
|
||||
document.getMetadata().putIfAbsent("knowledgeId", knowledgeSpec.getKnowledgeId());
|
||||
document.getMetadata().putIfAbsent("knowledgeName", knowledgeSpec.getName());
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取 AgentScope 要求的非空文档内容。
|
||||
* 创建与模型最终证据一致的知识库检索事件。
|
||||
*
|
||||
* @param document 知识文档
|
||||
* @return 文档内容
|
||||
* @param toolContext 工具执行上下文
|
||||
* @param knowledgeSpec 知识库声明
|
||||
* @param retrievalRequest 检索请求
|
||||
* @param documents 最终文档
|
||||
* @return 检索旁路事件
|
||||
*/
|
||||
private String safeContent(AgentKnowledgeDocument document) {
|
||||
return document.getContent() == null ? "" : document.getContent();
|
||||
private AgentRuntimeEvent retrievalEvent(AgentToolContext toolContext,
|
||||
AgentKnowledgeSpec knowledgeSpec,
|
||||
AgentKnowledgeRetrievalRequest retrievalRequest,
|
||||
List<AgentKnowledgeDocument> documents) {
|
||||
AgentRuntimeEvent event = AgentRuntimeEvent.of(AgentRuntimeEventType.KNOWLEDGE_RETRIEVAL);
|
||||
event.setToolCallId(toolContext.getToolCallId());
|
||||
event.getPayload().put("query", retrievalRequest.getQuery());
|
||||
event.getPayload().put("knowledgeId", knowledgeSpec.getKnowledgeId());
|
||||
event.getPayload().put("knowledgeName", knowledgeSpec.getName());
|
||||
event.getPayload().put("knowledgeType", knowledgeSpec.getMetadata().get("knowledgeType"));
|
||||
event.getPayload().put("faqCollection", knowledgeSpec.getMetadata().get("faqCollection"));
|
||||
event.getPayload().put("limit", retrievalRequest.getLimit());
|
||||
event.getPayload().put("scoreThreshold", retrievalRequest.getScoreThreshold());
|
||||
event.getPayload().put("documentCount", documents.size());
|
||||
event.getPayload().put("documents", documentSummaries(documents));
|
||||
event.getMetadata().put("toolName", AgentKnowledgeToolNames.build(knowledgeSpec.getRuntimeName()));
|
||||
event.getMetadata().put("knowledgeRuntimeName", knowledgeSpec.getRuntimeName());
|
||||
return event;
|
||||
}
|
||||
|
||||
/**
|
||||
* 将检索调用分发到多个知识源的聚合 Knowledge 实现。
|
||||
* 将最终文档转换为 UI 和完成消息引用可消费的稳定摘要。
|
||||
*
|
||||
* @param documents 最终文档
|
||||
* @return 文档摘要列表
|
||||
*/
|
||||
private class AggregateKnowledge implements Knowledge {
|
||||
|
||||
private final AgentRuntimeExecutionContext request;
|
||||
private final AgentRuntimeTurnContextHolder turnContextHolder;
|
||||
|
||||
private AggregateKnowledge(AgentRuntimeExecutionContext request, AgentRuntimeTurnContextHolder turnContextHolder) {
|
||||
this.request = request;
|
||||
this.turnContextHolder = turnContextHolder;
|
||||
}
|
||||
|
||||
/**
|
||||
* 忽略文档新增,因为知识库索引由 EasyFlow 负责。
|
||||
*
|
||||
* @param documents 文档列表
|
||||
* @return 完成信号
|
||||
*/
|
||||
@Override
|
||||
public Mono<Void> addDocuments(List<Document> documents) {
|
||||
return Mono.error(new UnsupportedOperationException(
|
||||
"Easy-Agents agent runtime knowledge does not support addDocuments. Use external knowledge service instead."));
|
||||
}
|
||||
|
||||
/**
|
||||
* 从已配置的知识源检索文档。
|
||||
*
|
||||
* @param query 查询
|
||||
* @param config 检索配置
|
||||
* @return 文档列表
|
||||
*/
|
||||
@Override
|
||||
public Mono<List<Document>> retrieve(String query, RetrieveConfig config) {
|
||||
return Mono.fromCallable(() -> retrieveAll(query, config));
|
||||
}
|
||||
|
||||
/**
|
||||
* 检索并合并所有已配置的知识源。
|
||||
*
|
||||
* @param query 查询
|
||||
* @param config 检索配置
|
||||
* @return 合并后的文档
|
||||
*/
|
||||
private List<Document> retrieveAll(String query, RetrieveConfig config) {
|
||||
List<AgentKnowledgeDocument> allDocuments = new ArrayList<>();
|
||||
int globalLimit = config == null || config.getLimit() <= 0 ? 5 : config.getLimit();
|
||||
double globalThreshold = config == null ? 0D : config.getScoreThreshold();
|
||||
for (AgentKnowledgeSpec spec : request.getAgentDefinition().getKnowledgeSpecs()) {
|
||||
AgentKnowledgeRetriever retriever = request.getKnowledgeRetrievers().get(spec.getKnowledgeId());
|
||||
if (retriever == null) {
|
||||
throw new AgentRuntimeException("Knowledge retriever is required: " + spec.getKnowledgeId());
|
||||
}
|
||||
AgentKnowledgeRetrievalRequest retrievalRequest = new AgentKnowledgeRetrievalRequest();
|
||||
retrievalRequest.setQuery(query);
|
||||
retrievalRequest.setLimit(spec.getLimit());
|
||||
retrievalRequest.setScoreThreshold(Math.max(spec.getScoreThreshold(), globalThreshold));
|
||||
retrievalRequest.setKnowledgeSpec(spec);
|
||||
AgentRuntimeExecutionContext currentRequest = currentRequest();
|
||||
retrievalRequest.setRuntimeContext(currentRequest.getRuntimeContext());
|
||||
retrievalRequest.getMetadata().put("traceId", currentRequest.getTraceId());
|
||||
retrievalRequest.getMetadata().put("sessionId", currentRequest.getSessionId());
|
||||
AgentKnowledgeRetrievalResult result = retriever.retrieve(retrievalRequest);
|
||||
if (result == null || result.getDocuments() == null) {
|
||||
emitKnowledgeRetrievalEvent(query, spec, retrievalRequest, new ArrayList<>());
|
||||
continue;
|
||||
}
|
||||
emitKnowledgeRetrievalEvent(query, spec, retrievalRequest, result.getDocuments());
|
||||
for (AgentKnowledgeDocument document : result.getDocuments()) {
|
||||
preserveKnowledgeMetadata(spec, document);
|
||||
allDocuments.add(document);
|
||||
}
|
||||
}
|
||||
allDocuments.sort(Comparator.comparing(
|
||||
AgentKnowledgeDocument::getScore,
|
||||
Comparator.nullsLast(Comparator.reverseOrder())
|
||||
));
|
||||
if (allDocuments.size() > globalLimit) {
|
||||
allDocuments = new ArrayList<>(allDocuments.subList(0, globalLimit));
|
||||
}
|
||||
return toDocuments(allDocuments);
|
||||
}
|
||||
|
||||
/**
|
||||
* 在单条文档上保留知识库级元数据。
|
||||
*
|
||||
* @param spec 知识库声明
|
||||
* @param document 文档
|
||||
*/
|
||||
private void preserveKnowledgeMetadata(AgentKnowledgeSpec spec, AgentKnowledgeDocument document) {
|
||||
Map<String, Object> knowledgeMetadata = new LinkedHashMap<>(spec.getMetadata());
|
||||
knowledgeMetadata.put("knowledgeId", spec.getKnowledgeId());
|
||||
knowledgeMetadata.put("knowledgeName", spec.getName());
|
||||
knowledgeMetadata.put("retrievalMode", spec.getRetrievalMode().name());
|
||||
knowledgeMetadata.putAll(document.getKnowledgeMetadata());
|
||||
document.setKnowledgeMetadata(knowledgeMetadata);
|
||||
document.getMetadata().putIfAbsent("knowledgeId", spec.getKnowledgeId());
|
||||
document.getMetadata().putIfAbsent("knowledgeName", spec.getName());
|
||||
}
|
||||
|
||||
/**
|
||||
* 发射知识库检索旁路事件,供聊天界面展示检索过程。
|
||||
*
|
||||
* <p>知识库检索本身属于 AgentScope RAG 主线路,返回的 Document 会继续进入
|
||||
* AgentScope 的上下文注入流程;这里发出的 {@code KNOWLEDGE_RETRIEVAL}
|
||||
* 只是旁路告知调用方,不会回写 memory,也不会参与模型消息序列。</p>
|
||||
*
|
||||
* @param query 查询
|
||||
* @param spec 知识库声明
|
||||
* @param retrievalRequest 检索请求
|
||||
* @param documents 检索文档
|
||||
*/
|
||||
private void emitKnowledgeRetrievalEvent(String query,
|
||||
AgentKnowledgeSpec spec,
|
||||
AgentKnowledgeRetrievalRequest retrievalRequest,
|
||||
List<AgentKnowledgeDocument> documents) {
|
||||
AgentRuntimeEvent event = currentEventBridge().event(AgentRuntimeEventType.KNOWLEDGE_RETRIEVAL);
|
||||
event.getPayload().put("query", query);
|
||||
event.getPayload().put("knowledgeId", spec.getKnowledgeId());
|
||||
event.getPayload().put("knowledgeName", spec.getName());
|
||||
event.getPayload().put("knowledgeType", spec.getMetadata().get("knowledgeType"));
|
||||
event.getPayload().put("faqCollection", spec.getMetadata().get("faqCollection"));
|
||||
event.getPayload().put("limit", retrievalRequest.getLimit());
|
||||
event.getPayload().put("scoreThreshold", retrievalRequest.getScoreThreshold());
|
||||
event.getPayload().put("documentCount", documents == null ? 0 : documents.size());
|
||||
event.getPayload().put("documents", documentSummaries(documents));
|
||||
currentEventBridge().emit(event);
|
||||
}
|
||||
|
||||
private AgentRuntimeExecutionContext currentRequest() {
|
||||
return turnContextHolder == null ? request : turnContextHolder.executionContext(request);
|
||||
}
|
||||
|
||||
private AgentRuntimeEventBridge currentEventBridge() {
|
||||
if (turnContextHolder != null && turnContextHolder.eventBridge().isPresent()) {
|
||||
return turnContextHolder.eventBridge().get();
|
||||
}
|
||||
return new AgentRuntimeEventBridge(request, turnContextHolder);
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建用于事件展示的命中片段,保留前端引注需要的原始 chunk 内容。
|
||||
*
|
||||
* @param documents 检索文档
|
||||
* @return 命中片段列表
|
||||
*/
|
||||
private List<Map<String, Object>> documentSummaries(List<AgentKnowledgeDocument> documents) {
|
||||
List<Map<String, Object>> summaries = new ArrayList<>();
|
||||
if (documents == null) {
|
||||
return summaries;
|
||||
}
|
||||
for (AgentKnowledgeDocument document : documents) {
|
||||
Map<String, Object> summary = new LinkedHashMap<>();
|
||||
summary.put("documentId", document.getDocumentId());
|
||||
summary.put("documentName", document.getDocumentName());
|
||||
summary.put("chunkId", document.getChunkId());
|
||||
summary.put("chunkContent", document.getContent());
|
||||
summary.put("score", document.getScore());
|
||||
summary.put("sourceUri", document.getSourceUri());
|
||||
summary.put("metadata", document.getMetadata());
|
||||
summaries.add(summary);
|
||||
}
|
||||
private List<Map<String, Object>> documentSummaries(List<AgentKnowledgeDocument> documents) {
|
||||
List<Map<String, Object>> summaries = new ArrayList<>();
|
||||
if (documents == null) {
|
||||
return summaries;
|
||||
}
|
||||
for (AgentKnowledgeDocument document : documents) {
|
||||
Map<String, Object> summary = new LinkedHashMap<>();
|
||||
summary.put("documentId", document.getDocumentId());
|
||||
summary.put("documentName", document.getDocumentName());
|
||||
summary.put("chunkId", document.getChunkId());
|
||||
summary.put("chunkContent", document.getContent());
|
||||
summary.put("score", document.getScore());
|
||||
summary.put("sourceUri", document.getSourceUri());
|
||||
summary.put("metadata", document.getMetadata());
|
||||
summaries.add(summary);
|
||||
}
|
||||
return summaries;
|
||||
}
|
||||
|
||||
/**
|
||||
* 格式化模型可见的结构化检索证据。
|
||||
*
|
||||
* @param knowledgeSpec 知识库声明
|
||||
* @param query 实际检索词
|
||||
* @param documents 最终文档
|
||||
* @return 模型上下文文本
|
||||
*/
|
||||
private String modelContent(AgentKnowledgeSpec knowledgeSpec,
|
||||
String query,
|
||||
List<AgentKnowledgeDocument> documents) {
|
||||
String knowledgeName = hasText(knowledgeSpec.getName())
|
||||
? knowledgeSpec.getName().trim()
|
||||
: knowledgeSpec.getRuntimeName().trim();
|
||||
if (documents.isEmpty()) {
|
||||
return "No relevant documents were found in knowledge base \""
|
||||
+ knowledgeName + "\" for query: " + query;
|
||||
}
|
||||
StringBuilder content = new StringBuilder()
|
||||
.append("Retrieved evidence from knowledge base \"")
|
||||
.append(knowledgeName)
|
||||
.append("\" for query: ")
|
||||
.append(query)
|
||||
.append("\n\n");
|
||||
for (int index = 0; index < documents.size(); index++) {
|
||||
AgentKnowledgeDocument document = documents.get(index);
|
||||
content.append("[Evidence ").append(index + 1).append("]\n");
|
||||
appendField(content, "Document", document.getDocumentName());
|
||||
appendField(content, "Document ID", document.getDocumentId());
|
||||
appendField(content, "Chunk ID", document.getChunkId());
|
||||
appendField(content, "Source", document.getSourceUri());
|
||||
if (document.getScore() != null) {
|
||||
content.append("Score: ").append(document.getScore()).append('\n');
|
||||
}
|
||||
content.append("Content:\n")
|
||||
.append(document.getContent() == null ? "" : document.getContent())
|
||||
.append("\n\n");
|
||||
}
|
||||
return content.toString().stripTrailing();
|
||||
}
|
||||
|
||||
/**
|
||||
* 追加非空证据字段。
|
||||
*
|
||||
* @param content 输出缓冲区
|
||||
* @param label 字段标签
|
||||
* @param value 字段值
|
||||
*/
|
||||
private void appendField(StringBuilder content, String label, String value) {
|
||||
if (hasText(value)) {
|
||||
content.append(label).append(": ").append(value.trim()).append('\n');
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 判断字符串是否包含非空白文本。
|
||||
*
|
||||
* @param value 待判断字符串
|
||||
* @return 包含文本时为 true
|
||||
*/
|
||||
private boolean hasText(String value) {
|
||||
return value != null && !value.isBlank();
|
||||
}
|
||||
|
||||
/**
|
||||
* 将可空值转换为字符串。
|
||||
*
|
||||
* @param value 原始值
|
||||
* @return 字符串;原始值为空时返回 null
|
||||
*/
|
||||
private String stringValue(Object value) {
|
||||
return value == null ? null : String.valueOf(value);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -13,7 +13,6 @@ import com.easyagents.agent.runtime.hitl.AgentPendingState;
|
||||
import com.easyagents.agent.runtime.hitl.AgentToolApprovalCoordinator;
|
||||
import com.easyagents.agent.runtime.hitl.AgentToolApprovalResolution;
|
||||
import com.easyagents.agent.runtime.hitl.AgentToolApprovalRejectedException;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeSpec;
|
||||
import com.easyagents.agent.runtime.knowledge.citation.AgentKnowledgeCitationMatcher;
|
||||
import com.easyagents.agent.runtime.knowledge.citation.HeuristicKnowledgeCitationMatcher;
|
||||
import com.easyagents.agent.runtime.message.*;
|
||||
@@ -22,6 +21,7 @@ import com.easyagents.agent.runtime.mcp.McpSkillRegistration;
|
||||
import com.easyagents.agent.runtime.mcp.McpSpecValidator;
|
||||
import com.easyagents.agent.runtime.mcp.McpToolkitAdapter;
|
||||
import com.easyagents.agent.runtime.persistence.session.noop.NoopAgentSessionStore;
|
||||
import com.easyagents.agent.runtime.prompt.SystemPromptComposer;
|
||||
import com.easyagents.agent.runtime.skill.AgentSkillBinding;
|
||||
import com.easyagents.agent.runtime.skill.AgentSkillRuntimeContext;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolInvoker;
|
||||
@@ -35,9 +35,6 @@ import io.agentscope.core.memory.Memory;
|
||||
import io.agentscope.core.memory.autocontext.AutoContextMemory;
|
||||
import io.agentscope.core.message.*;
|
||||
import io.agentscope.core.model.Model;
|
||||
import io.agentscope.core.rag.Knowledge;
|
||||
import io.agentscope.core.rag.RAGMode;
|
||||
import io.agentscope.core.rag.model.RetrieveConfig;
|
||||
import io.agentscope.core.session.Session;
|
||||
import io.agentscope.core.skill.SkillBox;
|
||||
import io.agentscope.core.state.SessionKey;
|
||||
@@ -58,18 +55,6 @@ import java.util.function.Supplier;
|
||||
*/
|
||||
public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
|
||||
private static final String ASYNC_TOOL_SYSTEM_PROMPT = """
|
||||
|
||||
Async tool protocol:
|
||||
- Async tools may expose submit, observe, result, cancel, and list sub-tools. Treat these sub-tools as one user-facing tool.
|
||||
- Do not ask the user to choose submit, observe, result, cancel, or list. These are internal execution phases.
|
||||
- For a normal user request to use an async tool, call its submit sub-tool first with the user-provided arguments by default.
|
||||
- After submit returns task_id, immediately call observe with that task_id to check progress.
|
||||
- If the task is completed and result is available, use the returned result to answer the user.
|
||||
- If the task is still running after observation, tell the user that the task is running and keep task_id/next_action for later tool calls.
|
||||
- Use result, list, or cancel directly only when the user explicitly asks to get a known task result, list tasks, or cancel a task.
|
||||
""";
|
||||
|
||||
private final AgentScopeModelFactory modelFactory;
|
||||
private final AgentScopeToolAdapter toolAdapter;
|
||||
private final AgentScopeKnowledgeAdapter knowledgeAdapter;
|
||||
@@ -360,7 +345,7 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
context.setRuntimeContext(runtimeContext.getRuntimeContext());
|
||||
context.setUserMessage(userMessage);
|
||||
context.setToolInvokers(runtimeContext.getToolInvokers());
|
||||
context.setKnowledgeRetrievers(runtimeContext.getKnowledgeRetrievers());
|
||||
context.setKnowledgeRegistrations(runtimeContext.getKnowledgeRegistrations());
|
||||
context.setSessionStore(runtimeContext.getSessionStore());
|
||||
context.setConversationRecorder(runtimeContext.getConversationRecorder());
|
||||
context.setMetadata(runtimeContext.getMetadata());
|
||||
@@ -381,7 +366,7 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
context.setAgentDefinition(runtimeContext.getAgentDefinition());
|
||||
context.setRuntimeContext(runtimeContext.getRuntimeContext());
|
||||
context.setToolInvokers(runtimeContext.getToolInvokers());
|
||||
context.setKnowledgeRetrievers(runtimeContext.getKnowledgeRetrievers());
|
||||
context.setKnowledgeRegistrations(runtimeContext.getKnowledgeRegistrations());
|
||||
context.setSessionStore(runtimeContext.getSessionStore());
|
||||
context.setConversationRecorder(runtimeContext.getConversationRecorder());
|
||||
Map<String, Object> metadata = new LinkedHashMap<>(runtimeContext.getMetadata());
|
||||
@@ -689,6 +674,7 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
saveSession();
|
||||
return Flux.just(cancelled(context));
|
||||
}
|
||||
saveSession();
|
||||
return Flux.just(failed(context, error));
|
||||
}
|
||||
|
||||
@@ -733,7 +719,8 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
* 将取消前已输出的助手内容补写入 AgentScope memory 并保存 session。
|
||||
*
|
||||
* <p>AgentScope 的正常完成路径会自行把最终助手消息写入 memory。取消订阅时不会触发
|
||||
* 完成路径,因此这里仅在已有非空助手内容时补写一次,确保下一轮对话能拿到中断前上下文。</p>
|
||||
* 完成路径,因此这里仅在已有非空助手内容时补写一次,并始终保存已经进入 memory 的
|
||||
* 用户消息,确保下一轮对话能拿到中断前上下文。</p>
|
||||
*
|
||||
* @param finalText 当前已累计的助手文本
|
||||
* @param finalMessage 当前已捕获的结构化助手消息
|
||||
@@ -741,10 +728,9 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
private void persistPartialAssistantOnCancel(StringBuilder finalText,
|
||||
AtomicReference<AgentMessage> finalMessage) {
|
||||
AgentMessage partialMessage = partialAssistantMessage(finalText, finalMessage);
|
||||
if (partialMessage == null) {
|
||||
return;
|
||||
if (partialMessage != null) {
|
||||
agent.getMemory().addMessage(messageAdapter.toMsg(partialMessage));
|
||||
}
|
||||
agent.getMemory().addMessage(messageAdapter.toMsg(partialMessage));
|
||||
saveSession();
|
||||
}
|
||||
|
||||
@@ -1109,7 +1095,8 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
context.setAgentDefinition(request.getAgentDefinition());
|
||||
context.setRuntimeContext(request.getRuntimeContext());
|
||||
context.setToolInvokers(request.getToolInvokers());
|
||||
context.setKnowledgeRetrievers(request.getKnowledgeRetrievers());
|
||||
context.setKnowledgeRegistrations(request.getKnowledgeRegistrations());
|
||||
context.setMemorySnapshot(request.getMemorySnapshot());
|
||||
context.setSessionStore(request.getSessionStore());
|
||||
context.setConversationRecorder(request.getConversationRecorder());
|
||||
context.setMetadata(request.getMetadata());
|
||||
@@ -1128,9 +1115,9 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
Toolkit toolkit = new Toolkit();
|
||||
AgentScopeToolkitBuildResult toolkitBuildResult = buildToolkit(context, toolkit);
|
||||
Map<String, List<AgentTool>> skillTools = toolkitBuildResult.skillTools();
|
||||
AgentScopeMemoryBuildResult memoryResult = memoryAdapter.createMemoryResult(null, definition.getMemoryPolicy(), model);
|
||||
AgentScopeMemoryBuildResult memoryResult = memoryAdapter.createMemoryResult(
|
||||
context.getMemorySnapshot(), definition.getMemoryPolicy(), model);
|
||||
Memory memory = memoryResult.getMemory();
|
||||
Knowledge knowledge = knowledgeAdapter.createAggregateKnowledge(context, turnContextHolder);
|
||||
SkillBox skillBox = skillAdapter.createSkillBox(definition.getSkillBoxSpec(), toolkit, skillTools,
|
||||
toolkitBuildResult.skillMcpRegistrations());
|
||||
// AutoContextInterceptor 是官方 AutoContextHook 的替代实现。这里仍只注册统一 runtime hook,
|
||||
@@ -1141,7 +1128,10 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
interceptors.add(new AutoContextInterceptor(eventBridge, memoryResult.getAutoContextConfig()));
|
||||
}
|
||||
interceptors.add(new MediaReferenceInterceptor(initRequest.getMediaResolver()));
|
||||
List<AgentToolSpec> runtimeToolSpecs = mergeToolSpecs(definition.getToolSpecs(), toolkitBuildResult.mcpToolSpecs(),
|
||||
List<AgentToolSpec> runtimeToolSpecs = mergeToolSpecs(
|
||||
definition.getToolSpecs(),
|
||||
toolkitBuildResult.knowledgeToolSpecs(),
|
||||
toolkitBuildResult.mcpToolSpecs(),
|
||||
toolkitBuildResult.operateToolSpecs());
|
||||
interceptors.add(new ToolHitlInterceptor(eventBridge, approvalCoordinator,
|
||||
runtimeToolSpecs));
|
||||
@@ -1156,7 +1146,7 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
ReActAgent.Builder builder = ReActAgent.builder()
|
||||
.name(definition.getAgentName())
|
||||
.description(definition.getDescription())
|
||||
.sysPrompt(systemPrompt(definition))
|
||||
.sysPrompt(SystemPromptComposer.compose(definition))
|
||||
.model(model)
|
||||
.toolkit(toolkit)
|
||||
.memory(memory)
|
||||
@@ -1165,41 +1155,12 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
.hook(new AgentScopeRuntimeHook(observationManager))
|
||||
.enablePendingToolRecovery(true)
|
||||
.statePersistence(AgentScopeSessionAdapter.toStatePersistence(definition.getPersistencePolicy()));
|
||||
if (knowledge != null) {
|
||||
builder.knowledge(knowledge)
|
||||
.ragMode(RAGMode.AGENTIC)
|
||||
.retrieveConfig(defaultRetrieveConfig(definition));
|
||||
}
|
||||
if (skillBox != null) {
|
||||
builder.skillBox(skillBox);
|
||||
}
|
||||
return builder.build();
|
||||
}
|
||||
|
||||
private String systemPrompt(AgentDefinition definition) {
|
||||
String prompt = definition.getSystemPrompt();
|
||||
if (!hasAsyncTool(definition)) {
|
||||
return prompt;
|
||||
}
|
||||
if (prompt == null || prompt.isBlank()) {
|
||||
return ASYNC_TOOL_SYSTEM_PROMPT.strip();
|
||||
}
|
||||
return prompt.stripTrailing() + ASYNC_TOOL_SYSTEM_PROMPT;
|
||||
}
|
||||
|
||||
private boolean hasAsyncTool(AgentDefinition definition) {
|
||||
if (definition == null || definition.getToolSpecs() == null) {
|
||||
return false;
|
||||
}
|
||||
for (AgentToolSpec toolSpec : definition.getToolSpecs()) {
|
||||
// AsyncToolSpecExpander marks all generated sub-tools with this runtime metadata.
|
||||
if (toolSpec != null && Boolean.TRUE.equals(toolSpec.getMetadata().get("asyncTool"))) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建 AgentScope Toolkit,并返回按 Skill ID 分组的工具。
|
||||
*
|
||||
@@ -1211,8 +1172,11 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
Toolkit toolkit) {
|
||||
Map<String, List<AgentTool>> skillTools = new LinkedHashMap<>();
|
||||
if (!context.getAgentDefinition().getExecutionOptions().isToolCallingEnabled()) {
|
||||
return new AgentScopeToolkitBuildResult(skillTools, List.of(), List.of(), List.of());
|
||||
return new AgentScopeToolkitBuildResult(skillTools, List.of(), List.of(), List.of(), List.of());
|
||||
}
|
||||
List<AgentToolSpec> knowledgeToolSpecs = knowledgeAdapter.createToolSpecs(context);
|
||||
validateRuntimeToolConflicts(context.getAgentDefinition().getToolSpecs(), knowledgeToolSpecs,
|
||||
List.of(), List.of());
|
||||
for (AgentToolSpec toolSpec : context.getAgentDefinition().getToolSpecs()) {
|
||||
AgentToolInvoker invoker = context.getToolInvokers().get(toolSpec.getName());
|
||||
AgentSkillBinding skillBinding = skillContext.getToolBinding(toolSpec.getName());
|
||||
@@ -1224,6 +1188,8 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
skillTools.computeIfAbsent(skillBinding.getSkillId(), key -> new ArrayList<>()).add(agentTool);
|
||||
}
|
||||
}
|
||||
knowledgeAdapter.registerTools(context, knowledgeToolSpecs, toolkit, toolAdapter,
|
||||
approvalCoordinator, turnContextHolder);
|
||||
McpRegistration mcpRegistration = mcpToolkitAdapter.register(
|
||||
context.getAgentDefinition().getMcpSpecs(), toolkit);
|
||||
mcpClients.addAll(mcpRegistration.getClients());
|
||||
@@ -1231,17 +1197,32 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
context.getAgentDefinition().getOperateToolSpecs(), toolkit);
|
||||
McpSpecValidator.validateToolConflicts(context.getAgentDefinition().getToolSpecs(),
|
||||
mcpRegistration.getToolSpecs(), context.getAgentDefinition().getOperateToolSpecs());
|
||||
return new AgentScopeToolkitBuildResult(skillTools, mcpRegistration.getToolSpecs(), operateToolSpecs,
|
||||
mcpRegistration.getSkillRegistrations());
|
||||
validateRuntimeToolConflicts(context.getAgentDefinition().getToolSpecs(), knowledgeToolSpecs,
|
||||
mcpRegistration.getToolSpecs(), operateToolSpecs);
|
||||
return new AgentScopeToolkitBuildResult(skillTools, knowledgeToolSpecs,
|
||||
mcpRegistration.getToolSpecs(), operateToolSpecs, mcpRegistration.getSkillRegistrations());
|
||||
}
|
||||
|
||||
/**
|
||||
* 合并所有运行时工具声明供统一治理与事件展示使用。
|
||||
*
|
||||
* @param toolSpecs 普通工具声明
|
||||
* @param knowledgeToolSpecs 知识库工具声明
|
||||
* @param mcpToolSpecs MCP 工具声明
|
||||
* @param operateToolSpecs 操作工具声明
|
||||
* @return 保持注册顺序的工具声明列表
|
||||
*/
|
||||
private List<AgentToolSpec> mergeToolSpecs(List<AgentToolSpec> toolSpecs,
|
||||
List<AgentToolSpec> knowledgeToolSpecs,
|
||||
List<AgentToolSpec> mcpToolSpecs,
|
||||
List<AgentToolSpec> operateToolSpecs) {
|
||||
List<AgentToolSpec> merged = new ArrayList<>();
|
||||
if (toolSpecs != null) {
|
||||
merged.addAll(toolSpecs);
|
||||
}
|
||||
if (knowledgeToolSpecs != null) {
|
||||
merged.addAll(knowledgeToolSpecs);
|
||||
}
|
||||
if (mcpToolSpecs != null) {
|
||||
merged.addAll(mcpToolSpecs);
|
||||
}
|
||||
@@ -1251,6 +1232,46 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
return merged;
|
||||
}
|
||||
|
||||
/**
|
||||
* 校验不同来源的运行时工具名称没有冲突。
|
||||
*
|
||||
* @param toolSpecs 普通工具声明
|
||||
* @param knowledgeToolSpecs 知识库工具声明
|
||||
* @param mcpToolSpecs MCP 工具声明
|
||||
* @param operateToolSpecs 操作工具声明
|
||||
* @throws AgentRuntimeException 工具名称重复时抛出
|
||||
*/
|
||||
private void validateRuntimeToolConflicts(List<AgentToolSpec> toolSpecs,
|
||||
List<AgentToolSpec> knowledgeToolSpecs,
|
||||
List<AgentToolSpec> mcpToolSpecs,
|
||||
List<AgentToolSpec> operateToolSpecs) {
|
||||
Set<String> names = new LinkedHashSet<>();
|
||||
for (List<AgentToolSpec> specs : List.of(
|
||||
safeToolSpecs(toolSpecs),
|
||||
safeToolSpecs(knowledgeToolSpecs),
|
||||
safeToolSpecs(mcpToolSpecs),
|
||||
safeToolSpecs(operateToolSpecs))) {
|
||||
for (AgentToolSpec spec : specs) {
|
||||
if (spec == null || spec.getName() == null || spec.getName().isBlank()) {
|
||||
continue;
|
||||
}
|
||||
if (!names.add(spec.getName())) {
|
||||
throw new AgentRuntimeException("Agent runtime tool name conflict: " + spec.getName());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 将可空工具列表转换为空安全列表。
|
||||
*
|
||||
* @param toolSpecs 工具声明
|
||||
* @return 非空工具声明列表
|
||||
*/
|
||||
private List<AgentToolSpec> safeToolSpecs(List<AgentToolSpec> toolSpecs) {
|
||||
return toolSpecs == null ? List.of() : toolSpecs;
|
||||
}
|
||||
|
||||
private void closeMcpClients() {
|
||||
for (McpClientWrapper client : mcpClients) {
|
||||
if (client == null) {
|
||||
@@ -1264,28 +1285,6 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
mcpClients.clear();
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建聚合知识库的默认检索配置。
|
||||
*
|
||||
* @param definition 智能体定义
|
||||
* @return 检索配置
|
||||
*/
|
||||
private RetrieveConfig defaultRetrieveConfig(AgentDefinition definition) {
|
||||
int limit = definition.getKnowledgeSpecs().stream()
|
||||
.mapToInt(AgentKnowledgeSpec::getLimit)
|
||||
.filter(value -> value > 0)
|
||||
.sum();
|
||||
double scoreThreshold = definition.getKnowledgeSpecs().stream()
|
||||
.mapToDouble(AgentKnowledgeSpec::getScoreThreshold)
|
||||
.filter(value -> value > 0D)
|
||||
.min()
|
||||
.orElse(0D);
|
||||
return RetrieveConfig.builder()
|
||||
.limit(limit <= 0 ? 5 : limit)
|
||||
.scoreThreshold(scoreThreshold)
|
||||
.build();
|
||||
}
|
||||
|
||||
public AgentInitRequest getInitRequest() {
|
||||
return initRequest;
|
||||
}
|
||||
@@ -1300,6 +1299,7 @@ public class AgentScopeReActRuntime implements AgentRuntime {
|
||||
}
|
||||
|
||||
private record AgentScopeToolkitBuildResult(Map<String, List<AgentTool>> skillTools,
|
||||
List<AgentToolSpec> knowledgeToolSpecs,
|
||||
List<AgentToolSpec> mcpToolSpecs,
|
||||
List<AgentToolSpec> operateToolSpecs,
|
||||
List<McpSkillRegistration> skillMcpRegistrations) {
|
||||
|
||||
@@ -170,7 +170,8 @@ public class AgentScopeToolAdapter {
|
||||
throw new AgentRuntimeException("Agent tool invoker is required: " + toolSpec.getName());
|
||||
}
|
||||
return new RuntimeAgentTool(toolSpec, invoker, request, approvalCoordinator, turnContextHolder,
|
||||
skillContext, skillBinding, emitNormalToolResult, true, true);
|
||||
skillContext, skillBinding, emitNormalToolResult, true, true,
|
||||
resolveInvocationClassLoader(invoker));
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -203,7 +204,8 @@ public class AgentScopeToolAdapter {
|
||||
throw new AgentRuntimeException("Agent tool invoker is required: " + toolSpec.getName());
|
||||
}
|
||||
return new RuntimeAgentTool(toolSpec, invoker, request, approvalCoordinator, turnContextHolder,
|
||||
skillContext, skillBinding, emitNormalToolResult, emitSkillStep, true);
|
||||
skillContext, skillBinding, emitNormalToolResult, emitSkillStep, true,
|
||||
resolveInvocationClassLoader(invoker));
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -238,7 +240,24 @@ public class AgentScopeToolAdapter {
|
||||
throw new AgentRuntimeException("Agent tool invoker is required: " + toolSpec.getName());
|
||||
}
|
||||
return new RuntimeAgentTool(toolSpec, invoker, request, approvalCoordinator, turnContextHolder,
|
||||
skillContext, skillBinding, emitNormalToolResult, emitSkillStep, handleApprovalInTool);
|
||||
skillContext, skillBinding, emitNormalToolResult, emitSkillStep, handleApprovalInTool,
|
||||
resolveInvocationClassLoader(invoker));
|
||||
}
|
||||
|
||||
/**
|
||||
* 解析工具执行时应使用的应用类加载器。
|
||||
*
|
||||
* <p>AgentScope 可能在 Reactor 工作线程执行工具。Spring Boot 可执行包中的业务类依赖
|
||||
* 注册工具时的应用类加载器,不能依赖工作线程可能继承到的系统类加载器。</p>
|
||||
*
|
||||
* @param invoker 工具调用器
|
||||
* @return 工具调用器所属类加载器;无法取得时回退到当前线程上下文类加载器
|
||||
*/
|
||||
private ClassLoader resolveInvocationClassLoader(AgentToolInvoker invoker) {
|
||||
ClassLoader invokerClassLoader = invoker.getClass().getClassLoader();
|
||||
return invokerClassLoader == null
|
||||
? Thread.currentThread().getContextClassLoader()
|
||||
: invokerClassLoader;
|
||||
}
|
||||
|
||||
private AgentRuntimeTurnContextHolder fixedHolder(AgentRuntimeExecutionContext request,
|
||||
@@ -258,7 +277,8 @@ public class AgentScopeToolAdapter {
|
||||
AgentSkillBinding skillBinding,
|
||||
boolean emitNormalToolResult,
|
||||
boolean emitSkillStep,
|
||||
boolean handleApprovalInTool) implements AgentTool {
|
||||
boolean handleApprovalInTool,
|
||||
ClassLoader invocationClassLoader) implements AgentTool {
|
||||
|
||||
/**
|
||||
* 获取工具名称。
|
||||
@@ -402,15 +422,28 @@ public class AgentScopeToolAdapter {
|
||||
* @return 工具结果块
|
||||
*/
|
||||
private ToolResultBlock invokeTool(ToolCallParam param, Map<String, Object> input) {
|
||||
AgentToolContext context = buildContext(param);
|
||||
AgentToolResult result = invoker.invoke(input, context);
|
||||
ToolResultBlock block = toToolResultBlock(param, result);
|
||||
// 有状态 runtime 中,普通工具结果由 AgentScope 原生 PostActingEvent
|
||||
// 旁路观察器统一发出;旧 sink 辅助路径没有统一 hook,因此仍允许 adapter 兼容发射。
|
||||
if (emitNormalToolResult || (emitSkillStep && activeSkillBinding() != null)) {
|
||||
emit(toolResultEvent(block));
|
||||
Thread currentThread = Thread.currentThread();
|
||||
ClassLoader originalClassLoader = currentThread.getContextClassLoader();
|
||||
boolean switchClassLoader = invocationClassLoader != null
|
||||
&& invocationClassLoader != originalClassLoader;
|
||||
if (switchClassLoader) {
|
||||
currentThread.setContextClassLoader(invocationClassLoader);
|
||||
}
|
||||
try {
|
||||
AgentToolContext context = buildContext(param);
|
||||
AgentToolResult result = invoker.invoke(input, context);
|
||||
ToolResultBlock block = toToolResultBlock(param, result);
|
||||
// 有状态 runtime 中,普通工具结果由 AgentScope 原生 PostActingEvent
|
||||
// 旁路观察器统一发出;旧 sink 辅助路径没有统一 hook,因此仍允许 adapter 兼容发射。
|
||||
if (emitNormalToolResult || (emitSkillStep && activeSkillBinding() != null)) {
|
||||
emit(toolResultEvent(block));
|
||||
}
|
||||
return block;
|
||||
} finally {
|
||||
if (switchClassLoader) {
|
||||
currentThread.setContextClassLoader(originalClassLoader);
|
||||
}
|
||||
}
|
||||
return block;
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -107,9 +107,9 @@ public class AgentRuntimeTurnContext {
|
||||
merged.setUserMessage(executionContext.getUserMessage());
|
||||
merged.setMemorySnapshot(executionContext.getMemorySnapshot());
|
||||
merged.setToolInvokers(fallback == null ? executionContext.getToolInvokers() : fallback.getToolInvokers());
|
||||
merged.setKnowledgeRetrievers(fallback == null
|
||||
? executionContext.getKnowledgeRetrievers()
|
||||
: fallback.getKnowledgeRetrievers());
|
||||
merged.setKnowledgeRegistrations(fallback == null
|
||||
? executionContext.getKnowledgeRegistrations()
|
||||
: fallback.getKnowledgeRegistrations());
|
||||
merged.setSessionStore(fallback == null ? executionContext.getSessionStore() : fallback.getSessionStore());
|
||||
merged.setConversationRecorder(fallback == null
|
||||
? executionContext.getConversationRecorder()
|
||||
|
||||
@@ -5,10 +5,12 @@ import com.easyagents.agent.runtime.event.AgentRuntimeEventBridge;
|
||||
import com.easyagents.agent.runtime.event.AgentRuntimeEventType;
|
||||
import com.easyagents.agent.runtime.event.AgentRuntimeObserver;
|
||||
import com.easyagents.agent.runtime.skill.AgentSkillRuntimeContext;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolCategory;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolSpec;
|
||||
import io.agentscope.core.hook.HookEvent;
|
||||
import io.agentscope.core.hook.PostActingEvent;
|
||||
import io.agentscope.core.hook.PreActingEvent;
|
||||
import io.agentscope.core.message.TextBlock;
|
||||
import io.agentscope.core.message.ToolResultBlock;
|
||||
import io.agentscope.core.message.ToolUseBlock;
|
||||
import reactor.core.publisher.Mono;
|
||||
@@ -130,12 +132,21 @@ public class ToolExecutionObserver implements AgentRuntimeObserver {
|
||||
|
||||
private void enrichToolPayload(AgentRuntimeEvent runtimeEvent, String toolName) {
|
||||
AgentToolSpec toolSpec = toolSpecs.get(toolName);
|
||||
if (toolSpec == null || toolSpec.getMetadata() == null || toolSpec.getMetadata().isEmpty()) {
|
||||
if (toolSpec == null) {
|
||||
return;
|
||||
}
|
||||
runtimeEvent.getPayload().put("toolCategory", toolSpec.getCategory().name());
|
||||
if (toolSpec.getMetadata() == null || toolSpec.getMetadata().isEmpty()) {
|
||||
return;
|
||||
}
|
||||
Map<String, Object> metadata = toolSpec.getMetadata();
|
||||
putIfPresent(runtimeEvent.getPayload(), metadata, "toolDisplayName");
|
||||
putIfPresent(runtimeEvent.getPayload(), metadata, "skillId");
|
||||
if (toolSpec.getCategory() == AgentToolCategory.KNOWLEDGE) {
|
||||
putIfPresent(runtimeEvent.getPayload(), metadata, "knowledgeId");
|
||||
putIfPresent(runtimeEvent.getPayload(), metadata, "knowledgeName");
|
||||
putIfPresent(runtimeEvent.getPayload(), metadata, "knowledgeRuntimeName");
|
||||
}
|
||||
}
|
||||
|
||||
private void putIfPresent(Map<String, Object> payload, Map<String, Object> metadata, String key) {
|
||||
@@ -149,7 +160,15 @@ public class ToolExecutionObserver implements AgentRuntimeObserver {
|
||||
return false;
|
||||
}
|
||||
Object success = result.getMetadata() == null ? null : result.getMetadata().get("success");
|
||||
return !(success instanceof Boolean) || Boolean.TRUE.equals(success);
|
||||
if (success instanceof Boolean) {
|
||||
return Boolean.TRUE.equals(success);
|
||||
}
|
||||
// AgentScope 1.x 将工具异常转换为不带 success metadata 的 "Error: ..." 文本结果。
|
||||
return result.getOutput().stream()
|
||||
.filter(TextBlock.class::isInstance)
|
||||
.map(TextBlock.class::cast)
|
||||
.map(TextBlock::getText)
|
||||
.noneMatch(text -> text != null && text.startsWith("Error: "));
|
||||
}
|
||||
|
||||
private boolean isSkillTool(String toolName) {
|
||||
|
||||
@@ -1,10 +0,0 @@
|
||||
package com.easyagents.agent.runtime.knowledge;
|
||||
|
||||
/**
|
||||
* 知识库检索策略。
|
||||
*/
|
||||
public enum AgentKnowledgePolicy {
|
||||
AGENTIC,
|
||||
GENERIC,
|
||||
DISABLED
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package com.easyagents.agent.runtime.knowledge;
|
||||
|
||||
import java.util.Objects;
|
||||
|
||||
/**
|
||||
* 知识库声明与检索器的运行时绑定。
|
||||
*/
|
||||
public final class AgentKnowledgeRegistration {
|
||||
|
||||
private final AgentKnowledgeSpec knowledgeSpec;
|
||||
private final AgentKnowledgeRetriever retriever;
|
||||
|
||||
/**
|
||||
* 创建知识库运行时绑定。
|
||||
*
|
||||
* @param knowledgeSpec 知识库声明
|
||||
* @param retriever 知识库检索器
|
||||
* @throws NullPointerException 声明或检索器为空时抛出
|
||||
*/
|
||||
public AgentKnowledgeRegistration(AgentKnowledgeSpec knowledgeSpec,
|
||||
AgentKnowledgeRetriever retriever) {
|
||||
this.knowledgeSpec = Objects.requireNonNull(knowledgeSpec, "knowledgeSpec");
|
||||
this.retriever = Objects.requireNonNull(retriever, "retriever");
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取知识库声明。
|
||||
*
|
||||
* @return 知识库声明
|
||||
*/
|
||||
public AgentKnowledgeSpec getKnowledgeSpec() {
|
||||
return knowledgeSpec;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取知识库检索器。
|
||||
*
|
||||
* @return 知识库检索器
|
||||
*/
|
||||
public AgentKnowledgeRetriever getRetriever() {
|
||||
return retriever;
|
||||
}
|
||||
}
|
||||
@@ -9,9 +9,9 @@ import java.util.Map;
|
||||
public class AgentKnowledgeSpec {
|
||||
|
||||
private String knowledgeId;
|
||||
private String runtimeName;
|
||||
private String name;
|
||||
private String description;
|
||||
private AgentKnowledgePolicy retrievalMode = AgentKnowledgePolicy.AGENTIC;
|
||||
private int limit = 5;
|
||||
private double scoreThreshold = 0D;
|
||||
private Map<String, Object> metadata = new LinkedHashMap<>();
|
||||
@@ -34,6 +34,24 @@ public class AgentKnowledgeSpec {
|
||||
this.knowledgeId = knowledgeId;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取调用方提供的英文运行名。
|
||||
*
|
||||
* @return 英文运行名
|
||||
*/
|
||||
public String getRuntimeName() {
|
||||
return runtimeName;
|
||||
}
|
||||
|
||||
/**
|
||||
* 设置调用方提供的英文运行名。
|
||||
*
|
||||
* @param runtimeName 英文运行名
|
||||
*/
|
||||
public void setRuntimeName(String runtimeName) {
|
||||
this.runtimeName = runtimeName;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取知识库名称。
|
||||
*
|
||||
@@ -70,24 +88,6 @@ public class AgentKnowledgeSpec {
|
||||
this.description = description;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取检索模式。
|
||||
*
|
||||
* @return 检索模式
|
||||
*/
|
||||
public AgentKnowledgePolicy getRetrievalMode() {
|
||||
return retrievalMode;
|
||||
}
|
||||
|
||||
/**
|
||||
* 设置检索模式。
|
||||
*
|
||||
* @param retrievalMode 检索模式
|
||||
*/
|
||||
public void setRetrievalMode(AgentKnowledgePolicy retrievalMode) {
|
||||
this.retrievalMode = retrievalMode == null ? AgentKnowledgePolicy.AGENTIC : retrievalMode;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取限制数量。
|
||||
*
|
||||
|
||||
@@ -0,0 +1,79 @@
|
||||
package com.easyagents.agent.runtime.knowledge;
|
||||
|
||||
import com.easyagents.agent.runtime.AgentRuntimeException;
|
||||
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.security.MessageDigest;
|
||||
import java.security.NoSuchAlgorithmException;
|
||||
import java.util.HexFormat;
|
||||
import java.util.regex.Pattern;
|
||||
|
||||
/**
|
||||
* Agent 知识库工具名称规范。
|
||||
*/
|
||||
public final class AgentKnowledgeToolNames {
|
||||
|
||||
/**
|
||||
* 知识库工具名称前缀。
|
||||
*/
|
||||
public static final String PREFIX = "retrieve_knowledge_";
|
||||
|
||||
/**
|
||||
* OpenAI-compatible Function Call 的通用名称长度上限。
|
||||
*/
|
||||
public static final int MAX_TOOL_NAME_LENGTH = 64;
|
||||
|
||||
private static final int HASH_LENGTH = 8;
|
||||
private static final Pattern SAFE_RUNTIME_NAME = Pattern.compile("^[A-Za-z0-9_-]+$");
|
||||
|
||||
/**
|
||||
* 禁止实例化工具类。
|
||||
*/
|
||||
private AgentKnowledgeToolNames() {
|
||||
}
|
||||
|
||||
/**
|
||||
* 根据调用方运行名生成稳定的知识库工具名。
|
||||
*
|
||||
* <p>超长名称保留可读前缀并追加稳定短哈希,避免不同模型服务对 Function Call
|
||||
* 名称长度限制不一致。</p>
|
||||
*
|
||||
* @param runtimeName 调用方提供的英文运行名
|
||||
* @return 合法且稳定的工具名
|
||||
* @throws AgentRuntimeException 运行名为空或包含非法字符时抛出
|
||||
*/
|
||||
public static String build(String runtimeName) {
|
||||
String normalized = runtimeName == null ? "" : runtimeName.trim();
|
||||
if (normalized.isEmpty()) {
|
||||
throw new AgentRuntimeException("Knowledge runtime name is required.");
|
||||
}
|
||||
if (!SAFE_RUNTIME_NAME.matcher(normalized).matches()) {
|
||||
throw new AgentRuntimeException(
|
||||
"Knowledge runtime name must contain only letters, numbers, underscores, or hyphens: "
|
||||
+ normalized);
|
||||
}
|
||||
String toolName = PREFIX + normalized;
|
||||
if (toolName.length() <= MAX_TOOL_NAME_LENGTH) {
|
||||
return toolName;
|
||||
}
|
||||
String hash = shortHash(normalized);
|
||||
int readableLength = MAX_TOOL_NAME_LENGTH - PREFIX.length() - HASH_LENGTH - 1;
|
||||
return PREFIX + normalized.substring(0, readableLength) + "_" + hash;
|
||||
}
|
||||
|
||||
/**
|
||||
* 计算用于超长工具名消歧的稳定短哈希。
|
||||
*
|
||||
* @param value 原始运行名
|
||||
* @return 八位十六进制哈希
|
||||
*/
|
||||
private static String shortHash(String value) {
|
||||
try {
|
||||
byte[] digest = MessageDigest.getInstance("SHA-256")
|
||||
.digest(value.getBytes(StandardCharsets.UTF_8));
|
||||
return HexFormat.of().formatHex(digest).substring(0, HASH_LENGTH);
|
||||
} catch (NoSuchAlgorithmException exception) {
|
||||
throw new IllegalStateException("SHA-256 is unavailable.", exception);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -34,9 +34,16 @@ public class HeuristicKnowledgeCitationMatcher implements AgentKnowledgeCitation
|
||||
if (normalizedAnswer.length() < MIN_NORMALIZED_ANSWER_LENGTH) {
|
||||
return List.of();
|
||||
}
|
||||
List<String> normalizedSegments = normalizeSegments(answerText);
|
||||
List<ScoredKnowledgeReference> scoredReferences = new ArrayList<>();
|
||||
for (AgentKnowledgeReference candidate : candidates) {
|
||||
double supportScore = supportScore(normalizedAnswer, normalize(candidate == null ? null : candidate.getChunkContent()));
|
||||
String normalizedContent = normalize(candidate == null ? null : candidate.getChunkContent());
|
||||
double supportScore = supportScore(normalizedAnswer, normalizedContent);
|
||||
// 长篇汇总回答中,每条证据通常只支撑一个段落或条目。继续只用整篇答案作分母,
|
||||
// 会随着主题增多把有效引用的重合度稀释到阈值以下。
|
||||
for (String normalizedSegment : normalizedSegments) {
|
||||
supportScore = Math.max(supportScore, supportScore(normalizedSegment, normalizedContent));
|
||||
}
|
||||
if (supportScore >= MIN_SUPPORT_SCORE) {
|
||||
scoredReferences.add(new ScoredKnowledgeReference(candidate, supportScore));
|
||||
}
|
||||
@@ -48,6 +55,22 @@ public class HeuristicKnowledgeCitationMatcher implements AgentKnowledgeCitation
|
||||
.toList();
|
||||
}
|
||||
|
||||
/**
|
||||
* 将答案切分为可独立核验的段落或句子并完成归一化。
|
||||
*
|
||||
* @param answerText 最终答案文本
|
||||
* @return 非空的归一化答案片段
|
||||
*/
|
||||
private List<String> normalizeSegments(String answerText) {
|
||||
if (answerText == null || answerText.isBlank()) {
|
||||
return List.of();
|
||||
}
|
||||
return Arrays.stream(answerText.split("[\\r\\n。!?!?;;]+"))
|
||||
.map(this::normalize)
|
||||
.filter(segment -> segment.length() >= MIN_NORMALIZED_ANSWER_LENGTH)
|
||||
.toList();
|
||||
}
|
||||
|
||||
/**
|
||||
* 计算答案与候选片段之间的文本支撑分。
|
||||
*
|
||||
|
||||
@@ -0,0 +1,105 @@
|
||||
package com.easyagents.agent.runtime.prompt;
|
||||
|
||||
import com.easyagents.agent.runtime.AgentDefinition;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolSpec;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
|
||||
/**
|
||||
* 组合智能体运行时使用的系统提示词。
|
||||
*/
|
||||
public final class SystemPromptComposer {
|
||||
|
||||
private static final String KNOWLEDGE_TOOL_PROTOCOL = """
|
||||
Knowledge tool protocol:
|
||||
- Knowledge tools provide grounded information for the scopes described in their tool descriptions.
|
||||
- Before answering a request involving facts, policies, procedures, entitlements, prices, product or service details, or other information within a knowledge tool's scope, call the most relevant knowledge tool first.
|
||||
- Do not rely only on general model knowledge for claims that fall within an available knowledge tool's scope.
|
||||
- If it is uncertain whether the request falls within a knowledge tool's scope, prefer making one retrieval call.
|
||||
- Build a concise, standalone query from the current request and only the necessary conversation context. Resolve pronouns and omitted subjects, and preserve relevant names, time, location, product, membership level, constraints, and user intent. Do not copy the entire conversation.
|
||||
- If the returned results are empty or do not address the request, reformulate the query and retry once when a materially different query is possible. Do not repeat the same query or retrieve indefinitely.
|
||||
- Base knowledge-backed claims only on information actually returned by the tools. Never imply that a knowledge base contains information that was not returned.
|
||||
- If retrieval remains insufficient, follow the Agent's configured system prompt for the response strategy.
|
||||
- Skip retrieval for greetings, casual conversation, pure writing or translation, and tasks that clearly do not depend on knowledge-base facts.
|
||||
""".strip();
|
||||
|
||||
private static final String ASYNC_TOOL_PROTOCOL = """
|
||||
Async tool protocol:
|
||||
- Async tools may expose submit, observe, result, cancel, and list sub-tools. Treat these sub-tools as one user-facing tool.
|
||||
- Do not ask the user to choose submit, observe, result, cancel, or list. These are internal execution phases.
|
||||
- For a normal user request to use an async tool, call its submit sub-tool first with the user-provided arguments by default.
|
||||
- After submit returns task_id, immediately call observe with that task_id to check progress.
|
||||
- If the task is completed and result is available, use the returned result to answer the user.
|
||||
- If the task is still running after observation, tell the user that the task is running and keep task_id/next_action for later tool calls.
|
||||
- Use result, list, or cancel directly only when the user explicitly asks to get a known task result, list tasks, or cancel a task.
|
||||
""".strip();
|
||||
|
||||
/**
|
||||
* 阻止工具类被实例化。
|
||||
*/
|
||||
private SystemPromptComposer() {
|
||||
}
|
||||
|
||||
/**
|
||||
* 按固定顺序组合用户提示词与运行时工具协议。
|
||||
*
|
||||
* @param definition 智能体定义
|
||||
* @return 最终系统提示词;没有任何提示词时返回原始空值
|
||||
*/
|
||||
public static String compose(AgentDefinition definition) {
|
||||
if (definition == null) {
|
||||
return null;
|
||||
}
|
||||
boolean knowledgeProtocolEnabled = hasEnabledKnowledgeTool(definition);
|
||||
boolean asyncProtocolEnabled = hasAsyncTool(definition);
|
||||
String userPrompt = definition.getSystemPrompt();
|
||||
if (!knowledgeProtocolEnabled && !asyncProtocolEnabled) {
|
||||
return userPrompt;
|
||||
}
|
||||
|
||||
List<String> promptSections = new ArrayList<>(3);
|
||||
if (userPrompt != null && !userPrompt.isBlank()) {
|
||||
promptSections.add(userPrompt.stripTrailing());
|
||||
}
|
||||
if (knowledgeProtocolEnabled) {
|
||||
promptSections.add(KNOWLEDGE_TOOL_PROTOCOL);
|
||||
}
|
||||
if (asyncProtocolEnabled) {
|
||||
promptSections.add(ASYNC_TOOL_PROTOCOL);
|
||||
}
|
||||
return String.join("\n\n", promptSections);
|
||||
}
|
||||
|
||||
/**
|
||||
* 判断当前定义是否会注册可调用的知识库工具。
|
||||
*
|
||||
* @param definition 智能体定义
|
||||
* @return 知识库工具可用时返回 true
|
||||
*/
|
||||
private static boolean hasEnabledKnowledgeTool(AgentDefinition definition) {
|
||||
return definition.getExecutionOptions() != null
|
||||
&& definition.getExecutionOptions().isToolCallingEnabled()
|
||||
&& definition.getKnowledgeSpecs() != null
|
||||
&& !definition.getKnowledgeSpecs().isEmpty();
|
||||
}
|
||||
|
||||
/**
|
||||
* 判断当前定义是否包含异步工具。
|
||||
*
|
||||
* @param definition 智能体定义
|
||||
* @return 包含异步工具时返回 true
|
||||
*/
|
||||
private static boolean hasAsyncTool(AgentDefinition definition) {
|
||||
if (definition.getToolSpecs() == null) {
|
||||
return false;
|
||||
}
|
||||
for (AgentToolSpec toolSpec : definition.getToolSpecs()) {
|
||||
// AsyncToolSpecExpander 会为生成的全部子工具写入该运行时元数据。
|
||||
if (toolSpec != null && Boolean.TRUE.equals(toolSpec.getMetadata().get("asyncTool"))) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,226 @@
|
||||
package com.easyagents.agent.runtime.agentscope;
|
||||
|
||||
import com.easyagents.agent.runtime.AgentDefinition;
|
||||
import com.easyagents.agent.runtime.AgentRuntimeExecutionContext;
|
||||
import com.easyagents.agent.runtime.event.AgentRuntimeEvent;
|
||||
import com.easyagents.agent.runtime.event.AgentRuntimeEventBridge;
|
||||
import com.easyagents.agent.runtime.event.AgentRuntimeEventType;
|
||||
import com.easyagents.agent.runtime.event.AgentRuntimeTurnContext;
|
||||
import com.easyagents.agent.runtime.event.AgentRuntimeTurnContextHolder;
|
||||
import com.easyagents.agent.runtime.hitl.AgentToolApprovalCoordinator;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeDocument;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeRegistration;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeRetrievalResult;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeSpec;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeToolNames;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolSpec;
|
||||
import io.agentscope.core.message.ToolResultBlock;
|
||||
import io.agentscope.core.message.ToolUseBlock;
|
||||
import io.agentscope.core.tool.ToolCallParam;
|
||||
import io.agentscope.core.tool.Toolkit;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
import reactor.core.publisher.Sinks;
|
||||
|
||||
import java.time.Duration;
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
|
||||
/**
|
||||
* {@link AgentScopeKnowledgeAdapter} 回归测试。
|
||||
*/
|
||||
public class AgentScopeKnowledgeAdapterTest {
|
||||
|
||||
/**
|
||||
* 验证每个知识库生成独立工具,工具名、描述和 Schema 向模型暴露完整信息。
|
||||
*/
|
||||
@Test
|
||||
public void shouldCreateOneToolPerKnowledge() {
|
||||
AgentKnowledgeSpec first = knowledgeSpec("knowledge-1", "homeinn_faq", 5);
|
||||
first.setName("如家 FAQ");
|
||||
first.setDescription("如家酒店入住、退房和会员服务常见问题");
|
||||
AgentKnowledgeSpec second = knowledgeSpec("knowledge-2", "hotel_policy", 3);
|
||||
AgentRuntimeExecutionContext context = executionContext(List.of(first, second));
|
||||
context.setKnowledgeRegistrations(List.of(
|
||||
registration(first, 1, 0.9D),
|
||||
registration(second, 1, 0.8D)));
|
||||
|
||||
List<AgentToolSpec> toolSpecs = new AgentScopeKnowledgeAdapter().createToolSpecs(context);
|
||||
|
||||
Assert.assertEquals(2, toolSpecs.size());
|
||||
Assert.assertEquals("retrieve_knowledge_homeinn_faq", toolSpecs.get(0).getName());
|
||||
Assert.assertTrue(toolSpecs.get(0).getDescription().contains("如家 FAQ"));
|
||||
Assert.assertTrue(toolSpecs.get(0).getDescription().contains("入住、退房"));
|
||||
Assert.assertEquals(List.of("query"), toolSpecs.get(0).getParametersSchema().get("required"));
|
||||
Assert.assertFalse(String.valueOf(toolSpecs.get(0).getParametersSchema()).contains("limit"));
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证选中一个知识库工具时只调用对应 Retriever,并使用绑定 limit 和阈值。
|
||||
*/
|
||||
@Test
|
||||
public void shouldDispatchOnlyToSelectedKnowledgeRetriever() {
|
||||
AgentKnowledgeSpec first = knowledgeSpec("knowledge-1", "homeinn_faq", 2);
|
||||
first.setScoreThreshold(0.5D);
|
||||
AgentKnowledgeSpec second = knowledgeSpec("knowledge-2", "hotel_policy", 2);
|
||||
AgentRuntimeExecutionContext context = executionContext(List.of(first, second));
|
||||
AtomicInteger firstCalls = new AtomicInteger();
|
||||
AtomicInteger secondCalls = new AtomicInteger();
|
||||
context.setKnowledgeRegistrations(List.of(
|
||||
new AgentKnowledgeRegistration(first, request -> {
|
||||
firstCalls.incrementAndGet();
|
||||
Assert.assertEquals("如家几点退房", request.getQuery());
|
||||
Assert.assertEquals(2, request.getLimit());
|
||||
Assert.assertEquals(0.5D, request.getScoreThreshold(), 0.0001D);
|
||||
return AgentKnowledgeRetrievalResult.of(List.of(
|
||||
document("knowledge-1", 0.9D),
|
||||
document("knowledge-1-low", 0.2D)));
|
||||
}),
|
||||
new AgentKnowledgeRegistration(second, request -> {
|
||||
secondCalls.incrementAndGet();
|
||||
return AgentKnowledgeRetrievalResult.of(List.of(document("knowledge-2", 0.8D)));
|
||||
})));
|
||||
RegisteredKnowledgeTools registered = register(context);
|
||||
|
||||
ToolResultBlock result = registered.toolkit().getTool("retrieve_knowledge_homeinn_faq")
|
||||
.callAsync(toolCall("retrieve_knowledge_homeinn_faq", "如家几点退房"))
|
||||
.block();
|
||||
|
||||
Assert.assertNotNull(result);
|
||||
Assert.assertEquals(1, firstCalls.get());
|
||||
Assert.assertEquals(0, secondCalls.get());
|
||||
Assert.assertEquals(1, result.getMetadata().get("documentCount"));
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证最终知识库事件与阈值过滤、排序和截断后的模型可见文档一致。
|
||||
*/
|
||||
@Test
|
||||
public void retrievalEventShouldMatchFinalToolDocuments() {
|
||||
AgentKnowledgeSpec spec = knowledgeSpec("knowledge-1", "homeinn_faq", 2);
|
||||
spec.setScoreThreshold(0.5D);
|
||||
AgentRuntimeExecutionContext context = executionContext(List.of(spec));
|
||||
context.setKnowledgeRegistrations(List.of(new AgentKnowledgeRegistration(spec, request ->
|
||||
AgentKnowledgeRetrievalResult.of(List.of(
|
||||
document("third", 0.7D),
|
||||
document("first", 0.9D),
|
||||
document("filtered", 0.2D),
|
||||
document("second", 0.8D))))));
|
||||
RegisteredKnowledgeTools registered = register(context);
|
||||
|
||||
ToolResultBlock result = registered.toolkit().getTool("retrieve_knowledge_homeinn_faq")
|
||||
.callAsync(toolCall("retrieve_knowledge_homeinn_faq", "query"))
|
||||
.block();
|
||||
List<AgentRuntimeEvent> events = registered.eventSink().asFlux()
|
||||
.filter(event -> event.getEventType() == AgentRuntimeEventType.KNOWLEDGE_RETRIEVAL)
|
||||
.take(1)
|
||||
.collectList()
|
||||
.block(Duration.ofSeconds(1));
|
||||
|
||||
Assert.assertNotNull(result);
|
||||
Assert.assertNotNull(events);
|
||||
Assert.assertEquals(1, events.size());
|
||||
AgentRuntimeEvent event = events.get(0);
|
||||
Assert.assertEquals(2, event.getPayload().get("documentCount"));
|
||||
@SuppressWarnings("unchecked")
|
||||
List<Map<String, Object>> eventDocuments =
|
||||
(List<Map<String, Object>>) event.getPayload().get("documents");
|
||||
@SuppressWarnings("unchecked")
|
||||
List<Map<String, Object>> resultDocuments =
|
||||
(List<Map<String, Object>>) result.getMetadata().get("documents");
|
||||
Assert.assertEquals(resultDocuments, eventDocuments);
|
||||
Assert.assertEquals("first", eventDocuments.get(0).get("documentId"));
|
||||
Assert.assertEquals("second", eventDocuments.get(1).get("documentId"));
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证超长运行名会稳定压缩到 Function Call 长度上限。
|
||||
*/
|
||||
@Test
|
||||
public void longRuntimeNameShouldProduceStableBoundedToolName() {
|
||||
String runtimeName = "knowledge_" + "a".repeat(80);
|
||||
|
||||
String first = AgentKnowledgeToolNames.build(runtimeName);
|
||||
String second = AgentKnowledgeToolNames.build(runtimeName);
|
||||
|
||||
Assert.assertEquals(first, second);
|
||||
Assert.assertEquals(AgentKnowledgeToolNames.MAX_TOOL_NAME_LENGTH, first.length());
|
||||
Assert.assertTrue(first.startsWith(AgentKnowledgeToolNames.PREFIX));
|
||||
}
|
||||
|
||||
private RegisteredKnowledgeTools register(AgentRuntimeExecutionContext context) {
|
||||
AgentScopeKnowledgeAdapter adapter = new AgentScopeKnowledgeAdapter();
|
||||
List<AgentToolSpec> toolSpecs = adapter.createToolSpecs(context);
|
||||
Toolkit toolkit = new Toolkit();
|
||||
Sinks.Many<AgentRuntimeEvent> eventSink = Sinks.many().replay().all();
|
||||
AgentRuntimeTurnContextHolder holder = new AgentRuntimeTurnContextHolder();
|
||||
AgentRuntimeEventBridge eventBridge = new AgentRuntimeEventBridge(context, holder);
|
||||
holder.set(new AgentRuntimeTurnContext(context, eventSink, eventBridge));
|
||||
adapter.registerTools(context, toolSpecs, toolkit, new AgentScopeToolAdapter(),
|
||||
AgentToolApprovalCoordinator.disabled(), holder);
|
||||
return new RegisteredKnowledgeTools(toolkit, eventSink);
|
||||
}
|
||||
|
||||
private ToolCallParam toolCall(String toolName, String query) {
|
||||
ToolUseBlock toolUseBlock = ToolUseBlock.builder()
|
||||
.id("call-1")
|
||||
.name(toolName)
|
||||
.input(Map.of("query", query))
|
||||
.build();
|
||||
return ToolCallParam.builder()
|
||||
.toolUseBlock(toolUseBlock)
|
||||
.input(toolUseBlock.getInput())
|
||||
.build();
|
||||
}
|
||||
|
||||
private AgentRuntimeExecutionContext executionContext(List<AgentKnowledgeSpec> specs) {
|
||||
AgentDefinition definition = new AgentDefinition();
|
||||
definition.setAgentId("agent-1");
|
||||
definition.setKnowledgeSpecs(specs);
|
||||
AgentRuntimeExecutionContext context = new AgentRuntimeExecutionContext();
|
||||
context.setAgentDefinition(definition);
|
||||
context.setTraceId("trace-1");
|
||||
context.setSessionId("session-1");
|
||||
return context;
|
||||
}
|
||||
|
||||
private AgentKnowledgeSpec knowledgeSpec(String knowledgeId, String runtimeName, int limit) {
|
||||
AgentKnowledgeSpec spec = new AgentKnowledgeSpec();
|
||||
spec.setKnowledgeId(knowledgeId);
|
||||
spec.setRuntimeName(runtimeName);
|
||||
spec.setName(knowledgeId);
|
||||
spec.setLimit(limit);
|
||||
return spec;
|
||||
}
|
||||
|
||||
private AgentKnowledgeRegistration registration(AgentKnowledgeSpec spec,
|
||||
int documentCount,
|
||||
double firstScore) {
|
||||
return new AgentKnowledgeRegistration(spec, request ->
|
||||
AgentKnowledgeRetrievalResult.of(documents(spec.getKnowledgeId(), documentCount, firstScore)));
|
||||
}
|
||||
|
||||
private List<AgentKnowledgeDocument> documents(String knowledgeId, int count, double firstScore) {
|
||||
List<AgentKnowledgeDocument> documents = new ArrayList<>();
|
||||
for (int index = 0; index < count; index++) {
|
||||
documents.add(document(knowledgeId + "-doc-" + index, firstScore - (index * 0.01D)));
|
||||
}
|
||||
return documents;
|
||||
}
|
||||
|
||||
private AgentKnowledgeDocument document(String documentId, double score) {
|
||||
AgentKnowledgeDocument document = new AgentKnowledgeDocument();
|
||||
document.setDocumentId(documentId);
|
||||
document.setDocumentName("FAQ");
|
||||
document.setChunkId(documentId + "-chunk");
|
||||
document.setContent("content-" + documentId);
|
||||
document.setScore(score);
|
||||
return document;
|
||||
}
|
||||
|
||||
private record RegisteredKnowledgeTools(Toolkit toolkit,
|
||||
Sinks.Many<AgentRuntimeEvent> eventSink) {
|
||||
}
|
||||
}
|
||||
@@ -10,6 +10,7 @@ import com.easyagents.agent.runtime.event.observer.SkillExecutionObserver;
|
||||
import com.easyagents.agent.runtime.event.observer.ToolExecutionObserver;
|
||||
import com.easyagents.agent.runtime.hitl.AgentResumeToken;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeDocument;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeRegistration;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeRetrievalResult;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeSpec;
|
||||
import com.easyagents.agent.runtime.memory.AgentMemoryCompressionParameter;
|
||||
@@ -23,6 +24,7 @@ import com.easyagents.agent.runtime.persistence.session.memory.InMemoryAgentSess
|
||||
import com.easyagents.agent.runtime.skill.AgentSkillBoxSpec;
|
||||
import com.easyagents.agent.runtime.skill.AgentSkillRuntimeContext;
|
||||
import com.easyagents.agent.runtime.skill.AgentSkillSpec;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolCategory;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolResult;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolSpec;
|
||||
import com.easyagents.agent.runtime.tool.operate.AgentOperateToolAdapter;
|
||||
@@ -153,6 +155,27 @@ public class AgentScopeStatefulRuntimeTest {
|
||||
Assert.assertTrue(runtime.getAgent().getSysPrompt().contains("immediately call observe"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldAppendKnowledgeToolProtocolPromptWhenKnowledgeToolsExist() {
|
||||
AgentInitRequest request = initRequest();
|
||||
AgentKnowledgeSpec knowledgeSpec = new AgentKnowledgeSpec();
|
||||
knowledgeSpec.setKnowledgeId("knowledge-faq");
|
||||
knowledgeSpec.setRuntimeName("faq");
|
||||
request.getAgentDefinition().setKnowledgeSpecs(List.of(knowledgeSpec));
|
||||
request.setKnowledgeRegistrations(List.of(new AgentKnowledgeRegistration(
|
||||
knowledgeSpec, retrievalRequest -> AgentKnowledgeRetrievalResult.of(List.of()))));
|
||||
AgentScopeReActRuntime runtime = fakeRuntime();
|
||||
|
||||
runtime.init(request);
|
||||
|
||||
Assert.assertTrue(runtime.getAgent().getSysPrompt()
|
||||
.startsWith("system\n\nKnowledge tool protocol:"));
|
||||
Assert.assertTrue(runtime.getAgent().getSysPrompt()
|
||||
.contains("call the most relevant knowledge tool first"));
|
||||
Assert.assertTrue(runtime.getAgent().getSysPrompt()
|
||||
.contains("reformulate the query and retry once"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldEmitSideEventWithRuntimeIdentityFromBridge() throws Exception {
|
||||
AgentRuntimeExecutionContext context = new AgentRuntimeExecutionContext();
|
||||
@@ -235,6 +258,26 @@ public class AgentScopeStatefulRuntimeTest {
|
||||
Assert.assertTrue(interceptors.stream().anyMatch(ToolHitlInterceptor.class::isInstance));
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证调用方提供的历史快照会在 Agent 首次构建时进入模型记忆。
|
||||
*/
|
||||
@Test
|
||||
public void shouldAttachInitialConversationHistoryToAgentMemory() {
|
||||
AgentInitRequest request = initRequest();
|
||||
AgentMemorySnapshot snapshot = new AgentMemorySnapshot();
|
||||
snapshot.addMessage(AgentMessage.text(AgentMessageRole.USER, "previous user question"));
|
||||
snapshot.addMessage(AgentMessage.text(AgentMessageRole.ASSISTANT, "previous assistant answer"));
|
||||
request.setMemorySnapshot(snapshot);
|
||||
AgentScopeReActRuntime runtime = fakeRuntime();
|
||||
|
||||
runtime.init(request);
|
||||
|
||||
List<Msg> messages = runtime.getAgent().getMemory().getMessages();
|
||||
Assert.assertEquals(2, messages.size());
|
||||
Assert.assertEquals("previous user question", messages.get(0).getTextContent());
|
||||
Assert.assertEquals("previous assistant answer", messages.get(1).getTextContent());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRegisterOperateToolsIntoToolkit() {
|
||||
AgentInitRequest request = initRequest();
|
||||
@@ -554,6 +597,82 @@ public class AgentScopeStatefulRuntimeTest {
|
||||
Assert.assertFalse(events.toString().contains("sentinel-secret"));
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证知识库工具复用统一工具生命周期,并携带可供调用方投影的知识库分类。
|
||||
*
|
||||
* @throws Exception 等待旁路事件失败时抛出
|
||||
*/
|
||||
@Test
|
||||
public void shouldEmitKnowledgeToolLifecycleFromToolExecutionObserver() throws Exception {
|
||||
AgentRuntimeExecutionContext context = executionContext();
|
||||
Sinks.Many<AgentRuntimeEvent> sink = Sinks.many().replay().all();
|
||||
AgentRuntimeEventBridge bridge = AgentRuntimeEventBridge.fixed(context, sink);
|
||||
AgentToolSpec toolSpec = new AgentToolSpec();
|
||||
toolSpec.setName("retrieve_knowledge_homeinn_faq");
|
||||
toolSpec.setCategory(AgentToolCategory.KNOWLEDGE);
|
||||
toolSpec.getMetadata().put("toolDisplayName", "如家 FAQ");
|
||||
toolSpec.getMetadata().put("knowledgeId", "knowledge-1");
|
||||
toolSpec.getMetadata().put("knowledgeName", "如家 FAQ");
|
||||
toolSpec.getMetadata().put("knowledgeRuntimeName", "homeinn_faq");
|
||||
ToolExecutionObserver observer = new ToolExecutionObserver(bridge, null, List.of(toolSpec));
|
||||
ReActAgent agent = initializedAgent();
|
||||
Toolkit toolkit = agent.getToolkit();
|
||||
ToolUseBlock toolUse = ToolUseBlock.builder()
|
||||
.id("knowledge-call-1")
|
||||
.name("retrieve_knowledge_homeinn_faq")
|
||||
.input(Map.of("query", "几点退房"))
|
||||
.build();
|
||||
ToolResultBlock toolResult = ToolResultBlock.of(
|
||||
"knowledge-call-1",
|
||||
"retrieve_knowledge_homeinn_faq",
|
||||
TextBlock.builder().text("中午十二点退房").build(),
|
||||
Map.of("success", true));
|
||||
|
||||
observer.observe(new PreActingEvent(agent, toolkit, toolUse)).block();
|
||||
observer.observe(new PostActingEvent(agent, toolkit, toolUse, toolResult)).block();
|
||||
List<AgentRuntimeEvent> events = sink.asFlux().take(2).collectList().toFuture()
|
||||
.get(3, TimeUnit.SECONDS);
|
||||
|
||||
Assert.assertEquals(AgentRuntimeEventType.TOOL_CALL, events.get(0).getEventType());
|
||||
Assert.assertEquals("KNOWLEDGE", events.get(0).getPayload().get("toolCategory"));
|
||||
Assert.assertEquals("knowledge-1", events.get(0).getPayload().get("knowledgeId"));
|
||||
Assert.assertEquals("RUNNING", events.get(0).getPayload().get("status"));
|
||||
Assert.assertEquals(AgentRuntimeEventType.TOOL_RESULT, events.get(1).getEventType());
|
||||
Assert.assertEquals("KNOWLEDGE", events.get(1).getPayload().get("toolCategory"));
|
||||
Assert.assertEquals("SUCCESS", events.get(1).getPayload().get("status"));
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证 AgentScope 的文本错误结果会投影为失败工具状态。
|
||||
*
|
||||
* @throws Exception 等待旁路事件失败时抛出
|
||||
*/
|
||||
@Test
|
||||
public void shouldMarkAgentScopeKnowledgeToolErrorAsFailed() throws Exception {
|
||||
AgentRuntimeExecutionContext context = executionContext();
|
||||
Sinks.Many<AgentRuntimeEvent> sink = Sinks.many().replay().all();
|
||||
AgentRuntimeEventBridge bridge = AgentRuntimeEventBridge.fixed(context, sink);
|
||||
AgentToolSpec toolSpec = new AgentToolSpec();
|
||||
toolSpec.setName("retrieve_knowledge_homeinn_faq");
|
||||
toolSpec.setCategory(AgentToolCategory.KNOWLEDGE);
|
||||
ToolExecutionObserver observer = new ToolExecutionObserver(bridge, null, List.of(toolSpec));
|
||||
ReActAgent agent = initializedAgent();
|
||||
ToolUseBlock toolUse = ToolUseBlock.builder()
|
||||
.id("knowledge-call-error")
|
||||
.name("retrieve_knowledge_homeinn_faq")
|
||||
.input(Map.of("query", "异常查询"))
|
||||
.build();
|
||||
ToolResultBlock toolResult = ToolResultBlock.error("retriever unavailable")
|
||||
.withIdAndName("knowledge-call-error", "retrieve_knowledge_homeinn_faq");
|
||||
|
||||
observer.observe(new PostActingEvent(agent, agent.getToolkit(), toolUse, toolResult)).block();
|
||||
AgentRuntimeEvent event = sink.asFlux().next().toFuture().get(3, TimeUnit.SECONDS);
|
||||
|
||||
Assert.assertEquals(AgentRuntimeEventType.TOOL_RESULT, event.getEventType());
|
||||
Assert.assertEquals("FAILED", event.getPayload().get("status"));
|
||||
Assert.assertEquals(Boolean.FALSE, event.getPayload().get("success"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldEmitSkillLifecycleEventsFromSkillExecutionObserver() throws Exception {
|
||||
AgentRuntimeExecutionContext context = executionContext();
|
||||
@@ -838,6 +957,59 @@ public class AgentScopeStatefulRuntimeTest {
|
||||
&& "partial answer".equals(message.getTextContent())));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldPersistUserMessageWhenModelFails() {
|
||||
InMemoryAgentSessionStore sessionStore = new InMemoryAgentSessionStore();
|
||||
AgentInitRequest request = initRequest();
|
||||
request.setSessionStore(sessionStore);
|
||||
AgentScopeReActRuntime runtime = runtimeWithError(new IllegalStateException("model unavailable"));
|
||||
runtime.init(request);
|
||||
|
||||
List<AgentRuntimeEvent> events = runtime.stream(
|
||||
AgentMessage.text(AgentMessageRole.USER, "remember this request"))
|
||||
.collectList()
|
||||
.block();
|
||||
|
||||
Assert.assertTrue(events.stream().anyMatch(event -> event.getEventType() == AgentRuntimeEventType.FAILED));
|
||||
Assert.assertTrue(sessionStore.exists("session-1"));
|
||||
AgentScopeReActRuntime restoredRuntime = fakeRuntime();
|
||||
restoredRuntime.init(request);
|
||||
Assert.assertTrue(restoredRuntime.getAgent().getMemory().getMessages().stream()
|
||||
.anyMatch(message -> message.getRole() == MsgRole.USER
|
||||
&& "remember this request".equals(message.getTextContent())));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldPersistUserMessageWhenReasoningOnlyStreamIsCancelled() throws Exception {
|
||||
InMemoryAgentSessionStore sessionStore = new InMemoryAgentSessionStore();
|
||||
AgentInitRequest request = initRequest();
|
||||
request.setSessionStore(sessionStore);
|
||||
AgentScopeReActRuntime runtime = runtimeWithModel(List.of(
|
||||
ChatResponse.builder()
|
||||
.id("reasoning-only")
|
||||
.content(List.of(ThinkingBlock.builder().thinking("still thinking").build()))
|
||||
.build()), Duration.ofSeconds(5));
|
||||
runtime.init(request);
|
||||
|
||||
CompletableFuture<AgentRuntimeEvent> firstReasoning = new CompletableFuture<>();
|
||||
reactor.core.Disposable disposable = runtime.stream(
|
||||
AgentMessage.text(AgentMessageRole.USER, "remember reasoning request"))
|
||||
.subscribe(event -> {
|
||||
if (event.getEventType() == AgentRuntimeEventType.REASONING_DELTA) {
|
||||
firstReasoning.complete(event);
|
||||
}
|
||||
}, firstReasoning::completeExceptionally);
|
||||
firstReasoning.get(3, TimeUnit.SECONDS);
|
||||
|
||||
disposable.dispose();
|
||||
awaitCondition(() -> sessionStore.exists("session-1"));
|
||||
AgentScopeReActRuntime restoredRuntime = fakeRuntime();
|
||||
restoredRuntime.init(request);
|
||||
Assert.assertTrue(restoredRuntime.getAgent().getMemory().getMessages().stream()
|
||||
.anyMatch(message -> message.getRole() == MsgRole.USER
|
||||
&& "remember reasoning request".equals(message.getTextContent())));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldNotDuplicateNormalToolEventsFromMainStream() {
|
||||
AgentInitRequest request = initRequest();
|
||||
@@ -1577,9 +1749,12 @@ public class AgentScopeStatefulRuntimeTest {
|
||||
AgentInitRequest request = initRequest();
|
||||
AgentKnowledgeSpec knowledgeSpec = new AgentKnowledgeSpec();
|
||||
knowledgeSpec.setKnowledgeId("knowledge-1");
|
||||
knowledgeSpec.setRuntimeName("faq");
|
||||
knowledgeSpec.setName("知识库");
|
||||
request.getAgentDefinition().setKnowledgeSpecs(List.of(knowledgeSpec));
|
||||
request.setKnowledgeRetrievers(Map.of("knowledge-1", retrievalRequest -> {
|
||||
AtomicBoolean knowledgeInvoked = new AtomicBoolean(false);
|
||||
request.setKnowledgeRegistrations(List.of(new AgentKnowledgeRegistration(knowledgeSpec, retrievalRequest -> {
|
||||
knowledgeInvoked.set(true);
|
||||
AgentKnowledgeDocument document = new AgentKnowledgeDocument();
|
||||
document.setDocumentId("doc-1");
|
||||
document.setDocumentName("说明文档");
|
||||
@@ -1587,9 +1762,27 @@ public class AgentScopeStatefulRuntimeTest {
|
||||
document.setContent("fake answer");
|
||||
document.setScore(0.9D);
|
||||
return AgentKnowledgeRetrievalResult.of(List.of(document));
|
||||
}));
|
||||
AgentScopeReActRuntime runtime = fakeRuntime();
|
||||
})));
|
||||
AgentScopeReActRuntime runtime = runtimeWithModel(List.of(
|
||||
ChatResponse.builder()
|
||||
.id("knowledge-tool-call")
|
||||
.content(List.of(ToolUseBlock.builder()
|
||||
.id("call-search")
|
||||
.name("retrieve_knowledge_faq")
|
||||
.input(Map.of("query", "fake answer"))
|
||||
.content("{\"query\":\"fake answer\"}")
|
||||
.build()))
|
||||
.finishReason("tool_calls")
|
||||
.build(),
|
||||
ChatResponse.builder()
|
||||
.id("knowledge-final")
|
||||
.content(List.of(TextBlock.builder().text("fake answer").build()))
|
||||
.finishReason("stop")
|
||||
.build()));
|
||||
runtime.init(request);
|
||||
Assert.assertNotNull(runtime.getAgent().getToolkit().getTool("retrieve_knowledge_faq"));
|
||||
Assert.assertTrue(runtime.getAgent().getToolkit().getToolSchemas().stream()
|
||||
.anyMatch(schema -> "retrieve_knowledge_faq".equals(schema.getName())));
|
||||
|
||||
List<AgentRuntimeEvent> events = runtime.stream(AgentMessage.text(AgentMessageRole.USER, "query knowledge"))
|
||||
.collectList()
|
||||
@@ -1600,19 +1793,44 @@ public class AgentScopeStatefulRuntimeTest {
|
||||
.findFirst()
|
||||
.orElseThrow();
|
||||
Assert.assertNotNull(completed.getMessage());
|
||||
if (completed.getMessage().getKnowledgeReferences().isEmpty()) {
|
||||
/*
|
||||
* 当前 runtime 将知识库注册为 AgentScope AGENTIC RAG,模型需要主动调用
|
||||
* retrieve_knowledge 才会产生 KNOWLEDGE_RETRIEVAL 旁路事件。fake model
|
||||
* 不会调用该工具时,不应强行猜引用。
|
||||
*/
|
||||
Assert.assertFalse(events.stream().anyMatch(event ->
|
||||
event.getEventType() == AgentRuntimeEventType.KNOWLEDGE_RETRIEVAL));
|
||||
return;
|
||||
}
|
||||
Assert.assertTrue("knowledgeInvoked=" + knowledgeInvoked.get() + ", events="
|
||||
+ events.stream().map(AgentRuntimeEvent::getEventType).toList(),
|
||||
events.stream().anyMatch(event ->
|
||||
event.getEventType() == AgentRuntimeEventType.KNOWLEDGE_RETRIEVAL));
|
||||
Assert.assertTrue(events.stream().anyMatch(event ->
|
||||
event.getEventType() == AgentRuntimeEventType.TOOL_CALL
|
||||
&& "KNOWLEDGE".equals(event.getPayload().get("toolCategory"))));
|
||||
Assert.assertTrue(events.stream().anyMatch(event ->
|
||||
event.getEventType() == AgentRuntimeEventType.TOOL_RESULT
|
||||
&& "KNOWLEDGE".equals(event.getPayload().get("toolCategory"))));
|
||||
Assert.assertEquals("chunk-1", completed.getMessage().getKnowledgeReferences().get(0).getChunkId());
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证知识库生成工具与普通工具同名时拒绝初始化,避免 Toolkit 静默覆盖。
|
||||
*/
|
||||
@Test(expected = AgentRuntimeException.class)
|
||||
public void shouldRejectKnowledgeToolNameConflictWithRegularTool() {
|
||||
AgentInitRequest request = initRequest();
|
||||
AgentKnowledgeSpec knowledgeSpec = new AgentKnowledgeSpec();
|
||||
knowledgeSpec.setKnowledgeId("knowledge-1");
|
||||
knowledgeSpec.setRuntimeName("faq");
|
||||
knowledgeSpec.setName("知识库");
|
||||
request.getAgentDefinition().setKnowledgeSpecs(List.of(knowledgeSpec));
|
||||
request.setKnowledgeRegistrations(List.of(new AgentKnowledgeRegistration(
|
||||
knowledgeSpec,
|
||||
retrievalRequest -> AgentKnowledgeRetrievalResult.of(List.of()))));
|
||||
AgentToolSpec regularTool = new AgentToolSpec();
|
||||
regularTool.setName("retrieve_knowledge_faq");
|
||||
regularTool.setDescription("conflicting tool");
|
||||
request.getAgentDefinition().setToolSpecs(List.of(regularTool));
|
||||
request.setToolInvokers(Map.of(
|
||||
regularTool.getName(),
|
||||
(arguments, context) -> AgentToolResult.success("done")));
|
||||
|
||||
fakeRuntime().init(request);
|
||||
}
|
||||
|
||||
private AgentScopeReActRuntime fakeRuntime() {
|
||||
return new AgentScopeReActRuntime(new FakeAgentScopeModelFactory(), new AgentScopeToolAdapter(),
|
||||
new AgentScopeKnowledgeAdapter(), new AgentScopeMemoryAdapter(), new AgentScopeSkillAdapter(),
|
||||
@@ -1636,6 +1854,37 @@ public class AgentScopeStatefulRuntimeTest {
|
||||
new AgentScopeMessageAdapter());
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建模型调用直接失败的运行时。
|
||||
*
|
||||
* @param error 模型异常
|
||||
* @return 测试运行时
|
||||
*/
|
||||
private AgentScopeReActRuntime runtimeWithError(Throwable error) {
|
||||
AgentScopeModelFactory modelFactory = new AgentScopeModelFactory() {
|
||||
@Override
|
||||
public Model create(AgentModelSpec modelSpec,
|
||||
com.easyagents.agent.runtime.model.AgentGenerationOptions generationOptions) {
|
||||
return new Model() {
|
||||
@Override
|
||||
public Flux<ChatResponse> stream(List<Msg> messages,
|
||||
List<ToolSchema> toolSchemas,
|
||||
GenerateOptions options) {
|
||||
return Flux.error(error);
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getModelName() {
|
||||
return modelSpec == null ? "fake-model" : modelSpec.getModelName();
|
||||
}
|
||||
};
|
||||
}
|
||||
};
|
||||
return new AgentScopeReActRuntime(modelFactory, new AgentScopeToolAdapter(),
|
||||
new AgentScopeKnowledgeAdapter(), new AgentScopeMemoryAdapter(), new AgentScopeSkillAdapter(),
|
||||
new AgentScopeMessageAdapter());
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建每次模型调用仅返回下一条预设响应的运行时。
|
||||
*
|
||||
|
||||
@@ -0,0 +1,107 @@
|
||||
package com.easyagents.agent.runtime.agentscope;
|
||||
|
||||
import com.easyagents.agent.runtime.AgentDefinition;
|
||||
import com.easyagents.agent.runtime.AgentRuntimeExecutionContext;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolInvoker;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolResult;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolSpec;
|
||||
import io.agentscope.core.message.ToolResultBlock;
|
||||
import io.agentscope.core.message.ToolUseBlock;
|
||||
import io.agentscope.core.tool.AgentTool;
|
||||
import io.agentscope.core.tool.ToolCallParam;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.lang.reflect.Proxy;
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
|
||||
/**
|
||||
* {@link AgentScopeToolAdapter} 回归测试。
|
||||
*/
|
||||
public class AgentScopeToolAdapterTest {
|
||||
|
||||
/**
|
||||
* 验证 Reactor 工作线程执行工具时使用调用器所属类加载器,并在结束后恢复原上下文。
|
||||
*/
|
||||
@Test
|
||||
public void shouldUseInvokerClassLoaderAndRestoreWorkerContext() {
|
||||
ClassLoader parentClassLoader = AgentToolInvoker.class.getClassLoader();
|
||||
ClassLoader invocationClassLoader = new ClassLoader(parentClassLoader) {
|
||||
};
|
||||
AtomicReference<ClassLoader> observedClassLoader = new AtomicReference<>();
|
||||
AgentToolInvoker invoker = (AgentToolInvoker) Proxy.newProxyInstance(
|
||||
invocationClassLoader,
|
||||
new Class<?>[]{AgentToolInvoker.class},
|
||||
(proxy, method, arguments) -> {
|
||||
observedClassLoader.set(Thread.currentThread().getContextClassLoader());
|
||||
return AgentToolResult.success("ok");
|
||||
});
|
||||
AgentTool tool = new AgentScopeToolAdapter().adapt(toolSpec(), invoker, executionContext());
|
||||
Thread currentThread = Thread.currentThread();
|
||||
ClassLoader originalClassLoader = currentThread.getContextClassLoader();
|
||||
ClassLoader workerClassLoader = new ClassLoader(originalClassLoader) {
|
||||
};
|
||||
|
||||
try {
|
||||
currentThread.setContextClassLoader(workerClassLoader);
|
||||
ToolResultBlock result = tool.callAsync(toolCall()).block();
|
||||
|
||||
Assert.assertNotNull(result);
|
||||
Assert.assertSame(invocationClassLoader, observedClassLoader.get());
|
||||
Assert.assertSame(workerClassLoader, currentThread.getContextClassLoader());
|
||||
} finally {
|
||||
currentThread.setContextClassLoader(originalClassLoader);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建测试工具声明。
|
||||
*
|
||||
* @return 工具声明
|
||||
*/
|
||||
private AgentToolSpec toolSpec() {
|
||||
AgentToolSpec toolSpec = new AgentToolSpec();
|
||||
toolSpec.setName("test_tool");
|
||||
toolSpec.setDescription("Test tool");
|
||||
toolSpec.setParametersSchema(Map.of(
|
||||
"type", "object",
|
||||
"properties", Map.of(),
|
||||
"additionalProperties", false));
|
||||
return toolSpec;
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建测试运行上下文。
|
||||
*
|
||||
* @return 运行上下文
|
||||
*/
|
||||
private AgentRuntimeExecutionContext executionContext() {
|
||||
AgentDefinition definition = new AgentDefinition();
|
||||
definition.setAgentId("agent-1");
|
||||
AgentRuntimeExecutionContext context = new AgentRuntimeExecutionContext();
|
||||
context.setAgentDefinition(definition);
|
||||
context.setRequestId("request-1");
|
||||
context.setTraceId("trace-1");
|
||||
context.setSessionId("session-1");
|
||||
return context;
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建测试工具调用参数。
|
||||
*
|
||||
* @return 工具调用参数
|
||||
*/
|
||||
private ToolCallParam toolCall() {
|
||||
ToolUseBlock toolUseBlock = ToolUseBlock.builder()
|
||||
.id("call-1")
|
||||
.name("test_tool")
|
||||
.input(Map.of())
|
||||
.content("{}")
|
||||
.build();
|
||||
return ToolCallParam.builder()
|
||||
.toolUseBlock(toolUseBlock)
|
||||
.input(Map.of())
|
||||
.build();
|
||||
}
|
||||
}
|
||||
@@ -72,6 +72,32 @@ public class HeuristicKnowledgeCitationMatcherTest {
|
||||
Assert.assertTrue(references.isEmpty());
|
||||
}
|
||||
|
||||
/**
|
||||
* 长篇多主题汇总回答应按独立条目匹配引用,避免整篇答案稀释局部证据。
|
||||
*/
|
||||
@Test
|
||||
public void shouldMatchReferencesFromLongMultiTopicAnswer() {
|
||||
AgentKnowledgeReference checkIn = reference("faq-check-in",
|
||||
"问题:最早入住酒店时间说明 答案:最早入住时间为入住日当天下午14:00,提前到店按房间状况安排。");
|
||||
AgentKnowledgeReference luggage = reference("faq-luggage",
|
||||
"问题:酒店寄存行李服务说明 答案:离店客人通常可免费寄存2天,第3天起收费。");
|
||||
AgentKnowledgeReference unrelated = reference("faq-unrelated",
|
||||
"问题:会员卡如何补办 答案:请携带本人证件到指定服务网点申请补办。");
|
||||
|
||||
String answer = "根据知识库整理,主要主题如下:\n"
|
||||
+ "一、预订与支付:官方渠道包括APP、微信小程序和客服电话。\n"
|
||||
+ "二、入住与退房:最早入住时间为下午14:00,提前到店按房态安排。\n"
|
||||
+ "三、酒店服务:离店后行李通常可免费寄存2天,第3天起收费。\n"
|
||||
+ "四、会员商城:可使用彩虹如愿豆兑换商品。";
|
||||
|
||||
List<AgentKnowledgeReference> references = matcher.match(
|
||||
answer, List.of(unrelated, checkIn, luggage));
|
||||
|
||||
Assert.assertEquals(2, references.size());
|
||||
Assert.assertTrue(references.stream().anyMatch(reference -> "faq-check-in".equals(reference.getDocumentId())));
|
||||
Assert.assertTrue(references.stream().anyMatch(reference -> "faq-luggage".equals(reference.getDocumentId())));
|
||||
}
|
||||
|
||||
/**
|
||||
* 空答案或空候选不应返回引用。
|
||||
*/
|
||||
|
||||
@@ -0,0 +1,149 @@
|
||||
package com.easyagents.agent.runtime.prompt;
|
||||
|
||||
import com.easyagents.agent.runtime.AgentDefinition;
|
||||
import com.easyagents.agent.runtime.knowledge.AgentKnowledgeSpec;
|
||||
import com.easyagents.agent.runtime.tool.AgentToolSpec;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
/**
|
||||
* 测试运行时系统提示词组合器。
|
||||
*/
|
||||
public class SystemPromptComposerTest {
|
||||
|
||||
@Test
|
||||
public void shouldKeepUserPromptWhenNoRuntimeProtocolIsRequired() {
|
||||
AgentDefinition definition = definition(" user prompt ");
|
||||
|
||||
Assert.assertEquals(" user prompt ", SystemPromptComposer.compose(definition));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldReturnKnowledgeProtocolWhenUserPromptIsBlank() {
|
||||
AgentDefinition definition = definition(" ");
|
||||
definition.setKnowledgeSpecs(List.of(knowledge("faq")));
|
||||
|
||||
String prompt = SystemPromptComposer.compose(definition);
|
||||
|
||||
Assert.assertTrue(prompt.startsWith("Knowledge tool protocol:"));
|
||||
Assert.assertFalse(prompt.startsWith("\n"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldAppendKnowledgeProtocolAfterUserPrompt() {
|
||||
AgentDefinition definition = definition("user prompt");
|
||||
definition.setKnowledgeSpecs(List.of(knowledge("faq")));
|
||||
|
||||
String prompt = SystemPromptComposer.compose(definition);
|
||||
|
||||
Assert.assertTrue(prompt.startsWith("user prompt\n\nKnowledge tool protocol:"));
|
||||
Assert.assertTrue(prompt.contains("prefer making one retrieval call"));
|
||||
Assert.assertTrue(prompt.contains("reformulate the query and retry once"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldAppendKnowledgeProtocolOnlyOnceForMultipleKnowledgeTools() {
|
||||
AgentDefinition definition = definition("user prompt");
|
||||
definition.setKnowledgeSpecs(List.of(knowledge("faq"), knowledge("policy")));
|
||||
|
||||
String prompt = SystemPromptComposer.compose(definition);
|
||||
|
||||
Assert.assertEquals(1, occurrences(prompt, "Knowledge tool protocol:"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldSkipKnowledgeProtocolWhenToolCallingIsDisabled() {
|
||||
AgentDefinition definition = definition("user prompt");
|
||||
definition.setKnowledgeSpecs(List.of(knowledge("faq")));
|
||||
definition.getExecutionOptions().setToolCallingEnabled(false);
|
||||
|
||||
Assert.assertEquals("user prompt", SystemPromptComposer.compose(definition));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldPreserveAsyncToolProtocolBehavior() {
|
||||
AgentDefinition definition = definition(null);
|
||||
definition.setToolSpecs(List.of(asyncTool()));
|
||||
|
||||
String prompt = SystemPromptComposer.compose(definition);
|
||||
|
||||
Assert.assertTrue(prompt.startsWith("Async tool protocol:"));
|
||||
Assert.assertTrue(prompt.contains("These are internal execution phases."));
|
||||
Assert.assertTrue(prompt.contains("call its submit sub-tool first with the user-provided arguments by default"));
|
||||
Assert.assertTrue(prompt.contains("immediately call observe"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldComposeUserKnowledgeAndAsyncProtocolsInStableOrder() {
|
||||
AgentDefinition definition = definition("user prompt");
|
||||
definition.setKnowledgeSpecs(List.of(knowledge("faq")));
|
||||
definition.setToolSpecs(List.of(asyncTool()));
|
||||
|
||||
String prompt = SystemPromptComposer.compose(definition);
|
||||
int userIndex = prompt.indexOf("user prompt");
|
||||
int knowledgeIndex = prompt.indexOf("Knowledge tool protocol:");
|
||||
int asyncIndex = prompt.indexOf("Async tool protocol:");
|
||||
|
||||
Assert.assertTrue(userIndex >= 0);
|
||||
Assert.assertTrue(knowledgeIndex > userIndex);
|
||||
Assert.assertTrue(asyncIndex > knowledgeIndex);
|
||||
Assert.assertEquals(1, occurrences(prompt, "Knowledge tool protocol:"));
|
||||
Assert.assertEquals(1, occurrences(prompt, "Async tool protocol:"));
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建测试智能体定义。
|
||||
*
|
||||
* @param systemPrompt 用户系统提示词
|
||||
* @return 智能体定义
|
||||
*/
|
||||
private AgentDefinition definition(String systemPrompt) {
|
||||
AgentDefinition definition = new AgentDefinition();
|
||||
definition.setSystemPrompt(systemPrompt);
|
||||
return definition;
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建测试知识库定义。
|
||||
*
|
||||
* @param runtimeName 知识库运行名
|
||||
* @return 知识库定义
|
||||
*/
|
||||
private AgentKnowledgeSpec knowledge(String runtimeName) {
|
||||
AgentKnowledgeSpec spec = new AgentKnowledgeSpec();
|
||||
spec.setKnowledgeId("knowledge-" + runtimeName);
|
||||
spec.setRuntimeName(runtimeName);
|
||||
return spec;
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建测试异步工具定义。
|
||||
*
|
||||
* @return 异步工具定义
|
||||
*/
|
||||
private AgentToolSpec asyncTool() {
|
||||
AgentToolSpec toolSpec = new AgentToolSpec();
|
||||
toolSpec.setName("demo_submit");
|
||||
toolSpec.getMetadata().put("asyncTool", true);
|
||||
return toolSpec;
|
||||
}
|
||||
|
||||
/**
|
||||
* 统计文本片段出现次数。
|
||||
*
|
||||
* @param value 待检查文本
|
||||
* @param fragment 目标片段
|
||||
* @return 出现次数
|
||||
*/
|
||||
private int occurrences(String value, String fragment) {
|
||||
int count = 0;
|
||||
int offset = 0;
|
||||
while ((offset = value.indexOf(fragment, offset)) >= 0) {
|
||||
count++;
|
||||
offset += fragment.length();
|
||||
}
|
||||
return count;
|
||||
}
|
||||
}
|
||||
@@ -43,6 +43,11 @@ public class StoreOptions extends Metadata {
|
||||
public void setEmbeddingOptions(EmbeddingOptions embeddingOptions) {
|
||||
throw new IllegalStateException("Can not set embeddingOptions to the default instance.");
|
||||
}
|
||||
|
||||
@Override
|
||||
public void setTimeoutMillis(Long timeoutMillis) {
|
||||
throw new IllegalStateException("Can not set timeoutMillis to the default instance.");
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
@@ -65,6 +70,11 @@ public class StoreOptions extends Metadata {
|
||||
*/
|
||||
private EmbeddingOptions embeddingOptions = EmbeddingOptions.DEFAULT;
|
||||
|
||||
/**
|
||||
* Optional upper bound for one store operation.
|
||||
*/
|
||||
private Long timeoutMillis;
|
||||
|
||||
|
||||
public String getCollectionName() {
|
||||
return collectionName;
|
||||
@@ -111,6 +121,17 @@ public class StoreOptions extends Metadata {
|
||||
this.embeddingOptions = embeddingOptions;
|
||||
}
|
||||
|
||||
public Long getTimeoutMillis() {
|
||||
return timeoutMillis;
|
||||
}
|
||||
|
||||
public void setTimeoutMillis(Long timeoutMillis) {
|
||||
if (timeoutMillis != null && timeoutMillis <= 0L) {
|
||||
throw new IllegalArgumentException("timeoutMillis must be greater than zero");
|
||||
}
|
||||
this.timeoutMillis = timeoutMillis;
|
||||
}
|
||||
|
||||
|
||||
public static StoreOptions ofCollectionName(String collectionName) {
|
||||
StoreOptions storeOptions = new StoreOptions();
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
/*
|
||||
* Copyright (c) 2023-2026, Easy-Agents (fuhai999@gmail.com).
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
*/
|
||||
package com.easyagents.core.store;
|
||||
|
||||
/**
|
||||
* Indicates that a store operation exhausted its caller-provided time budget.
|
||||
*/
|
||||
public class StoreTimeoutException extends RuntimeException {
|
||||
|
||||
public StoreTimeoutException(String message) {
|
||||
super(message);
|
||||
}
|
||||
|
||||
public StoreTimeoutException(String message, Throwable cause) {
|
||||
super(message, cause);
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package com.easyagents.document.core.async;
|
||||
|
||||
import com.easyagents.core.util.StringUtil;
|
||||
import com.easyagents.document.core.exception.DocumentParseException;
|
||||
import com.easyagents.document.core.exception.DocumentAsyncTaskNotFoundException;
|
||||
import com.easyagents.document.core.entity.ParseResponse;
|
||||
import com.easyagents.document.core.entity.ParseTaskInfo;
|
||||
import com.easyagents.document.core.entity.ParseTaskStatus;
|
||||
@@ -135,7 +136,7 @@ public class DocumentAsyncTaskManager {
|
||||
}
|
||||
DocumentAsyncTaskRecord record = repository.find(taskId);
|
||||
if (record == null) {
|
||||
throw new DocumentParseException("Document async task not found: " + taskId);
|
||||
throw new DocumentAsyncTaskNotFoundException(taskId);
|
||||
}
|
||||
return record;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
package com.easyagents.document.core.exception;
|
||||
|
||||
/**
|
||||
* 进程内异步文档任务不存在。
|
||||
*
|
||||
* <p>本地 Office 解析任务允许使用内存仓库;进程重启后,调用方可以
|
||||
* 通过该异常识别执行实例已经丢失,并从持久化业务任务重新提交。</p>
|
||||
*
|
||||
* @author Codex
|
||||
* @since 2026-09-02
|
||||
*/
|
||||
public class DocumentAsyncTaskNotFoundException extends DocumentParseException {
|
||||
|
||||
private final String taskId;
|
||||
|
||||
public DocumentAsyncTaskNotFoundException(String taskId) {
|
||||
super("Document async task not found: " + taskId);
|
||||
this.taskId = taskId;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取已丢失的任务 ID。
|
||||
*
|
||||
* @return 任务 ID
|
||||
*/
|
||||
public String getTaskId() {
|
||||
return taskId;
|
||||
}
|
||||
}
|
||||
@@ -4,6 +4,7 @@ import com.easyagents.document.core.entity.ParseResponse;
|
||||
import com.easyagents.document.core.entity.ParseResult;
|
||||
import com.easyagents.document.core.entity.ParseTaskInfo;
|
||||
import com.easyagents.document.core.entity.ParseTaskStatus;
|
||||
import com.easyagents.document.core.exception.DocumentAsyncTaskNotFoundException;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
@@ -18,6 +19,21 @@ import java.util.concurrent.Executor;
|
||||
*/
|
||||
public class DocumentAsyncTaskManagerTest {
|
||||
|
||||
@Test
|
||||
public void shouldExposeMissingInMemoryTaskAsRecoverableSignal() {
|
||||
DocumentAsyncTaskManager manager = new DocumentAsyncTaskManager(
|
||||
new InMemoryDocumentAsyncTaskRepository(),
|
||||
Runnable::run
|
||||
);
|
||||
|
||||
try {
|
||||
manager.queryTaskInfo("lost-task");
|
||||
Assert.fail("expected DocumentAsyncTaskNotFoundException");
|
||||
} catch (DocumentAsyncTaskNotFoundException error) {
|
||||
Assert.assertEquals("lost-task", error.getTaskId());
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldTrackTaskLifecycleAndResult() {
|
||||
Executor directExecutor = new Executor() {
|
||||
|
||||
@@ -14,6 +14,7 @@ import com.easyagents.document.core.entity.ParseRequest;
|
||||
import com.easyagents.document.core.entity.ParseResponse;
|
||||
import com.easyagents.document.core.entity.ParseResult;
|
||||
import com.easyagents.document.core.entity.XlsxParseRequest;
|
||||
import com.easyagents.document.core.exception.DocumentParseException;
|
||||
import com.easyagents.document.core.support.AbstractAsyncDocumentParseService;
|
||||
import com.easyagents.document.xlsx.XlsxDocumentProvider;
|
||||
import com.easyagents.document.xlsx.model.XlsxCellArtifact;
|
||||
@@ -52,6 +53,10 @@ import java.util.concurrent.Executors;
|
||||
public class MineruXlsxDocumentParseService extends AbstractAsyncDocumentParseService<XlsxParseRequest> implements XlsxDocumentProvider {
|
||||
|
||||
public static final String PROVIDER_NAME = "mineru";
|
||||
private static final byte[] OLE2_SIGNATURE = new byte[] {
|
||||
(byte) 0xd0, (byte) 0xcf, 0x11, (byte) 0xe0,
|
||||
(byte) 0xa1, (byte) 0xb1, 0x1a, (byte) 0xe1
|
||||
};
|
||||
|
||||
private final MineruProperties properties;
|
||||
private final MineruClient client;
|
||||
@@ -145,6 +150,9 @@ public class MineruXlsxDocumentParseService extends AbstractAsyncDocumentParseSe
|
||||
|
||||
@Override
|
||||
protected ParseResponse doParse(XlsxParseRequest request, DocumentAsyncTaskUpdater updater) {
|
||||
for (ParseFile file : request.getFiles()) {
|
||||
validateXlsxContent(file);
|
||||
}
|
||||
ParseResponse response = new ParseResponse();
|
||||
List<ParseResult> results = new ArrayList<ParseResult>();
|
||||
String backend = null;
|
||||
@@ -211,6 +219,50 @@ public class MineruXlsxDocumentParseService extends AbstractAsyncDocumentParseSe
|
||||
return aggregate;
|
||||
}
|
||||
|
||||
/**
|
||||
* 在 POI 打开工作簿前校验 XLSX 容器签名,避免向用户暴露底层格式异常。
|
||||
*
|
||||
* @param file 待解析文件
|
||||
*/
|
||||
private void validateXlsxContent(ParseFile file) {
|
||||
byte[] content = file == null ? null : file.getContent();
|
||||
if (hasZipSignature(content)) {
|
||||
return;
|
||||
}
|
||||
String fileName = file == null || !StringUtil.hasText(file.getFileName())
|
||||
? "当前文件"
|
||||
: "文件“" + file.getFileName() + "”";
|
||||
String reason = startsWith(content, OLE2_SIGNATURE)
|
||||
? "可能是旧版 XLS 或已加密文件"
|
||||
: "文件内容与 .xlsx 扩展名不一致或文件已损坏";
|
||||
throw new DocumentParseException(
|
||||
fileName + "不是标准 XLSX," + reason
|
||||
+ "。请解除保护后用 Excel/WPS 另存为 XLSX(修改文件后缀无效)"
|
||||
);
|
||||
}
|
||||
|
||||
private boolean hasZipSignature(byte[] content) {
|
||||
return content != null
|
||||
&& content.length >= 4
|
||||
&& content[0] == 'P'
|
||||
&& content[1] == 'K'
|
||||
&& ((content[2] == 3 && content[3] == 4)
|
||||
|| (content[2] == 5 && content[3] == 6)
|
||||
|| (content[2] == 7 && content[3] == 8));
|
||||
}
|
||||
|
||||
private boolean startsWith(byte[] content, byte[] signature) {
|
||||
if (content == null || content.length < signature.length) {
|
||||
return false;
|
||||
}
|
||||
for (int index = 0; index < signature.length; index++) {
|
||||
if (content[index] != signature[index]) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
private SheetExtraction extractSheet(XSSFSheet sheet,
|
||||
int sheetIndex,
|
||||
DataFormatter formatter,
|
||||
|
||||
@@ -16,6 +16,7 @@ import com.easyagents.document.core.entity.ParseTaskStatus;
|
||||
import com.easyagents.document.core.entity.XlsxParseRequest;
|
||||
import com.easyagents.document.core.exception.DocumentParseException;
|
||||
import com.easyagents.document.xlsx.model.XlsxParseArtifact;
|
||||
import org.apache.poi.hssf.usermodel.HSSFWorkbook;
|
||||
import org.apache.poi.ss.usermodel.ClientAnchor;
|
||||
import org.apache.poi.xssf.usermodel.XSSFDrawing;
|
||||
import org.apache.poi.xssf.usermodel.XSSFSheet;
|
||||
@@ -119,6 +120,31 @@ public class MineruXlsxDocumentParseServiceTest {
|
||||
Assert.assertEquals("image/jpeg", result.getImages().get(0).getMimeType());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectLegacyXlsContentWithActionableMessage() throws Exception {
|
||||
RecordingClient client = new RecordingClient(defaultProperties());
|
||||
MineruMapper mapper = new MineruMapper(defaultProperties());
|
||||
MineruXlsxDocumentParseService service = new MineruXlsxDocumentParseService(
|
||||
defaultProperties(),
|
||||
client,
|
||||
mapper,
|
||||
new DocumentAsyncTaskManager(new InMemoryDocumentAsyncTaskRepository(), directExecutor())
|
||||
);
|
||||
XlsxParseRequest request = new XlsxParseRequest();
|
||||
request.addFile(ParseFile.of("legacy.xlsx", buildLegacyWorkbookBytes()));
|
||||
|
||||
DocumentParseException error = Assert.assertThrows(
|
||||
DocumentParseException.class,
|
||||
() -> service.parse(request)
|
||||
);
|
||||
|
||||
Assert.assertEquals(
|
||||
"文件“legacy.xlsx”不是标准 XLSX,可能是旧版 XLS 或已加密文件。"
|
||||
+ "请解除保护后用 Excel/WPS 另存为 XLSX(修改文件后缀无效)",
|
||||
error.getMessage()
|
||||
);
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldAppendImageReferenceForImageOnlySheet() throws Exception {
|
||||
RecordingClient client = new RecordingClient(defaultProperties());
|
||||
@@ -255,6 +281,15 @@ public class MineruXlsxDocumentParseServiceTest {
|
||||
return writeWorkbook(workbook);
|
||||
}
|
||||
|
||||
private byte[] buildLegacyWorkbookBytes() throws Exception {
|
||||
try (HSSFWorkbook workbook = new HSSFWorkbook();
|
||||
ByteArrayOutputStream outputStream = new ByteArrayOutputStream()) {
|
||||
workbook.createSheet("Sheet1").createRow(0).createCell(0).setCellValue("旧版表格");
|
||||
workbook.write(outputStream);
|
||||
return outputStream.toByteArray();
|
||||
}
|
||||
}
|
||||
|
||||
private void addPicture(XSSFWorkbook workbook,
|
||||
XSSFSheet sheet,
|
||||
int rowIndex,
|
||||
|
||||
@@ -634,7 +634,8 @@ public class Chain {
|
||||
NodeStateField
|
||||
.EXECUTION_ATTEMPT_KEY);
|
||||
}
|
||||
if (node.getCondition() == null) {
|
||||
if (node.getJoinMode() == NodeJoinMode.ANY
|
||||
&& node.getCondition() == null) {
|
||||
s.recordTrigger(triggerEdgeId);
|
||||
fields.add(NodeStateField.TRIGGER_COUNT);
|
||||
fields.add(NodeStateField.TRIGGER_EDGE_IDS);
|
||||
@@ -795,7 +796,7 @@ public class Chain {
|
||||
|
||||
private boolean shouldSkipNode(Node node, String edgeId) {
|
||||
NodeCondition condition = node.getCondition();
|
||||
if (condition == null) {
|
||||
if (node.getJoinMode() == NodeJoinMode.ANY && condition == null) {
|
||||
return false;
|
||||
}
|
||||
return executeWithLock(stateInstanceId, 10, TimeUnit.SECONDS, () -> {
|
||||
@@ -805,7 +806,11 @@ public class Chain {
|
||||
});
|
||||
|
||||
Map<String, Object> prevResult = Collections.emptyMap();
|
||||
boolean shouldSkipNode = !condition.check(this, newState, prevResult);
|
||||
boolean joinPending = node.getJoinMode() == NodeJoinMode.ALL
|
||||
&& !newState.isUpstreamFullyExecuted();
|
||||
boolean shouldSkipNode = joinPending
|
||||
|| (condition != null
|
||||
&& !condition.check(this, newState, prevResult));
|
||||
if (shouldSkipNode) {
|
||||
updateStateSafely(state -> {
|
||||
return state.addUncheckedNodeId(node.id)
|
||||
@@ -1746,16 +1751,91 @@ public class Chain {
|
||||
stateInstanceId,
|
||||
10L,
|
||||
TimeUnit.SECONDS,
|
||||
() -> {
|
||||
ChainState current =
|
||||
chainStateRepository.load(stateInstanceId);
|
||||
if (current == null
|
||||
|| current.getStatus() != ChainStatus.SUSPEND) {
|
||||
return false;
|
||||
}
|
||||
resumeSuspended(variables);
|
||||
return true;
|
||||
});
|
||||
() -> resumeSuspended(variables));
|
||||
}
|
||||
|
||||
private void validateResumeVariables(
|
||||
ChainState state, Map<String, Object> variables) {
|
||||
List<Parameter> parameters = state.getSuspendForParameters();
|
||||
if (parameters == null || parameters.isEmpty()) {
|
||||
return;
|
||||
}
|
||||
Map<String, Object> submitted = variables == null
|
||||
? Collections.emptyMap()
|
||||
: variables;
|
||||
Set<String> expectedKeys = parameters.stream()
|
||||
.map(Parameter::getName)
|
||||
.filter(Objects::nonNull)
|
||||
.collect(java.util.stream.Collectors.toCollection(
|
||||
LinkedHashSet::new));
|
||||
Set<String> extraKeys = new LinkedHashSet<>(submitted.keySet());
|
||||
extraKeys.removeAll(expectedKeys);
|
||||
if (!extraKeys.isEmpty()) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数包含未声明字段");
|
||||
}
|
||||
for (Parameter parameter : parameters) {
|
||||
String name = parameter.getName();
|
||||
Object value = submitted.get(name);
|
||||
String label = StringUtil.getFirstWithText(
|
||||
parameter.getFormLabel(), name);
|
||||
if (!submitted.containsKey(name) || isBlankResumeValue(value)) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数[" + label + "]不能为空");
|
||||
}
|
||||
validateResumeOption(parameter, value, label);
|
||||
}
|
||||
}
|
||||
|
||||
private void validateResumeOption(
|
||||
Parameter parameter, Object value, String label) {
|
||||
List<ParameterOption> options = parameter.getOptions();
|
||||
if (options == null || options.isEmpty()) {
|
||||
return;
|
||||
}
|
||||
Set<String> allowedValues = options.stream()
|
||||
.map(ParameterOption::getValue)
|
||||
.filter(Objects::nonNull)
|
||||
.collect(java.util.stream.Collectors.toCollection(
|
||||
LinkedHashSet::new));
|
||||
if ("checkbox".equals(parameter.getFormType())) {
|
||||
if (!(value instanceof Collection<?> selected)) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数[" + label + "]必须提交字符串数组");
|
||||
}
|
||||
Set<String> unique = new LinkedHashSet<>();
|
||||
for (Object item : selected) {
|
||||
if (!(item instanceof String selectedValue)
|
||||
|| !allowedValues.contains(selectedValue)) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数[" + label + "]包含未配置选项");
|
||||
}
|
||||
if (!unique.add(selectedValue)) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数[" + label + "]不能重复选择同一选项");
|
||||
}
|
||||
}
|
||||
return;
|
||||
}
|
||||
if (!(value instanceof String selectedValue)) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数[" + label + "]必须提交单个字符串值");
|
||||
}
|
||||
if (!allowedValues.contains(selectedValue)) {
|
||||
throw new ChainResumeException(
|
||||
"确认参数[" + label + "]包含未配置选项");
|
||||
}
|
||||
}
|
||||
|
||||
private boolean isBlankResumeValue(Object value) {
|
||||
if (value == null) {
|
||||
return true;
|
||||
}
|
||||
if (value instanceof String text) {
|
||||
return text.trim().isEmpty();
|
||||
}
|
||||
return value instanceof Collection<?> collection
|
||||
&& collection.isEmpty();
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -1772,36 +1852,51 @@ public class Chain {
|
||||
*
|
||||
* @param variables 恢复时注入的变量
|
||||
*/
|
||||
private void resumeSuspended(Map<String, Object> variables) {
|
||||
ChainState newState = updateStateSafely(state -> {
|
||||
if (variables != null) {
|
||||
state.getMemory().putAll(variables);
|
||||
return EnumSet.of(ChainStateField.MEMORY);
|
||||
} else {
|
||||
private boolean resumeSuspended(Map<String, Object> variables) {
|
||||
AtomicBoolean resumed = new AtomicBoolean(false);
|
||||
AtomicReference<Set<String>> suspendedNodeIds =
|
||||
new AtomicReference<>(Collections.emptySet());
|
||||
updateStateSafely(state -> {
|
||||
resumed.set(false);
|
||||
suspendedNodeIds.set(Collections.emptySet());
|
||||
if (state.getStatus() != ChainStatus.SUSPEND) {
|
||||
return null;
|
||||
}
|
||||
validateResumeVariables(state, variables);
|
||||
if (state.getSuspendNodeIds() != null) {
|
||||
suspendedNodeIds.set(
|
||||
new LinkedHashSet<>(state.getSuspendNodeIds()));
|
||||
}
|
||||
EnumSet<ChainStateField> updatedFields = EnumSet.of(
|
||||
ChainStateField.STATUS,
|
||||
ChainStateField.SUSPEND_NODE_IDS,
|
||||
ChainStateField.SUSPEND_FOR_PARAMETERS);
|
||||
if (variables != null && !variables.isEmpty()) {
|
||||
state.getMemory().putAll(variables);
|
||||
updatedFields.add(ChainStateField.MEMORY);
|
||||
}
|
||||
state.setStatus(ChainStatus.RUNNING);
|
||||
state.setSuspendNodeIds(null);
|
||||
state.setSuspendForParameters(null);
|
||||
resumed.set(true);
|
||||
return updatedFields;
|
||||
});
|
||||
if (!resumed.get()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
notifyEvent(new ChainResumeEvent(this, variables));
|
||||
setStatusAndNotifyEvent(ChainStatus.RUNNING);
|
||||
notifyEvent(new ChainStatusChangeEvent(
|
||||
this, ChainStatus.RUNNING, ChainStatus.SUSPEND));
|
||||
|
||||
Set<String> suspendNodeIds = newState.getSuspendNodeIds();
|
||||
if (suspendNodeIds != null && !suspendNodeIds.isEmpty()) {
|
||||
// 移除 suspend 状态,方便二次 suspend 时,不带有旧数据
|
||||
updateStateSafely(state -> {
|
||||
state.setSuspendNodeIds(null);
|
||||
state.setSuspendForParameters(null);
|
||||
return EnumSet.of(ChainStateField.SUSPEND_NODE_IDS, ChainStateField.SUSPEND_FOR_PARAMETERS);
|
||||
});
|
||||
|
||||
for (String id : suspendNodeIds) {
|
||||
Node node = definition.getNodeById(id);
|
||||
if (node == null) {
|
||||
throw new ChainException("Node not found: " + id);
|
||||
}
|
||||
scheduleNode(node, null, TriggerType.RESUME, 0L);
|
||||
for (String id : suspendedNodeIds.get()) {
|
||||
Node node = definition.getNodeById(id);
|
||||
if (node == null) {
|
||||
throw new ChainException("Node not found: " + id);
|
||||
}
|
||||
scheduleNode(node, null, TriggerType.RESUME, 0L);
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
public void resume() {
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
/**
|
||||
* 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.chain;
|
||||
|
||||
/**
|
||||
* 工作流挂起参数不满足当前恢复请求时抛出的异常。
|
||||
*/
|
||||
public class ChainResumeException extends ChainException {
|
||||
|
||||
public ChainResumeException(String message) {
|
||||
super(message);
|
||||
}
|
||||
}
|
||||
@@ -42,6 +42,7 @@ public abstract class Node implements Serializable {
|
||||
|
||||
protected NodeCondition condition;
|
||||
protected NodeValidator validator;
|
||||
protected NodeJoinMode joinMode = NodeJoinMode.ANY;
|
||||
|
||||
// 循环执行相关属性
|
||||
protected boolean loopEnable = false; // 是否启用循环执行
|
||||
@@ -70,6 +71,10 @@ public abstract class Node implements Serializable {
|
||||
}
|
||||
|
||||
public void setParentId(String parentId) {
|
||||
if (StringUtil.hasText(parentId) && getJoinMode() == NodeJoinMode.ALL) {
|
||||
throw new IllegalArgumentException(
|
||||
"joinMode 'all' is not supported for loop child nodes");
|
||||
}
|
||||
this.parentId = parentId;
|
||||
}
|
||||
|
||||
@@ -121,6 +126,31 @@ public abstract class Node implements Serializable {
|
||||
this.validator = validator;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取节点的直接入边汇聚模式。
|
||||
*
|
||||
* <p>旧序列化对象缺少该字段时返回 {@link NodeJoinMode#ANY}。</p>
|
||||
*
|
||||
* @return 汇聚模式
|
||||
*/
|
||||
public NodeJoinMode getJoinMode() {
|
||||
return joinMode == null ? NodeJoinMode.ANY : joinMode;
|
||||
}
|
||||
|
||||
/**
|
||||
* 设置节点的直接入边汇聚模式。
|
||||
*
|
||||
* @param joinMode 汇聚模式
|
||||
*/
|
||||
public void setJoinMode(NodeJoinMode joinMode) {
|
||||
NodeJoinMode resolved = joinMode == null ? NodeJoinMode.ANY : joinMode;
|
||||
if (resolved == NodeJoinMode.ALL && StringUtil.hasText(parentId)) {
|
||||
throw new IllegalArgumentException(
|
||||
"joinMode 'all' is not supported for loop child nodes");
|
||||
}
|
||||
this.joinMode = resolved;
|
||||
}
|
||||
|
||||
// protected void addOutwardEdge(Edge edge) {
|
||||
// if (this.outwardEdges == null) {
|
||||
// this.outwardEdges = new ArrayList<>();
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
/**
|
||||
* Copyright (c) 2025-2026, Michael Yang 杨福海 (fuhai999@gmail.com).
|
||||
* <p>
|
||||
* Licensed under the GNU Lesser General Public License (LGPL) ,Version 3.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
* <p>
|
||||
* http://www.gnu.org/licenses/lgpl-3.0.txt
|
||||
* <p>
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
package com.easyagents.flow.core.chain;
|
||||
|
||||
import java.util.Locale;
|
||||
|
||||
/**
|
||||
* 多入边节点的触发汇聚模式。
|
||||
*/
|
||||
public enum NodeJoinMode {
|
||||
|
||||
/** 任意一条直接入边到达即可执行。 */
|
||||
ANY("any"),
|
||||
/** 全部直接入边到达后才执行。 */
|
||||
ALL("all");
|
||||
|
||||
private final String value;
|
||||
|
||||
NodeJoinMode(String value) {
|
||||
this.value = value;
|
||||
}
|
||||
|
||||
public String getValue() {
|
||||
return value;
|
||||
}
|
||||
|
||||
/**
|
||||
* 按工作流 JSON 值解析汇聚模式。
|
||||
*
|
||||
* @param value 配置值
|
||||
* @return 汇聚模式
|
||||
* @throws IllegalArgumentException 配置为空或不受支持
|
||||
*/
|
||||
public static NodeJoinMode ofValue(String value) {
|
||||
if (value == null || value.trim().isEmpty()) {
|
||||
throw new IllegalArgumentException("joinMode must be 'any' or 'all'");
|
||||
}
|
||||
String normalized = value.trim().toLowerCase(Locale.ROOT);
|
||||
for (NodeJoinMode mode : values()) {
|
||||
if (mode.value.equals(normalized)) {
|
||||
return mode;
|
||||
}
|
||||
}
|
||||
throw new IllegalArgumentException(
|
||||
"Unsupported joinMode: " + value + "; expected 'any' or 'all'");
|
||||
}
|
||||
}
|
||||
@@ -49,6 +49,11 @@ public class Parameter implements Serializable, Cloneable {
|
||||
*/
|
||||
protected List<Object> enums;
|
||||
|
||||
/**
|
||||
* 显示文案与实际值分离的结构化选项。
|
||||
*/
|
||||
protected List<ParameterOption> options;
|
||||
|
||||
/**
|
||||
* 用户输入的表单类型,例如:"input" "textarea" "select" "radio" "checkbox" 等等
|
||||
*/
|
||||
@@ -242,6 +247,14 @@ public class Parameter implements Serializable, Cloneable {
|
||||
}
|
||||
}
|
||||
|
||||
public List<ParameterOption> getOptions() {
|
||||
return options;
|
||||
}
|
||||
|
||||
public void setOptions(List<ParameterOption> options) {
|
||||
this.options = options;
|
||||
}
|
||||
|
||||
public String getFormType() {
|
||||
return formType;
|
||||
}
|
||||
@@ -298,6 +311,7 @@ public class Parameter implements Serializable, Cloneable {
|
||||
", flattenAggregation=" + flattenAggregation +
|
||||
", children=" + children +
|
||||
", enums=" + enums +
|
||||
", options=" + options +
|
||||
", formType='" + formType + '\'' +
|
||||
", formLabel='" + formLabel + '\'' +
|
||||
", formPlaceholder='" + formPlaceholder + '\'' +
|
||||
@@ -320,6 +334,13 @@ public class Parameter implements Serializable, Cloneable {
|
||||
clone.enums = new ArrayList<>(this.enums.size());
|
||||
clone.enums.addAll(this.enums);
|
||||
}
|
||||
if (this.options != null) {
|
||||
clone.options = new ArrayList<>(this.options.size());
|
||||
for (ParameterOption option : this.options) {
|
||||
clone.options.add(new ParameterOption(
|
||||
option.getLabel(), option.getValue()));
|
||||
}
|
||||
}
|
||||
return clone;
|
||||
} catch (CloneNotSupportedException e) {
|
||||
throw new AssertionError();
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
/**
|
||||
* 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.chain;
|
||||
|
||||
import java.io.Serializable;
|
||||
|
||||
/**
|
||||
* 用户输入参数的结构化选项。
|
||||
*/
|
||||
public class ParameterOption implements Serializable {
|
||||
|
||||
private static final long serialVersionUID = 1L;
|
||||
|
||||
private String label;
|
||||
private String value;
|
||||
|
||||
public ParameterOption() {
|
||||
}
|
||||
|
||||
public ParameterOption(String label, String value) {
|
||||
setLabel(label);
|
||||
setValue(value);
|
||||
}
|
||||
|
||||
public String getLabel() {
|
||||
return label;
|
||||
}
|
||||
|
||||
public void setLabel(String label) {
|
||||
this.label = trim(label);
|
||||
}
|
||||
|
||||
public String getValue() {
|
||||
return value;
|
||||
}
|
||||
|
||||
public void setValue(String value) {
|
||||
this.value = trim(value);
|
||||
}
|
||||
|
||||
private static String trim(String value) {
|
||||
return value == null ? null : value.trim();
|
||||
}
|
||||
|
||||
@Override
|
||||
public String toString() {
|
||||
return "ParameterOption{" +
|
||||
"label='" + label + '\'' +
|
||||
", value='" + value + '\'' +
|
||||
'}';
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -15,21 +15,55 @@
|
||||
*/
|
||||
package com.easyagents.flow.core.node;
|
||||
|
||||
|
||||
import com.easyagents.flow.core.chain.Chain;
|
||||
import com.easyagents.flow.core.chain.ChainSuspendException;
|
||||
import com.easyagents.flow.core.chain.DataType;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.chain.ParameterOption;
|
||||
import com.easyagents.flow.core.chain.RefType;
|
||||
import com.easyagents.flow.core.chain.repository.ChainStateField;
|
||||
import com.easyagents.flow.core.chain.runtime.Trigger;
|
||||
import com.easyagents.flow.core.chain.runtime.TriggerContext;
|
||||
import com.easyagents.flow.core.chain.runtime.TriggerType;
|
||||
|
||||
import java.util.*;
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collections;
|
||||
import java.util.EnumSet;
|
||||
import java.util.HashSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Objects;
|
||||
import java.util.Set;
|
||||
|
||||
public class ConfirmNode extends BaseNode {
|
||||
private static final long serialVersionUID = 1L;
|
||||
|
||||
public static final String DEFAULT_OUTPUT_NAME = "selection";
|
||||
public static final int MAX_OPTIONS = 100;
|
||||
public static final int MAX_MESSAGE_LENGTH = 2000;
|
||||
public static final int MAX_OPTION_LENGTH = 200;
|
||||
public static final Set<String> SUPPORTED_CONFIGURATION_KEYS = Set.of(
|
||||
"condition",
|
||||
"description",
|
||||
"expand",
|
||||
"joinMode",
|
||||
"loopBreakCondition",
|
||||
"loopEnable",
|
||||
"loopIntervalMs",
|
||||
"maxLoopCount",
|
||||
"maxRetryCount",
|
||||
"message",
|
||||
"multiple",
|
||||
"options",
|
||||
"outputDefs",
|
||||
"resetRetryCountAfterNormal",
|
||||
"retryEnable",
|
||||
"retryIntervalMs",
|
||||
"title");
|
||||
|
||||
private String message;
|
||||
private List<Parameter> confirms;
|
||||
private boolean multiple;
|
||||
private List<String> options;
|
||||
|
||||
public String getMessage() {
|
||||
return message;
|
||||
@@ -39,115 +73,152 @@ public class ConfirmNode extends BaseNode {
|
||||
this.message = message;
|
||||
}
|
||||
|
||||
public List<Parameter> getConfirms() {
|
||||
return confirms;
|
||||
public boolean isMultiple() {
|
||||
return multiple;
|
||||
}
|
||||
|
||||
public void setConfirms(List<Parameter> confirms) {
|
||||
if (confirms != null) {
|
||||
for (Parameter confirm : confirms) {
|
||||
confirm.setRefType(RefType.INPUT);
|
||||
confirm.setRequired(true); // 必填,才能正确通过 getParameterValuesOnly 获取参数值
|
||||
confirm.setName(confirm.getName());
|
||||
}
|
||||
}
|
||||
this.confirms = confirms;
|
||||
public void setMultiple(boolean multiple) {
|
||||
this.multiple = multiple;
|
||||
}
|
||||
|
||||
public List<String> getOptions() {
|
||||
return options;
|
||||
}
|
||||
|
||||
public void setOptions(List<String> options) {
|
||||
this.options = options;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Map<String, Object> execute(Chain chain) {
|
||||
validateConfiguration();
|
||||
String outputName = resolveOutputName();
|
||||
Parameter parameter = buildParameter();
|
||||
|
||||
List<Parameter> confirmParameters = new ArrayList<>();
|
||||
addConfirmParameter(confirmParameters);
|
||||
|
||||
if (confirms != null) {
|
||||
for (Parameter confirm : confirms) {
|
||||
Parameter clone = confirm.clone();
|
||||
clone.setName(confirm.getName() + "__" + getId());
|
||||
clone.setRefType(RefType.INPUT);
|
||||
confirmParameters.add(clone);
|
||||
}
|
||||
// 确认值只能来自经过 Chain.resumeIfSuspended 校验后创建的恢复触发器。
|
||||
// 启动参数与普通节点内存中的同名值均不能绕过人工确认。
|
||||
if (!isValidatedResumeTrigger(chain)) {
|
||||
chain.updateStateSafely(state -> {
|
||||
if (state.getMemory().remove(parameter.getName()) == null) {
|
||||
return null;
|
||||
}
|
||||
return EnumSet.of(ChainStateField.MEMORY);
|
||||
});
|
||||
}
|
||||
|
||||
Map<String, Object> values;
|
||||
try {
|
||||
values = chain.getExecutionState()
|
||||
.resolveParameters(this, confirmParameters);
|
||||
// 移除 confirm 参数,方便在其他节点二次确认,或者在 for 循环中第二次获取
|
||||
.resolveParameters(this, Collections.singletonList(parameter));
|
||||
chain.updateStateSafely(state -> {
|
||||
for (Parameter confirmParameter : confirmParameters) {
|
||||
state.getMemory().remove(confirmParameter.getName());
|
||||
if (!state.getMemory().containsKey(parameter.getName())) {
|
||||
return null;
|
||||
}
|
||||
state.getMemory().remove(parameter.getName());
|
||||
return EnumSet.of(ChainStateField.MEMORY);
|
||||
});
|
||||
} catch (ChainSuspendException e) {
|
||||
} catch (ChainSuspendException exception) {
|
||||
chain.updateStateSafely(state -> {
|
||||
state.setMessage(message);
|
||||
return EnumSet.of(ChainStateField.MESSAGE);
|
||||
});
|
||||
|
||||
if (confirms != null) {
|
||||
List<Parameter> newParameters = new ArrayList<>();
|
||||
for (Parameter confirm : confirms) {
|
||||
Parameter clone = confirm.clone();
|
||||
clone.setName(confirm.getName() + "__" + getId());
|
||||
clone.setRefType(RefType.REF); // 固定为 REF
|
||||
newParameters.add(clone);
|
||||
}
|
||||
|
||||
// 获取参数值,不会触发 ChainSuspendException 错误
|
||||
Map<String, Object> parameterValues =
|
||||
chain.getExecutionState().resolveParameters(
|
||||
this,
|
||||
newParameters,
|
||||
null,
|
||||
true);
|
||||
|
||||
// 设置 enums,方便前端给用户进行选择
|
||||
for (Parameter confirmParameter : confirmParameters) {
|
||||
if (confirmParameter.getEnums() == null) {
|
||||
Object enumsObject = parameterValues.get(confirmParameter.getName());
|
||||
confirmParameter.setEnumsObject(enumsObject);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
throw e;
|
||||
throw exception;
|
||||
}
|
||||
|
||||
return Collections.singletonMap(
|
||||
outputName,
|
||||
values.get(parameter.getName()));
|
||||
}
|
||||
|
||||
Map<String, Object> results = new HashMap<>(values.size());
|
||||
values.forEach((key, value) -> {
|
||||
int index = key.lastIndexOf("__");
|
||||
if (index >= 0) {
|
||||
results.put(key.substring(0, index), value);
|
||||
} else {
|
||||
results.put(key, value);
|
||||
/**
|
||||
* 获取并校验当前节点配置的唯一输出名称。
|
||||
*
|
||||
* @return 用户配置的输出名称
|
||||
*/
|
||||
public String resolveOutputName() {
|
||||
if (outputDefs == null || outputDefs.size() != 1
|
||||
|| outputDefs.get(0) == null) {
|
||||
throw new IllegalArgumentException(
|
||||
"用户确认节点必须配置一个输出参数");
|
||||
}
|
||||
Parameter output = outputDefs.get(0);
|
||||
String outputName = output.getName();
|
||||
if (outputName == null || outputName.trim().isEmpty()) {
|
||||
throw new IllegalArgumentException(
|
||||
"用户确认节点输出参数名称不能为空");
|
||||
}
|
||||
DataType expectedType = multiple
|
||||
? DataType.Array_String
|
||||
: DataType.String;
|
||||
if (output.getDataType() != expectedType) {
|
||||
throw new IllegalArgumentException(
|
||||
"用户确认节点输出参数类型必须与选择方式一致");
|
||||
}
|
||||
return outputName;
|
||||
}
|
||||
|
||||
private boolean isValidatedResumeTrigger(Chain chain) {
|
||||
Trigger trigger = TriggerContext.getCurrentTrigger();
|
||||
return trigger != null
|
||||
&& trigger.getType() == TriggerType.RESUME
|
||||
&& Objects.equals(
|
||||
chain.getStateInstanceId(),
|
||||
trigger.getStateInstanceId())
|
||||
&& Objects.equals(getId(), trigger.getNodeId());
|
||||
}
|
||||
|
||||
public void validateConfiguration() {
|
||||
requireText(message, MAX_MESSAGE_LENGTH, "确认提示内容");
|
||||
if (options == null || options.isEmpty()) {
|
||||
throw new IllegalArgumentException("用户确认节点至少需要一个选项");
|
||||
}
|
||||
if (options.size() > MAX_OPTIONS) {
|
||||
throw new IllegalArgumentException(
|
||||
"用户确认节点最多支持 " + MAX_OPTIONS + " 个选项");
|
||||
}
|
||||
|
||||
Set<String> normalizedOptions = new HashSet<>();
|
||||
for (String option : options) {
|
||||
String normalized = requireText(option, MAX_OPTION_LENGTH, "选项内容");
|
||||
if (!normalizedOptions.add(normalized)) {
|
||||
throw new IllegalArgumentException("用户确认节点选项内容重复: " + normalized);
|
||||
}
|
||||
});
|
||||
|
||||
return results;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
private void addConfirmParameter(List<Parameter> parameters) {
|
||||
// “确认 和 取消” 的参数
|
||||
private Parameter buildParameter() {
|
||||
Parameter parameter = new Parameter();
|
||||
parameter.setId(DEFAULT_OUTPUT_NAME);
|
||||
parameter.setName(DEFAULT_OUTPUT_NAME + "__" + getId());
|
||||
parameter.setDataType(multiple ? DataType.Array_String : DataType.String);
|
||||
parameter.setRefType(RefType.INPUT);
|
||||
parameter.setId("confirm");
|
||||
parameter.setName("confirm__" + getId());
|
||||
parameter.setRequired(true);
|
||||
|
||||
List<Object> selectionData = new ArrayList<>();
|
||||
selectionData.add("yes");
|
||||
selectionData.add("no");
|
||||
|
||||
parameter.setEnums(selectionData);
|
||||
parameter.setContentType("text");
|
||||
parameter.setFormType("confirm");
|
||||
parameters.add(parameter);
|
||||
parameter.setFormType(multiple ? "checkbox" : "radio");
|
||||
parameter.setFormLabel("选择内容");
|
||||
parameter.setOptions(buildRuntimeOptions());
|
||||
return parameter;
|
||||
}
|
||||
|
||||
private List<ParameterOption> buildRuntimeOptions() {
|
||||
List<ParameterOption> runtimeOptions = new ArrayList<>(options.size());
|
||||
for (String option : options) {
|
||||
String normalized = option.trim();
|
||||
runtimeOptions.add(new ParameterOption(normalized, normalized));
|
||||
}
|
||||
return runtimeOptions;
|
||||
}
|
||||
|
||||
private static String requireText(
|
||||
String value, int maxLength, String fieldName) {
|
||||
String normalized = value == null ? "" : value.trim();
|
||||
if (normalized.isEmpty()) {
|
||||
throw new IllegalArgumentException(fieldName + "不能为空");
|
||||
}
|
||||
if (normalized.length() > maxLength) {
|
||||
throw new IllegalArgumentException(
|
||||
fieldName + "不能超过 " + maxLength + " 个字符");
|
||||
}
|
||||
return normalized;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 + '\'' +
|
||||
|
||||
@@ -21,6 +21,7 @@ import com.alibaba.fastjson.JSONObject;
|
||||
import com.easyagents.flow.core.chain.DataType;
|
||||
import com.easyagents.flow.core.chain.JsCodeCondition;
|
||||
import com.easyagents.flow.core.chain.Node;
|
||||
import com.easyagents.flow.core.chain.NodeJoinMode;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.chain.RefType;
|
||||
import com.easyagents.flow.core.node.BaseNode;
|
||||
@@ -123,6 +124,10 @@ public abstract class BaseNodeParser<T extends BaseNode> implements NodeParser<T
|
||||
|
||||
if (!data.isEmpty()) {
|
||||
|
||||
if (data.containsKey("joinMode")) {
|
||||
node.setJoinMode(NodeJoinMode.ofValue(data.getString("joinMode")));
|
||||
}
|
||||
|
||||
addParameters(node, data);
|
||||
addOutputDefs(node, data);
|
||||
|
||||
|
||||
@@ -15,11 +15,12 @@
|
||||
*/
|
||||
package com.easyagents.flow.core.parser.impl;
|
||||
|
||||
import com.alibaba.fastjson.JSONArray;
|
||||
import com.alibaba.fastjson.JSONObject;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.node.ConfirmNode;
|
||||
import com.easyagents.flow.core.parser.BaseNodeParser;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
|
||||
public class ConfirmNodeParser extends BaseNodeParser<ConfirmNode> {
|
||||
@@ -28,12 +29,36 @@ public class ConfirmNodeParser extends BaseNodeParser<ConfirmNode> {
|
||||
public ConfirmNode doParse(JSONObject root, JSONObject data, JSONObject chainJSONObject) {
|
||||
|
||||
ConfirmNode confirmNode = new ConfirmNode();
|
||||
confirmNode.setMessage(data.getString("message"));
|
||||
|
||||
List<Parameter> confirms = getParameters(data, "confirms");
|
||||
if (confirms != null && !confirms.isEmpty()) {
|
||||
confirmNode.setConfirms(confirms);
|
||||
for (String key : data.keySet()) {
|
||||
if (!ConfirmNode.SUPPORTED_CONFIGURATION_KEYS.contains(key)) {
|
||||
throw new IllegalArgumentException(
|
||||
"用户确认节点包含无效配置字段: " + key);
|
||||
}
|
||||
}
|
||||
Object message = data.get("message");
|
||||
if (!(message instanceof String)) {
|
||||
throw new IllegalArgumentException("用户确认节点提示内容必须为字符串");
|
||||
}
|
||||
confirmNode.setMessage((String) message);
|
||||
|
||||
Object multiple = data.get("multiple");
|
||||
if (!(multiple instanceof Boolean)) {
|
||||
throw new IllegalArgumentException("用户确认节点选择方式必须为布尔值");
|
||||
}
|
||||
confirmNode.setMultiple((Boolean) multiple);
|
||||
|
||||
Object optionsValue = data.get("options");
|
||||
if (!(optionsValue instanceof JSONArray options)) {
|
||||
throw new IllegalArgumentException("用户确认节点选项必须为数组");
|
||||
}
|
||||
List<String> confirmOptions = new ArrayList<>(options.size());
|
||||
for (Object option : options) {
|
||||
if (!(option instanceof String)) {
|
||||
throw new IllegalArgumentException("用户确认节点选项内容必须为字符串");
|
||||
}
|
||||
confirmOptions.add((String) option);
|
||||
}
|
||||
confirmNode.setOptions(confirmOptions);
|
||||
|
||||
return confirmNode;
|
||||
}
|
||||
|
||||
@@ -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"));
|
||||
|
||||
@@ -15,10 +15,15 @@
|
||||
*/
|
||||
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.ChainStatus;
|
||||
import com.easyagents.flow.core.chain.DataType;
|
||||
import com.easyagents.flow.core.chain.Edge;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.chain.RefType;
|
||||
import com.easyagents.flow.core.chain.ChainState;
|
||||
import com.easyagents.flow.core.chain.event.ChainEndEvent;
|
||||
import com.easyagents.flow.core.chain.repository.ChainDefinitionSnapshotRepository;
|
||||
@@ -35,6 +40,7 @@ import com.easyagents.flow.core.node.EndNode;
|
||||
import com.easyagents.flow.core.node.BaseNode;
|
||||
import com.easyagents.flow.core.node.ConfirmNode;
|
||||
import com.easyagents.flow.core.node.StartNode;
|
||||
import com.easyagents.flow.core.parser.ChainParser;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
@@ -61,6 +67,127 @@ import java.util.concurrent.atomic.AtomicReference;
|
||||
*/
|
||||
public class ChainExecutorConcurrencyTest {
|
||||
|
||||
/**
|
||||
* 验证启动变量不能伪造确认节点的恢复参数并绕过人工确认。
|
||||
*/
|
||||
@Test
|
||||
public void shouldSuspendConfirmNodeDespitePrefilledStartVariable()
|
||||
throws Exception {
|
||||
ScheduledExecutorService schedulerPool =
|
||||
Executors.newSingleThreadScheduledExecutor();
|
||||
ExecutorService workerPool = Executors.newFixedThreadPool(2);
|
||||
TriggerScheduler triggerScheduler = new TriggerScheduler(
|
||||
new InMemoryTriggerStore(), schedulerPool, workerPool, 10L);
|
||||
ChainDefinition definition = createConfirmDefinition();
|
||||
InMemoryChainStateRepository stateRepository =
|
||||
new InMemoryChainStateRepository();
|
||||
ChainExecutor executor = new ChainExecutor(
|
||||
ignored -> definition,
|
||||
stateRepository,
|
||||
new InMemoryNodeStateRepository(),
|
||||
triggerScheduler);
|
||||
|
||||
try {
|
||||
String instanceId = executor.executeAsync(
|
||||
definition.getId(),
|
||||
Map.of("selection__confirm", "未配置值"));
|
||||
|
||||
ChainState state = awaitStatus(
|
||||
stateRepository, instanceId, ChainStatus.SUSPEND);
|
||||
Assert.assertFalse(
|
||||
state.getMemory().containsKey("selection__confirm"));
|
||||
Assert.assertEquals(
|
||||
"selection__confirm",
|
||||
state.getSuspendForParameters().get(0).getName());
|
||||
} finally {
|
||||
triggerScheduler.shutdown();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证设计器最终契约可解析、挂起、恢复,并把用户选择按配置名称交给结束节点。
|
||||
*/
|
||||
@Test
|
||||
public void shouldFlowConfiguredConfirmOutputToEndNode()
|
||||
throws Exception {
|
||||
ScheduledExecutorService schedulerPool =
|
||||
Executors.newSingleThreadScheduledExecutor();
|
||||
ExecutorService workerPool = Executors.newFixedThreadPool(2);
|
||||
TriggerScheduler triggerScheduler = new TriggerScheduler(
|
||||
new InMemoryTriggerStore(), schedulerPool, workerPool, 10L);
|
||||
ChainDefinition definition = createParsedConfirmDefinition();
|
||||
InMemoryChainStateRepository stateRepository =
|
||||
new InMemoryChainStateRepository();
|
||||
ChainExecutor executor = new ChainExecutor(
|
||||
ignored -> definition,
|
||||
stateRepository,
|
||||
new InMemoryNodeStateRepository(),
|
||||
triggerScheduler);
|
||||
|
||||
try {
|
||||
String instanceId = executor.executeAsync(
|
||||
definition.getId(), Collections.emptyMap());
|
||||
ChainState suspended = awaitStatus(
|
||||
stateRepository, instanceId, ChainStatus.SUSPEND);
|
||||
|
||||
Assert.assertEquals(
|
||||
"selection__confirm",
|
||||
suspended.getSuspendForParameters().get(0).getName());
|
||||
Assert.assertTrue(executor.resumeAsyncIfSuspended(
|
||||
instanceId,
|
||||
Map.of("selection__confirm", "继续")));
|
||||
|
||||
ChainState completed = awaitStatus(
|
||||
stateRepository, instanceId, ChainStatus.SUCCEEDED);
|
||||
Assert.assertEquals("继续", completed.getExecuteResult().get("result"));
|
||||
} finally {
|
||||
triggerScheduler.shutdown();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证多选确认结果以字符串数组形式流转到结束节点。
|
||||
*/
|
||||
@Test
|
||||
public void shouldFlowMultipleConfirmOutputToEndNode()
|
||||
throws Exception {
|
||||
ScheduledExecutorService schedulerPool =
|
||||
Executors.newSingleThreadScheduledExecutor();
|
||||
ExecutorService workerPool = Executors.newFixedThreadPool(2);
|
||||
TriggerScheduler triggerScheduler = new TriggerScheduler(
|
||||
new InMemoryTriggerStore(), schedulerPool, workerPool, 10L);
|
||||
ChainDefinition definition = createParsedConfirmDefinition(true);
|
||||
InMemoryChainStateRepository stateRepository =
|
||||
new InMemoryChainStateRepository();
|
||||
ChainExecutor executor = new ChainExecutor(
|
||||
ignored -> definition,
|
||||
stateRepository,
|
||||
new InMemoryNodeStateRepository(),
|
||||
triggerScheduler);
|
||||
|
||||
try {
|
||||
String instanceId = executor.executeAsync(
|
||||
definition.getId(), Collections.emptyMap());
|
||||
ChainState suspended = awaitStatus(
|
||||
stateRepository, instanceId, ChainStatus.SUSPEND);
|
||||
|
||||
Assert.assertEquals(
|
||||
DataType.Array_String,
|
||||
suspended.getSuspendForParameters().get(0).getDataType());
|
||||
List<String> selection = List.of("继续", "停止");
|
||||
Assert.assertTrue(executor.resumeAsyncIfSuspended(
|
||||
instanceId,
|
||||
Map.of("selection__confirm", selection)));
|
||||
|
||||
ChainState completed = awaitStatus(
|
||||
stateRepository, instanceId, ChainStatus.SUCCEEDED);
|
||||
Assert.assertEquals(
|
||||
selection, completed.getExecuteResult().get("result"));
|
||||
} finally {
|
||||
triggerScheduler.shutdown();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证同步 Tool 入口遇到人工挂起会快速失败,不会无限占用调用线程。
|
||||
*
|
||||
@@ -563,6 +690,12 @@ public class ChainExecutorConcurrencyTest {
|
||||
start.setId("start");
|
||||
ConfirmNode confirm = new ConfirmNode();
|
||||
confirm.setId("confirm");
|
||||
confirm.setMessage("请选择是否继续");
|
||||
confirm.setOptions(List.of("继续", "停止"));
|
||||
confirm.setOutputDefs(Collections.singletonList(
|
||||
new Parameter(
|
||||
ConfirmNode.DEFAULT_OUTPUT_NAME,
|
||||
DataType.String)));
|
||||
EndNode end = new EndNode();
|
||||
end.setId("end");
|
||||
Edge first = new Edge();
|
||||
@@ -581,6 +714,91 @@ public class ChainExecutorConcurrencyTest {
|
||||
return definition;
|
||||
}
|
||||
|
||||
private ChainDefinition createParsedConfirmDefinition() {
|
||||
return createParsedConfirmDefinition(false);
|
||||
}
|
||||
|
||||
private ChainDefinition createParsedConfirmDefinition(boolean multiple) {
|
||||
JSONArray nodes = new JSONArray();
|
||||
nodes.add(nodeJson("start", "startNode", new JSONObject()));
|
||||
|
||||
JSONObject confirmData = new JSONObject();
|
||||
confirmData.put("message", "请选择是否继续");
|
||||
confirmData.put("multiple", multiple);
|
||||
confirmData.put("options", new JSONArray(List.of("继续", "停止")));
|
||||
String outputType = multiple ? "Array<String>" : "String";
|
||||
confirmData.put("outputDefs", new JSONArray(List.of(
|
||||
parameterJson("templateType", outputType, null))));
|
||||
nodes.add(nodeJson("confirm", "confirmNode", confirmData));
|
||||
|
||||
JSONObject endData = new JSONObject();
|
||||
endData.put("outputDefs", new JSONArray(List.of(
|
||||
parameterJson(
|
||||
"result", outputType, "confirm.templateType"))));
|
||||
nodes.add(nodeJson("end", "endNode", endData));
|
||||
|
||||
JSONArray edges = new JSONArray();
|
||||
edges.add(edgeJson("start-to-confirm", "start", "confirm"));
|
||||
edges.add(edgeJson("confirm-to-end", "confirm", "end"));
|
||||
JSONObject flow = new JSONObject();
|
||||
flow.put("nodes", nodes);
|
||||
flow.put("edges", edges);
|
||||
|
||||
ChainDefinition definition = ChainParser.builder()
|
||||
.withDefaultParsers(true)
|
||||
.build()
|
||||
.parse(flow.toJSONString());
|
||||
definition.setId("confirm-output-flow-test");
|
||||
return definition;
|
||||
}
|
||||
|
||||
private JSONObject nodeJson(
|
||||
String id, String type, JSONObject data) {
|
||||
JSONObject node = new JSONObject();
|
||||
node.put("id", id);
|
||||
node.put("type", type);
|
||||
node.put("data", data);
|
||||
return node;
|
||||
}
|
||||
|
||||
private JSONObject edgeJson(
|
||||
String id, String source, String target) {
|
||||
JSONObject edge = new JSONObject();
|
||||
edge.put("id", id);
|
||||
edge.put("source", source);
|
||||
edge.put("target", target);
|
||||
return edge;
|
||||
}
|
||||
|
||||
private JSONObject parameterJson(
|
||||
String name, String dataType, String ref) {
|
||||
JSONObject parameter = new JSONObject();
|
||||
parameter.put("name", name);
|
||||
parameter.put("dataType", dataType);
|
||||
if (ref != null) {
|
||||
parameter.put("ref", ref);
|
||||
parameter.put("refType", RefType.REF.toString());
|
||||
}
|
||||
return parameter;
|
||||
}
|
||||
|
||||
private ChainState awaitStatus(
|
||||
InMemoryChainStateRepository repository,
|
||||
String instanceId,
|
||||
ChainStatus expected) throws InterruptedException {
|
||||
long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(3);
|
||||
ChainState state;
|
||||
do {
|
||||
state = repository.load(instanceId);
|
||||
if (state != null && state.getStatus() == expected) {
|
||||
return state;
|
||||
}
|
||||
Thread.sleep(10L);
|
||||
} while (System.nanoTime() < deadline);
|
||||
Assert.fail("workflow did not reach status " + expected);
|
||||
return state;
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建用于取消传播验证的工作流。
|
||||
*
|
||||
|
||||
@@ -2,15 +2,29 @@ package com.easyagents.flow.core.test;
|
||||
|
||||
import com.easyagents.flow.core.chain.Chain;
|
||||
import com.easyagents.flow.core.chain.ChainDefinition;
|
||||
import com.easyagents.flow.core.chain.ChainResumeException;
|
||||
import com.easyagents.flow.core.chain.ChainState;
|
||||
import com.easyagents.flow.core.chain.ChainStatus;
|
||||
import com.easyagents.flow.core.chain.EventManager;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.chain.ParameterOption;
|
||||
import com.easyagents.flow.core.chain.repository.ChainStateField;
|
||||
import com.easyagents.flow.core.chain.repository.InMemoryChainStateRepository;
|
||||
import com.easyagents.flow.core.chain.repository.InMemoryNodeStateRepository;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.util.Arrays;
|
||||
import java.util.Collections;
|
||||
import java.util.EnumSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.CountDownLatch;
|
||||
import java.util.concurrent.ExecutorService;
|
||||
import java.util.concurrent.Executors;
|
||||
import java.util.concurrent.Future;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
|
||||
/**
|
||||
* {@link Chain} 暂停恢复状态守卫测试。
|
||||
@@ -72,6 +86,167 @@ public class ChainResumeGuardTest {
|
||||
.containsKey("unexpected"));
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证确认选项只能按挂起时声明的字段和值恢复。
|
||||
*/
|
||||
@Test
|
||||
public void shouldValidateDeclaredResumeOptionsBeforeStateTransition() {
|
||||
InMemoryChainStateRepository stateRepository =
|
||||
new InMemoryChainStateRepository();
|
||||
Chain chain = createChain(stateRepository, "resume-options");
|
||||
Parameter single = optionParameter(
|
||||
"templateType__confirm", "会议类型", "radio");
|
||||
Parameter multiple = optionParameter(
|
||||
"participants__confirm", "参会人员", "checkbox");
|
||||
chain.getExecutionState().setSuspendForParameters(
|
||||
Arrays.asList(single, multiple));
|
||||
chain.suspend();
|
||||
|
||||
ChainResumeException invalidOption = assertRejected(
|
||||
chain,
|
||||
Map.of(
|
||||
"templateType__confirm", "UNKNOWN",
|
||||
"participants__confirm", List.of("REVIEW")));
|
||||
ChainResumeException extraField = assertRejected(
|
||||
chain,
|
||||
Map.of(
|
||||
"templateType__confirm", "AGENDA",
|
||||
"participants__confirm", List.of("REVIEW", "REVIEW")));
|
||||
assertRejected(
|
||||
chain,
|
||||
Map.of(
|
||||
"templateType__confirm", "AGENDA",
|
||||
"participants__confirm", List.of("REVIEW"),
|
||||
"extra", "value"));
|
||||
Assert.assertFalse(invalidOption.getMessage().contains("UNKNOWN"));
|
||||
Assert.assertFalse(extraField.getMessage().contains("extra"));
|
||||
|
||||
Assert.assertEquals(
|
||||
ChainStatus.SUSPEND,
|
||||
stateRepository.load("resume-options").getStatus());
|
||||
Assert.assertTrue(
|
||||
stateRepository.load("resume-options").getMemory().isEmpty());
|
||||
|
||||
boolean resumed = chain.resumeIfSuspended(Map.of(
|
||||
"templateType__confirm", "AGENDA",
|
||||
"participants__confirm", List.of("REVIEW", "BRIEFING")));
|
||||
|
||||
Assert.assertTrue(resumed);
|
||||
Assert.assertEquals(
|
||||
"AGENDA",
|
||||
stateRepository.load("resume-options")
|
||||
.getMemory()
|
||||
.get("templateType__confirm"));
|
||||
Assert.assertEquals(
|
||||
List.of("REVIEW", "BRIEFING"),
|
||||
stateRepository.load("resume-options")
|
||||
.getMemory()
|
||||
.get("participants__confirm"));
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证并发恢复同一暂停实例时,只有一个请求可以完成状态转换。
|
||||
*/
|
||||
@Test
|
||||
public void shouldAllowOnlyOneConcurrentResume() throws Exception {
|
||||
InMemoryChainStateRepository stateRepository =
|
||||
new InMemoryChainStateRepository();
|
||||
Chain first = createChain(stateRepository, "resume-concurrent");
|
||||
Chain second = createChain(stateRepository, "resume-concurrent");
|
||||
Parameter parameter = optionParameter(
|
||||
"templateType__confirm", "会议类型", "radio");
|
||||
first.getExecutionState().setSuspendForParameters(
|
||||
List.of(parameter));
|
||||
first.suspend();
|
||||
|
||||
CountDownLatch ready = new CountDownLatch(2);
|
||||
CountDownLatch start = new CountDownLatch(1);
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
Future<Boolean> firstResult = executor.submit(() -> {
|
||||
ready.countDown();
|
||||
start.await();
|
||||
return first.resumeIfSuspended(
|
||||
Map.of("templateType__confirm", "AGENDA"));
|
||||
});
|
||||
Future<Boolean> secondResult = executor.submit(() -> {
|
||||
ready.countDown();
|
||||
start.await();
|
||||
return second.resumeIfSuspended(
|
||||
Map.of("templateType__confirm", "REVIEW"));
|
||||
});
|
||||
|
||||
Assert.assertTrue(ready.await(5, TimeUnit.SECONDS));
|
||||
start.countDown();
|
||||
int resumedCount = (firstResult.get(5, TimeUnit.SECONDS) ? 1 : 0)
|
||||
+ (secondResult.get(5, TimeUnit.SECONDS) ? 1 : 0);
|
||||
|
||||
Assert.assertEquals(1, resumedCount);
|
||||
Assert.assertEquals(
|
||||
ChainStatus.RUNNING,
|
||||
stateRepository.load("resume-concurrent").getStatus());
|
||||
Object selected = stateRepository.load("resume-concurrent")
|
||||
.getMemory()
|
||||
.get("templateType__confirm");
|
||||
Assert.assertTrue(
|
||||
"AGENDA".equals(selected) || "REVIEW".equals(selected));
|
||||
} finally {
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证恢复变量、状态和暂停上下文通过一次原子更新完成。
|
||||
*/
|
||||
@Test
|
||||
public void shouldCommitResumeTransitionInSingleStateUpdate() {
|
||||
CountingChainStateRepository stateRepository =
|
||||
new CountingChainStateRepository();
|
||||
Chain chain = createChain(stateRepository, "resume-single-update");
|
||||
chain.getExecutionState().setSuspendForParameters(List.of(
|
||||
optionParameter("templateType__confirm", "会议类型", "radio")));
|
||||
chain.suspend();
|
||||
stateRepository.resetUpdateCount();
|
||||
|
||||
boolean resumed = chain.resumeIfSuspended(
|
||||
Map.of("templateType__confirm", "AGENDA"));
|
||||
|
||||
ChainState state = stateRepository.load("resume-single-update");
|
||||
Assert.assertTrue(resumed);
|
||||
Assert.assertEquals(1, stateRepository.getUpdateCount());
|
||||
Assert.assertEquals(ChainStatus.RUNNING, state.getStatus());
|
||||
Assert.assertNull(state.getSuspendNodeIds());
|
||||
Assert.assertNull(state.getSuspendForParameters());
|
||||
Assert.assertEquals("AGENDA", state.getMemory().get(
|
||||
"templateType__confirm"));
|
||||
}
|
||||
|
||||
private ChainResumeException assertRejected(
|
||||
Chain chain, Map<String, Object> variables) {
|
||||
try {
|
||||
chain.resumeIfSuspended(variables);
|
||||
Assert.fail("invalid resume variables must be rejected");
|
||||
return null;
|
||||
} catch (ChainResumeException expected) {
|
||||
Assert.assertNotNull(expected.getMessage());
|
||||
return expected;
|
||||
}
|
||||
}
|
||||
|
||||
private Parameter optionParameter(
|
||||
String name, String label, String formType) {
|
||||
Parameter parameter = new Parameter();
|
||||
parameter.setName(name);
|
||||
parameter.setFormLabel(label);
|
||||
parameter.setFormType(formType);
|
||||
parameter.setRequired(true);
|
||||
parameter.setOptions(List.of(
|
||||
new ParameterOption("第一议题", "AGENDA"),
|
||||
new ParameterOption("审议类", "REVIEW"),
|
||||
new ParameterOption("听取类", "BRIEFING")));
|
||||
return parameter;
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建使用进程内状态仓储的最小工作流。
|
||||
*
|
||||
@@ -93,4 +268,25 @@ public class ChainResumeGuardTest {
|
||||
chain.setEventManager(new EventManager());
|
||||
return chain;
|
||||
}
|
||||
|
||||
private static class CountingChainStateRepository
|
||||
extends InMemoryChainStateRepository {
|
||||
private final AtomicInteger updateCount = new AtomicInteger();
|
||||
|
||||
@Override
|
||||
public boolean tryUpdate(
|
||||
ChainState chainState,
|
||||
EnumSet<ChainStateField> fields) {
|
||||
updateCount.incrementAndGet();
|
||||
return super.tryUpdate(chainState, fields);
|
||||
}
|
||||
|
||||
private int getUpdateCount() {
|
||||
return updateCount.get();
|
||||
}
|
||||
|
||||
private void resetUpdateCount() {
|
||||
updateCount.set(0);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,235 @@
|
||||
/**
|
||||
* 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.ChainSuspendException;
|
||||
import com.easyagents.flow.core.chain.DataType;
|
||||
import com.easyagents.flow.core.chain.EventManager;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.chain.repository.InMemoryChainStateRepository;
|
||||
import com.easyagents.flow.core.chain.repository.InMemoryNodeStateRepository;
|
||||
import com.easyagents.flow.core.chain.runtime.Trigger;
|
||||
import com.easyagents.flow.core.chain.runtime.TriggerContext;
|
||||
import com.easyagents.flow.core.chain.runtime.TriggerType;
|
||||
import com.easyagents.flow.core.node.ConfirmNode;
|
||||
import com.easyagents.flow.core.parser.impl.ConfirmNodeParser;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.util.Collections;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/**
|
||||
* 用户确认节点选择与输出契约测试。
|
||||
*/
|
||||
public class ConfirmNodeTest {
|
||||
|
||||
@Test
|
||||
public void shouldBuildSingleChoiceAndReturnSelectedContent() {
|
||||
ConfirmNode node = parse(false);
|
||||
node.setId("confirm-1");
|
||||
Chain chain = createChain();
|
||||
|
||||
Parameter parameter = suspend(node, chain);
|
||||
Assert.assertEquals("selection__confirm-1", parameter.getName());
|
||||
Assert.assertEquals("radio", parameter.getFormType());
|
||||
Assert.assertEquals("选择内容", parameter.getFormLabel());
|
||||
Assert.assertEquals(DataType.String, parameter.getDataType());
|
||||
Assert.assertEquals("第一议题", parameter.getOptions().get(0).getLabel());
|
||||
Assert.assertEquals("第一议题", parameter.getOptions().get(0).getValue());
|
||||
|
||||
chain.getExecutionState().getMemory().put(parameter.getName(), "审议类");
|
||||
Map<String, Object> result = executeAsResume(node, chain);
|
||||
|
||||
Assert.assertEquals(Collections.singletonMap("selection", "审议类"), result);
|
||||
Assert.assertFalse(chain.getExecutionState().getMemory()
|
||||
.containsKey(parameter.getName()));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldBuildMultipleChoiceAndReturnSelectedContents() {
|
||||
ConfirmNode node = parse(true);
|
||||
node.setId("confirm-1");
|
||||
Chain chain = createChain();
|
||||
|
||||
Parameter parameter = suspend(node, chain);
|
||||
Assert.assertEquals("checkbox", parameter.getFormType());
|
||||
Assert.assertEquals(DataType.Array_String, parameter.getDataType());
|
||||
|
||||
List<String> selected = List.of("第一议题", "听取类");
|
||||
chain.getExecutionState().getMemory().put(parameter.getName(), selected);
|
||||
|
||||
Assert.assertEquals(selected, executeAsResume(node, chain).get("selection"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldUseConfiguredOutputNameWithoutChangingSuspendParameter() {
|
||||
ConfirmNode node = parse(false, "templateChoice");
|
||||
node.setId("confirm-1");
|
||||
Chain chain = createChain();
|
||||
|
||||
Parameter parameter = suspend(node, chain);
|
||||
Assert.assertEquals("selection__confirm-1", parameter.getName());
|
||||
|
||||
chain.getExecutionState().getMemory().put(
|
||||
parameter.getName(), "第一议题");
|
||||
Assert.assertEquals(
|
||||
Collections.singletonMap("templateChoice", "第一议题"),
|
||||
executeAsResume(node, chain));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldIgnorePrefilledValueWithoutResumeTrigger() {
|
||||
ConfirmNode node = parse(false);
|
||||
node.setId("confirm-1");
|
||||
Chain chain = createChain();
|
||||
chain.getExecutionState().getMemory().put(
|
||||
"selection__confirm-1", "未配置值");
|
||||
|
||||
Parameter parameter = suspend(node, chain);
|
||||
|
||||
Assert.assertEquals("selection__confirm-1", parameter.getName());
|
||||
Assert.assertFalse(chain.getExecutionState().getMemory()
|
||||
.containsKey(parameter.getName()));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectDuplicateNormalizedOptionContents() {
|
||||
ConfirmNode node = parse(false);
|
||||
node.setOptions(List.of("审议类", " 审议类 "));
|
||||
|
||||
try {
|
||||
node.validateConfiguration();
|
||||
Assert.fail("duplicate option contents must be rejected");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertTrue(expected.getMessage().contains("选项内容重复"));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectNonStringOptionDuringParsing() {
|
||||
JSONObject data = data(false);
|
||||
data.getJSONArray("options").add(1);
|
||||
|
||||
try {
|
||||
new ConfirmNodeParser().doParse(
|
||||
new JSONObject(), data, new JSONObject());
|
||||
Assert.fail("non-string option must be rejected");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertTrue(expected.getMessage().contains("必须为字符串"));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectStringEncodedOptionsDuringParsing() {
|
||||
JSONObject data = data(false);
|
||||
data.put("options", "[\"第一议题\"]");
|
||||
|
||||
try {
|
||||
new ConfirmNodeParser().doParse(
|
||||
new JSONObject(), data, new JSONObject());
|
||||
Assert.fail("string encoded options must be rejected");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertTrue(expected.getMessage().contains("必须为数组"));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectImplicitModeConversionDuringParsing() {
|
||||
JSONObject data = data(false);
|
||||
data.put("multiple", "false");
|
||||
|
||||
try {
|
||||
new ConfirmNodeParser().doParse(
|
||||
new JSONObject(), data, new JSONObject());
|
||||
Assert.fail("non-boolean mode must be rejected");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertTrue(expected.getMessage().contains("必须为布尔值"));
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectUnknownConfigurationFieldDuringParsing() {
|
||||
JSONObject data = data(false);
|
||||
data.put("async", true);
|
||||
|
||||
try {
|
||||
new ConfirmNodeParser().doParse(
|
||||
new JSONObject(), data, new JSONObject());
|
||||
Assert.fail("unknown confirm configuration must be rejected");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertTrue(expected.getMessage().contains("无效配置字段"));
|
||||
}
|
||||
}
|
||||
|
||||
private static ConfirmNode parse(boolean multiple) {
|
||||
return parse(multiple, ConfirmNode.DEFAULT_OUTPUT_NAME);
|
||||
}
|
||||
|
||||
private static ConfirmNode parse(boolean multiple, String outputName) {
|
||||
ConfirmNode node = new ConfirmNodeParser().doParse(
|
||||
new JSONObject(), data(multiple), new JSONObject());
|
||||
Parameter output = new Parameter();
|
||||
output.setName(outputName);
|
||||
output.setDataType(multiple
|
||||
? DataType.Array_String
|
||||
: DataType.String);
|
||||
node.setOutputDefs(Collections.singletonList(output));
|
||||
node.validateConfiguration();
|
||||
return node;
|
||||
}
|
||||
|
||||
private static JSONObject data(boolean multiple) {
|
||||
JSONObject data = new JSONObject();
|
||||
data.put("message", "请选择会议纪要模板");
|
||||
data.put("multiple", multiple);
|
||||
JSONArray options = new JSONArray();
|
||||
options.addAll(List.of("第一议题", "审议类", "听取类"));
|
||||
data.put("options", options);
|
||||
return data;
|
||||
}
|
||||
|
||||
private static Parameter suspend(ConfirmNode node, Chain chain) {
|
||||
try {
|
||||
node.execute(chain);
|
||||
throw new AssertionError("confirm node must suspend");
|
||||
} catch (ChainSuspendException expected) {
|
||||
Assert.assertEquals(1, expected.getSuspendParameters().size());
|
||||
return expected.getSuspendParameters().get(0);
|
||||
}
|
||||
}
|
||||
|
||||
private static Map<String, Object> executeAsResume(
|
||||
ConfirmNode node, Chain chain) {
|
||||
Trigger trigger = new Trigger();
|
||||
trigger.setType(TriggerType.RESUME);
|
||||
trigger.setStateInstanceId(chain.getStateInstanceId());
|
||||
trigger.setNodeId(node.getId());
|
||||
TriggerContext.setCurrentTrigger(trigger);
|
||||
try {
|
||||
return node.execute(chain);
|
||||
} finally {
|
||||
TriggerContext.clearCurrentTrigger();
|
||||
}
|
||||
}
|
||||
|
||||
private static Chain createChain() {
|
||||
ChainDefinition definition = new ChainDefinition();
|
||||
definition.setId("confirm-node-test");
|
||||
definition.setNodes(Collections.emptyList());
|
||||
definition.setEdges(Collections.emptyList());
|
||||
Chain chain = new Chain(definition, "confirm-node-instance");
|
||||
chain.setChainStateRepository(new InMemoryChainStateRepository());
|
||||
chain.setNodeStateRepository(new InMemoryNodeStateRepository());
|
||||
chain.setEventManager(new EventManager());
|
||||
return chain;
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,407 @@
|
||||
package com.easyagents.flow.core.test;
|
||||
|
||||
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.chain.ChainStatus;
|
||||
import com.easyagents.flow.core.chain.Edge;
|
||||
import com.easyagents.flow.core.chain.Node;
|
||||
import com.easyagents.flow.core.chain.NodeCondition;
|
||||
import com.easyagents.flow.core.chain.NodeJoinMode;
|
||||
import com.easyagents.flow.core.chain.NodeState;
|
||||
import com.easyagents.flow.core.chain.Parameter;
|
||||
import com.easyagents.flow.core.chain.RefType;
|
||||
import com.easyagents.flow.core.chain.repository.InMemoryChainStateRepository;
|
||||
import com.easyagents.flow.core.chain.repository.InMemoryNodeStateRepository;
|
||||
import com.easyagents.flow.core.chain.runtime.ChainExecutor;
|
||||
import com.easyagents.flow.core.chain.runtime.InMemoryTriggerStore;
|
||||
import com.easyagents.flow.core.chain.runtime.TriggerScheduler;
|
||||
import com.easyagents.flow.core.node.BaseNode;
|
||||
import com.easyagents.flow.core.node.EndNode;
|
||||
import com.easyagents.flow.core.node.StartNode;
|
||||
import com.easyagents.flow.core.parser.ChainParser;
|
||||
import com.easyagents.flow.core.parser.impl.EndNodeParser;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.util.Collections;
|
||||
import java.util.HashMap;
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.CountDownLatch;
|
||||
import java.util.concurrent.ExecutorService;
|
||||
import java.util.concurrent.Executors;
|
||||
import java.util.concurrent.ScheduledExecutorService;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
import java.util.function.BooleanSupplier;
|
||||
|
||||
/**
|
||||
* 验证普通节点的直接入边汇聚模式。
|
||||
*/
|
||||
public class NodeJoinModeTest {
|
||||
|
||||
@Test
|
||||
public void shouldParseJoinModeAndRejectInvalidValues() {
|
||||
ChainParser parser = ChainParser.builder()
|
||||
.withDefaultParsers(true)
|
||||
.build();
|
||||
|
||||
Assert.assertEquals(
|
||||
NodeJoinMode.ANY,
|
||||
parseEndNode(parser, null, null).getJoinMode());
|
||||
Assert.assertEquals(
|
||||
NodeJoinMode.ALL,
|
||||
parseEndNode(parser, "all", null).getJoinMode());
|
||||
Assert.assertEquals(
|
||||
NodeJoinMode.ANY,
|
||||
parseEndNode(parser, "ANY", null).getJoinMode());
|
||||
|
||||
assertInvalidJoinMode(() -> parseEndNode(parser, "first", null));
|
||||
assertInvalidJoinMode(() -> parseEndNode(parser, "", null));
|
||||
assertInvalidJoinMode(() -> parseEndNode(parser, "all", "loop"));
|
||||
|
||||
ProbeJoinNode node = new ProbeJoinNode(new AtomicInteger());
|
||||
node.setJoinMode(NodeJoinMode.ALL);
|
||||
assertInvalidJoinMode(() -> node.setParentId("loop"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldWaitForEveryInboundEdgeBeforeExecutingAndCheckingCondition()
|
||||
throws Exception {
|
||||
JoinFixture fixture = createJoinFixture(NodeJoinMode.ALL, true);
|
||||
String instanceId = null;
|
||||
try {
|
||||
instanceId = fixture.executor.executeAsync(
|
||||
fixture.definition.getId(), Collections.emptyMap());
|
||||
Assert.assertTrue(fixture.branchACompleted.await(2, TimeUnit.SECONDS));
|
||||
Assert.assertTrue(fixture.branchBStarted.await(2, TimeUnit.SECONDS));
|
||||
|
||||
String currentInstanceId = instanceId;
|
||||
await(() -> hasTriggerEdge(
|
||||
fixture.nodeStateRepository,
|
||||
currentInstanceId,
|
||||
"join",
|
||||
"a-join"));
|
||||
Assert.assertEquals(0, fixture.joinExecutions.get());
|
||||
Assert.assertEquals(0, fixture.conditionChecks.get());
|
||||
|
||||
fixture.releaseBranchB.countDown();
|
||||
ChainState finalState = awaitTerminal(
|
||||
fixture.chainStateRepository, instanceId);
|
||||
|
||||
Assert.assertEquals(1, fixture.joinExecutions.get());
|
||||
Assert.assertEquals(1, fixture.conditionChecks.get());
|
||||
Assert.assertEquals("A+B", finalState.getExecuteResult().get("combined"));
|
||||
Assert.assertEquals(Boolean.TRUE, finalState.getExecuteResult().get("sawBoth"));
|
||||
} finally {
|
||||
fixture.releaseBranchB.countDown();
|
||||
fixture.scheduler.shutdown();
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldKeepAnyModeFirstArrivalBehavior() throws Exception {
|
||||
JoinFixture fixture = createJoinFixture(NodeJoinMode.ANY, false);
|
||||
try {
|
||||
String instanceId = fixture.executor.executeAsync(
|
||||
fixture.definition.getId(), Collections.emptyMap());
|
||||
Assert.assertTrue(fixture.branchACompleted.await(2, TimeUnit.SECONDS));
|
||||
Assert.assertTrue(fixture.branchBStarted.await(2, TimeUnit.SECONDS));
|
||||
|
||||
await(() -> fixture.joinExecutions.get() > 0);
|
||||
ChainState finalState = awaitTerminal(
|
||||
fixture.chainStateRepository, instanceId);
|
||||
|
||||
Assert.assertEquals(1, fixture.joinExecutions.get());
|
||||
Assert.assertEquals(Boolean.FALSE, finalState.getExecuteResult().get("sawBoth"));
|
||||
} finally {
|
||||
fixture.releaseBranchB.countDown();
|
||||
fixture.scheduler.shutdown();
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldKeepDefaultAnyRetryAndLoopBehavior() throws Exception {
|
||||
ScheduledExecutorService schedulerPool =
|
||||
Executors.newSingleThreadScheduledExecutor();
|
||||
ExecutorService workerPool = Executors.newFixedThreadPool(3);
|
||||
TriggerScheduler scheduler = new TriggerScheduler(
|
||||
new InMemoryTriggerStore(), schedulerPool, workerPool, 1000L);
|
||||
AtomicInteger executions = new AtomicInteger();
|
||||
ChainDefinition definition = new ChainDefinition();
|
||||
definition.setId("join-mode-retry-loop");
|
||||
StartNode start = new StartNode();
|
||||
start.setId("start");
|
||||
RetryLoopNode worker = new RetryLoopNode(executions);
|
||||
worker.setId("worker");
|
||||
worker.setRetryEnable(true);
|
||||
worker.setMaxRetryCount(1);
|
||||
worker.setRetryIntervalMs(0L);
|
||||
worker.setLoopEnable(true);
|
||||
worker.setMaxLoopCount(2);
|
||||
worker.setLoopIntervalMs(0L);
|
||||
EndNode end = endNode("worker", "count", "count");
|
||||
definition.addNode(start);
|
||||
definition.addNode(worker);
|
||||
definition.addNode(end);
|
||||
definition.addEdge(edge("start-worker", "start", "worker"));
|
||||
definition.addEdge(edge("worker-end", "worker", "end"));
|
||||
|
||||
ChainExecutor executor = new ChainExecutor(
|
||||
ignored -> definition,
|
||||
new InMemoryChainStateRepository(),
|
||||
new InMemoryNodeStateRepository(),
|
||||
scheduler);
|
||||
try {
|
||||
Map<String, Object> result = executor.execute(
|
||||
definition.getId(), Collections.emptyMap(), 5L, TimeUnit.SECONDS);
|
||||
|
||||
Assert.assertEquals(NodeJoinMode.ANY, worker.getJoinMode());
|
||||
Assert.assertEquals(3, executions.get());
|
||||
Assert.assertEquals(3, result.get("count"));
|
||||
} finally {
|
||||
scheduler.shutdown();
|
||||
}
|
||||
}
|
||||
|
||||
private JoinFixture createJoinFixture(
|
||||
NodeJoinMode joinMode, boolean withCondition) {
|
||||
JoinFixture fixture = new JoinFixture();
|
||||
fixture.schedulerPool = Executors.newScheduledThreadPool(2);
|
||||
fixture.workerPool = Executors.newFixedThreadPool(4);
|
||||
fixture.scheduler = new TriggerScheduler(
|
||||
new InMemoryTriggerStore(),
|
||||
fixture.schedulerPool,
|
||||
fixture.workerPool,
|
||||
1000L);
|
||||
fixture.chainStateRepository = new InMemoryChainStateRepository();
|
||||
fixture.nodeStateRepository = new InMemoryNodeStateRepository();
|
||||
fixture.definition = new ChainDefinition();
|
||||
fixture.definition.setId("join-mode-" + joinMode.getValue());
|
||||
|
||||
StartNode start = new StartNode();
|
||||
start.setId("start");
|
||||
BranchNode branchA = new BranchNode(
|
||||
"A", fixture.branchACompleted, null, null);
|
||||
branchA.setId("a");
|
||||
BranchNode branchB = new BranchNode(
|
||||
"B", null, fixture.branchBStarted, fixture.releaseBranchB);
|
||||
branchB.setId("b");
|
||||
ProbeJoinNode join = new ProbeJoinNode(fixture.joinExecutions);
|
||||
join.setId("join");
|
||||
join.setJoinMode(joinMode);
|
||||
if (withCondition) {
|
||||
join.setCondition(new BothOutputsCondition(fixture.conditionChecks));
|
||||
}
|
||||
EndNode end = endNode("join", "combined", "combined");
|
||||
end.addOutputDef(outputRef("sawBoth", "join.sawBoth"));
|
||||
|
||||
fixture.definition.addNode(start);
|
||||
fixture.definition.addNode(branchA);
|
||||
fixture.definition.addNode(branchB);
|
||||
fixture.definition.addNode(join);
|
||||
fixture.definition.addNode(end);
|
||||
fixture.definition.addEdge(edge("start-a", "start", "a"));
|
||||
fixture.definition.addEdge(edge("start-b", "start", "b"));
|
||||
fixture.definition.addEdge(edge("a-join", "a", "join"));
|
||||
fixture.definition.addEdge(edge("b-join", "b", "join"));
|
||||
fixture.definition.addEdge(edge("join-end", "join", "end"));
|
||||
|
||||
fixture.executor = new ChainExecutor(
|
||||
ignored -> fixture.definition,
|
||||
fixture.chainStateRepository,
|
||||
fixture.nodeStateRepository,
|
||||
fixture.scheduler);
|
||||
return fixture;
|
||||
}
|
||||
|
||||
private Node parseEndNode(
|
||||
ChainParser parser, String joinMode, String parentId) {
|
||||
JSONObject data = new JSONObject();
|
||||
if (joinMode != null) {
|
||||
data.put("joinMode", joinMode);
|
||||
}
|
||||
JSONObject nodeJson = new JSONObject();
|
||||
nodeJson.put("id", "end");
|
||||
nodeJson.put("type", "endNode");
|
||||
nodeJson.put("parentId", parentId);
|
||||
nodeJson.put("data", data);
|
||||
return new EndNodeParser().parse(
|
||||
nodeJson, new JSONObject(), parser);
|
||||
}
|
||||
|
||||
private void assertInvalidJoinMode(Runnable action) {
|
||||
try {
|
||||
action.run();
|
||||
Assert.fail("Expected invalid join mode");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertTrue(expected.getMessage().contains("joinMode"));
|
||||
}
|
||||
}
|
||||
|
||||
private static boolean hasTriggerEdge(
|
||||
InMemoryNodeStateRepository repository,
|
||||
String instanceId,
|
||||
String nodeId,
|
||||
String edgeId) {
|
||||
NodeState state = repository.load(instanceId, nodeId);
|
||||
return state != null && state.getTriggerEdgeIds().contains(edgeId);
|
||||
}
|
||||
|
||||
private static ChainState awaitTerminal(
|
||||
InMemoryChainStateRepository repository,
|
||||
String instanceId) throws Exception {
|
||||
await(() -> {
|
||||
ChainState state = repository.load(instanceId);
|
||||
return state != null
|
||||
&& state.getStatus() != null
|
||||
&& state.getStatus().isTerminal();
|
||||
});
|
||||
ChainState state = repository.load(instanceId);
|
||||
Assert.assertEquals(ChainStatus.SUCCEEDED, state.getStatus());
|
||||
return state;
|
||||
}
|
||||
|
||||
private static void await(BooleanSupplier condition) throws Exception {
|
||||
long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(3L);
|
||||
while (!condition.getAsBoolean() && System.nanoTime() < deadline) {
|
||||
Thread.sleep(10L);
|
||||
}
|
||||
Assert.assertTrue("condition was not met before timeout", condition.getAsBoolean());
|
||||
}
|
||||
|
||||
private static EndNode endNode(
|
||||
String sourceNodeId, String sourceName, String outputName) {
|
||||
EndNode end = new EndNode();
|
||||
end.setId("end");
|
||||
end.addOutputDef(outputRef(
|
||||
outputName, sourceNodeId + "." + sourceName));
|
||||
return end;
|
||||
}
|
||||
|
||||
private static Parameter outputRef(String name, String ref) {
|
||||
Parameter parameter = new Parameter();
|
||||
parameter.setName(name);
|
||||
parameter.setRef(ref);
|
||||
parameter.setRefType(RefType.REF);
|
||||
return parameter;
|
||||
}
|
||||
|
||||
private static Edge edge(String id, String source, String target) {
|
||||
Edge edge = new Edge();
|
||||
edge.setId(id);
|
||||
edge.setSource(source);
|
||||
edge.setTarget(target);
|
||||
return edge;
|
||||
}
|
||||
|
||||
private static final class BranchNode extends BaseNode {
|
||||
private final String value;
|
||||
private final CountDownLatch completed;
|
||||
private final CountDownLatch started;
|
||||
private final CountDownLatch release;
|
||||
|
||||
private BranchNode(
|
||||
String value,
|
||||
CountDownLatch completed,
|
||||
CountDownLatch started,
|
||||
CountDownLatch release) {
|
||||
this.value = value;
|
||||
this.completed = completed;
|
||||
this.started = started;
|
||||
this.release = release;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Map<String, Object> execute(Chain chain) {
|
||||
if (started != null) {
|
||||
started.countDown();
|
||||
}
|
||||
if (release != null) {
|
||||
try {
|
||||
if (!release.await(3L, TimeUnit.SECONDS)) {
|
||||
throw new IllegalStateException("branch release timed out");
|
||||
}
|
||||
} catch (InterruptedException error) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException("branch interrupted", error);
|
||||
}
|
||||
}
|
||||
if (completed != null) {
|
||||
completed.countDown();
|
||||
}
|
||||
return Collections.singletonMap("value", value);
|
||||
}
|
||||
}
|
||||
|
||||
private static final class ProbeJoinNode extends BaseNode {
|
||||
private final AtomicInteger executions;
|
||||
|
||||
private ProbeJoinNode(AtomicInteger executions) {
|
||||
this.executions = executions;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Map<String, Object> execute(Chain chain) {
|
||||
executions.incrementAndGet();
|
||||
Object a = chain.getExecutionState().getMemory().get("a.value");
|
||||
Object b = chain.getExecutionState().getMemory().get("b.value");
|
||||
Map<String, Object> result = new HashMap<>();
|
||||
result.put("combined", String.valueOf(a) + "+" + String.valueOf(b));
|
||||
result.put("sawBoth", a != null && b != null);
|
||||
return result;
|
||||
}
|
||||
}
|
||||
|
||||
private static final class BothOutputsCondition implements NodeCondition {
|
||||
private static final long serialVersionUID = 1L;
|
||||
private final AtomicInteger checks;
|
||||
|
||||
private BothOutputsCondition(AtomicInteger checks) {
|
||||
this.checks = checks;
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean check(
|
||||
Chain chain,
|
||||
NodeState context,
|
||||
Map<String, Object> executeResult) {
|
||||
checks.incrementAndGet();
|
||||
Map<String, Object> memory = chain.getExecutionState().getMemory();
|
||||
return memory.containsKey("a.value") && memory.containsKey("b.value");
|
||||
}
|
||||
}
|
||||
|
||||
private static final class RetryLoopNode extends BaseNode {
|
||||
private final AtomicInteger executions;
|
||||
|
||||
private RetryLoopNode(AtomicInteger executions) {
|
||||
this.executions = executions;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Map<String, Object> execute(Chain chain) {
|
||||
int count = executions.incrementAndGet();
|
||||
if (count == 1) {
|
||||
throw new IllegalStateException("retry once");
|
||||
}
|
||||
return Collections.singletonMap("count", count);
|
||||
}
|
||||
}
|
||||
|
||||
private static final class JoinFixture {
|
||||
private final CountDownLatch branchACompleted = new CountDownLatch(1);
|
||||
private final CountDownLatch branchBStarted = new CountDownLatch(1);
|
||||
private final CountDownLatch releaseBranchB = new CountDownLatch(1);
|
||||
private final AtomicInteger joinExecutions = new AtomicInteger();
|
||||
private final AtomicInteger conditionChecks = new AtomicInteger();
|
||||
private ScheduledExecutorService schedulerPool;
|
||||
private ExecutorService workerPool;
|
||||
private TriggerScheduler scheduler;
|
||||
private InMemoryChainStateRepository chainStateRepository;
|
||||
private InMemoryNodeStateRepository nodeStateRepository;
|
||||
private ChainDefinition definition;
|
||||
private ChainExecutor executor;
|
||||
}
|
||||
}
|
||||
@@ -2,14 +2,26 @@ package com.easyagents.rag.retrieval;
|
||||
|
||||
import com.easyagents.core.document.Document;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
|
||||
/**
|
||||
* 将不同检索路径的原始分数转换为统一的零到一最终相关度。
|
||||
*/
|
||||
public final class RagScoreNormalizer {
|
||||
|
||||
/**
|
||||
* 禁止实例化工具类。
|
||||
*/
|
||||
private RagScoreNormalizer() {
|
||||
}
|
||||
|
||||
/**
|
||||
* 按检索模式归一化文档最终分数。
|
||||
*
|
||||
* @param documents 待归一化文档
|
||||
* @param retrievalMode 检索模式
|
||||
* @param reranked 是否已经过重排模型
|
||||
*/
|
||||
public static void normalize(List<Document> documents, RetrievalMode retrievalMode, boolean reranked) {
|
||||
if (documents == null || documents.isEmpty()) {
|
||||
return;
|
||||
@@ -49,41 +61,29 @@ public final class RagScoreNormalizer {
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 保留重排模型返回的绝对相关度,并将异常范围限制到零到一。
|
||||
*
|
||||
* @param documents 待归一化的文档
|
||||
*/
|
||||
private static void normalizeRerankScores(List<Document> documents) {
|
||||
List<Double> rawScores = new ArrayList<Double>(documents.size());
|
||||
boolean allPresent = true;
|
||||
Double min = null;
|
||||
Double max = null;
|
||||
for (Document document : documents) {
|
||||
Double rawScore = readRawScore(document, RagRetrievalMetadataKeys.RERANK_SCORE, document == null ? null : document.getScore());
|
||||
rawScores.add(rawScore);
|
||||
if (rawScore == null) {
|
||||
allPresent = false;
|
||||
continue;
|
||||
if (document != null) {
|
||||
// Rerank 适配器返回绝对相关度,按查询结果集再次缩放会把低相关第一名错误抬高到 1。
|
||||
document.setScore(clamp01(rawScore));
|
||||
}
|
||||
min = min == null ? rawScore : Math.min(min, rawScore);
|
||||
max = max == null ? rawScore : Math.max(max, rawScore);
|
||||
}
|
||||
|
||||
if (allPresent && min != null && max != null && Double.compare(max, min) != 0) {
|
||||
for (int i = 0; i < documents.size(); i++) {
|
||||
Double rawScore = rawScores.get(i);
|
||||
documents.get(i).setScore(clamp01((rawScore - min) / (max - min)));
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
if (documents.size() == 1) {
|
||||
documents.get(0).setScore(1D);
|
||||
return;
|
||||
}
|
||||
|
||||
int size = documents.size();
|
||||
for (int i = 0; i < size; i++) {
|
||||
documents.get(i).setScore(clamp01(1D - ((double) i / (double) (size - 1))));
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 优先读取指定元数据中的原始分数。
|
||||
*
|
||||
* @param document 文档
|
||||
* @param metadataKey 分数元数据键
|
||||
* @param fallback 元数据不可用时的回退分数
|
||||
* @return 原始分数;文档为空时返回 null
|
||||
*/
|
||||
private static Double readRawScore(Document document, String metadataKey, Double fallback) {
|
||||
if (document == null) {
|
||||
return null;
|
||||
@@ -102,6 +102,12 @@ public final class RagScoreNormalizer {
|
||||
return fallback;
|
||||
}
|
||||
|
||||
/**
|
||||
* 将可空分数限制到零到一范围。
|
||||
*
|
||||
* @param value 原始分数
|
||||
* @return 有效最终分数
|
||||
*/
|
||||
private static double clamp01(Double value) {
|
||||
if (value == null || value.isNaN() || value.isInfinite()) {
|
||||
return 0D;
|
||||
|
||||
@@ -7,8 +7,14 @@ import org.junit.Test;
|
||||
import java.util.Arrays;
|
||||
import java.util.List;
|
||||
|
||||
/**
|
||||
* {@link RagScoreNormalizer} 回归测试。
|
||||
*/
|
||||
public class RagScoreNormalizerTest {
|
||||
|
||||
/**
|
||||
* 验证关键词分数按有界函数归一化。
|
||||
*/
|
||||
@Test
|
||||
public void shouldNormalizeKeywordScoresToZeroAndOneRange() {
|
||||
Document first = document(1, 9D, RagRetrievalMetadataKeys.KEYWORD_SCORE);
|
||||
@@ -20,6 +26,9 @@ public class RagScoreNormalizerTest {
|
||||
Assert.assertEquals(0D, second.getScore(), 0.0001D);
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证混合检索 RRF 分数按理论上界归一化。
|
||||
*/
|
||||
@Test
|
||||
public void shouldNormalizeHybridFusionScoreByRrfUpperBound() {
|
||||
Document document = document(1, 2D / (RrfFusionStrategy.DEFAULT_RRF_K + 1D), RagRetrievalMetadataKeys.FUSION_SCORE);
|
||||
@@ -29,36 +38,60 @@ public class RagScoreNormalizerTest {
|
||||
Assert.assertEquals(1D, document.getScore(), 0.0001D);
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证重排模型返回的绝对相关度保持不变。
|
||||
*/
|
||||
@Test
|
||||
public void shouldNormalizeRerankScoresByMinMax() {
|
||||
public void shouldPreserveRerankRelevanceScores() {
|
||||
List<Document> documents = Arrays.asList(
|
||||
document(1, 10D, RagRetrievalMetadataKeys.RERANK_SCORE),
|
||||
document(2, 20D, RagRetrievalMetadataKeys.RERANK_SCORE),
|
||||
document(3, 30D, RagRetrievalMetadataKeys.RERANK_SCORE)
|
||||
document(1, 0.1D, RagRetrievalMetadataKeys.RERANK_SCORE),
|
||||
document(2, 0.5D, RagRetrievalMetadataKeys.RERANK_SCORE),
|
||||
document(3, 0.9D, RagRetrievalMetadataKeys.RERANK_SCORE)
|
||||
);
|
||||
|
||||
RagScoreNormalizer.normalize(documents, RetrievalMode.HYBRID, true);
|
||||
|
||||
Assert.assertEquals(0.1D, documents.get(0).getScore(), 0.0001D);
|
||||
Assert.assertEquals(0.5D, documents.get(1).getScore(), 0.0001D);
|
||||
Assert.assertEquals(0.9D, documents.get(2).getScore(), 0.0001D);
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证单条低分重排结果不会被抬高为高相关结果。
|
||||
*/
|
||||
@Test
|
||||
public void shouldKeepSingleLowRerankScoreLow() {
|
||||
Document document = document(1, 0.1D, RagRetrievalMetadataKeys.RERANK_SCORE);
|
||||
|
||||
RagScoreNormalizer.normalize(Arrays.asList(document), RetrievalMode.HYBRID, true);
|
||||
|
||||
Assert.assertEquals(0.1D, document.getScore(), 0.0001D);
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证越界重排分数会被限制到零到一范围。
|
||||
*/
|
||||
@Test
|
||||
public void shouldClampRerankScoresToZeroAndOneRange() {
|
||||
List<Document> documents = Arrays.asList(
|
||||
document(1, -0.2D, RagRetrievalMetadataKeys.RERANK_SCORE),
|
||||
document(2, 1.2D, RagRetrievalMetadataKeys.RERANK_SCORE)
|
||||
);
|
||||
|
||||
RagScoreNormalizer.normalize(documents, RetrievalMode.HYBRID, true);
|
||||
|
||||
Assert.assertEquals(0D, documents.get(0).getScore(), 0.0001D);
|
||||
Assert.assertEquals(0.5D, documents.get(1).getScore(), 0.0001D);
|
||||
Assert.assertEquals(1D, documents.get(2).getScore(), 0.0001D);
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldFallbackToRankBasedNormalizationWhenRerankScoresAreEqual() {
|
||||
List<Document> documents = Arrays.asList(
|
||||
document(1, 5D, RagRetrievalMetadataKeys.RERANK_SCORE),
|
||||
document(2, 5D, RagRetrievalMetadataKeys.RERANK_SCORE),
|
||||
document(3, 5D, RagRetrievalMetadataKeys.RERANK_SCORE)
|
||||
);
|
||||
|
||||
RagScoreNormalizer.normalize(documents, RetrievalMode.HYBRID, true);
|
||||
|
||||
Assert.assertEquals(1D, documents.get(0).getScore(), 0.0001D);
|
||||
Assert.assertEquals(0.5D, documents.get(1).getScore(), 0.0001D);
|
||||
Assert.assertEquals(0D, documents.get(2).getScore(), 0.0001D);
|
||||
Assert.assertEquals(1D, documents.get(1).getScore(), 0.0001D);
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建携带原始分数元数据的测试文档。
|
||||
*
|
||||
* @param id 文档 ID
|
||||
* @param score 原始分数
|
||||
* @param metadataKey 分数元数据键
|
||||
* @return 测试文档
|
||||
*/
|
||||
private Document document(Object id, Double score, String metadataKey) {
|
||||
Document document = new Document();
|
||||
document.setId(id);
|
||||
|
||||
@@ -199,10 +199,13 @@ public class LuceneSearcher implements DocumentSearcher, AutoCloseable {
|
||||
@Override
|
||||
public List<Document> searchDocuments(KeywordSearchRequest request) {
|
||||
List<Document> results = new ArrayList<>();
|
||||
if (request == null || request.getKeyword() == null || request.getKeyword().trim().isEmpty()) {
|
||||
return results;
|
||||
}
|
||||
try (IndexReader reader = DirectoryReader.open(directory)) {
|
||||
IndexSearcher searcher = new IndexSearcher(reader);
|
||||
Query query = buildQuery(request);
|
||||
TopDocs topDocs = searcher.search(query, request == null ? 10 : request.getCount());
|
||||
TopDocs topDocs = searcher.search(query, request.getCount());
|
||||
for (ScoreDoc scoreDoc : topDocs.scoreDocs) {
|
||||
org.apache.lucene.document.Document doc = searcher.doc(scoreDoc.doc);
|
||||
Document resultDoc = new Document();
|
||||
@@ -224,29 +227,24 @@ public class LuceneSearcher implements DocumentSearcher, AutoCloseable {
|
||||
return results;
|
||||
}
|
||||
|
||||
Query buildQuery(KeywordSearchRequest request) {
|
||||
try {
|
||||
String keyword = request == null ? null : request.getKeyword();
|
||||
Query buildQuery(KeywordSearchRequest request) throws ParseException {
|
||||
String escapedKeyword = QueryParser.escape(request.getKeyword());
|
||||
|
||||
QueryParser titleQueryParser = new QueryParser("title", analyzer);
|
||||
Query titleQuery = titleQueryParser.parse(keyword);
|
||||
BooleanClause titleBooleanClause = new BooleanClause(titleQuery, BooleanClause.Occur.SHOULD);
|
||||
QueryParser titleQueryParser = new QueryParser("title", analyzer);
|
||||
Query titleQuery = titleQueryParser.parse(escapedKeyword);
|
||||
BooleanClause titleBooleanClause = new BooleanClause(titleQuery, BooleanClause.Occur.SHOULD);
|
||||
|
||||
QueryParser contentQueryParser = new QueryParser("content", analyzer);
|
||||
Query contentQuery = contentQueryParser.parse(keyword);
|
||||
BooleanClause contentBooleanClause = new BooleanClause(contentQuery, BooleanClause.Occur.SHOULD);
|
||||
QueryParser contentQueryParser = new QueryParser("content", analyzer);
|
||||
Query contentQuery = contentQueryParser.parse(escapedKeyword);
|
||||
BooleanClause contentBooleanClause = new BooleanClause(contentQuery, BooleanClause.Occur.SHOULD);
|
||||
|
||||
BooleanQuery.Builder builder = new BooleanQuery.Builder();
|
||||
builder.add(titleBooleanClause)
|
||||
.add(contentBooleanClause);
|
||||
if (request != null && request.getKnowledgeId() != null && !request.getKnowledgeId().trim().isEmpty()) {
|
||||
builder.add(new TermQuery(new Term(KeywordSearchMetadataKeys.KNOWLEDGE_ID, request.getKnowledgeId().trim())), BooleanClause.Occur.MUST);
|
||||
}
|
||||
return builder.build();
|
||||
} catch (ParseException e) {
|
||||
LOG.error(e.toString(), e);
|
||||
BooleanQuery.Builder builder = new BooleanQuery.Builder();
|
||||
builder.add(titleBooleanClause)
|
||||
.add(contentBooleanClause);
|
||||
if (request.getKnowledgeId() != null && !request.getKnowledgeId().trim().isEmpty()) {
|
||||
builder.add(new TermQuery(new Term(KeywordSearchMetadataKeys.KNOWLEDGE_ID, request.getKnowledgeId().trim())), BooleanClause.Occur.MUST);
|
||||
}
|
||||
return null;
|
||||
return builder.build();
|
||||
}
|
||||
|
||||
private static Analyzer createAnalyzer() {
|
||||
|
||||
@@ -57,6 +57,45 @@ public class LuceneSearcherTest {
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证用户输入中的 Lucene 特殊字符按普通文本检索。
|
||||
*
|
||||
* @throws Exception 临时目录或 Lucene 资源操作失败时抛出
|
||||
*/
|
||||
@Test
|
||||
public void shouldTreatLuceneSpecialCharactersAsPlainText() throws Exception {
|
||||
Path tempDir = Files.createTempDirectory("lucene-searcher-special-character-test");
|
||||
LuceneConfig config = new LuceneConfig();
|
||||
config.setIndexDirPath(tempDir.toString());
|
||||
try (LuceneSearcher searcher = new LuceneSearcher(config)) {
|
||||
Document document = new Document();
|
||||
document.setId("special-character");
|
||||
document.setContent("A/C.\\nPLEASE");
|
||||
Assert.assertTrue(searcher.addDocument(document));
|
||||
|
||||
List<Document> results = searcher.searchDocuments("A/C.\\nPLEASE", 10);
|
||||
|
||||
Assert.assertEquals(1, results.size());
|
||||
Assert.assertEquals("special-character", String.valueOf(results.get(0).getId()));
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证空查询直接返回空结果。
|
||||
*
|
||||
* @throws Exception 临时目录或 Lucene 资源操作失败时抛出
|
||||
*/
|
||||
@Test
|
||||
public void shouldReturnEmptyForMissingKeyword() throws Exception {
|
||||
Path tempDir = Files.createTempDirectory("lucene-searcher-empty-keyword-test");
|
||||
LuceneConfig config = new LuceneConfig();
|
||||
config.setIndexDirPath(tempDir.toString());
|
||||
try (LuceneSearcher searcher = new LuceneSearcher(config)) {
|
||||
Assert.assertTrue(searcher.searchDocuments((KeywordSearchRequest) null).isEmpty());
|
||||
Assert.assertTrue(searcher.searchDocuments(" ", 10).isEmpty());
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 验证多个导入线程可共享同一个 IndexWriter 完成批量写入。
|
||||
*
|
||||
|
||||
@@ -25,7 +25,6 @@
|
||||
<dependency>
|
||||
<groupId>io.milvus</groupId>
|
||||
<artifactId>milvus-sdk-java</artifactId>
|
||||
<version>2.4.1</version>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>junit</groupId>
|
||||
|
||||
@@ -0,0 +1,617 @@
|
||||
/*
|
||||
* Copyright (c) 2023-2026, Easy-Agents (fuhai999@gmail.com).
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
*/
|
||||
package com.easyagents.store.milvus;
|
||||
|
||||
import com.easyagents.core.util.StringUtil;
|
||||
import io.grpc.Context;
|
||||
import io.milvus.pool.MilvusClientV2Pool;
|
||||
import io.milvus.pool.PoolConfig;
|
||||
import io.milvus.v2.client.ConnectConfig;
|
||||
import io.milvus.v2.client.MilvusClientV2;
|
||||
import io.milvus.v2.client.RetryConfig;
|
||||
|
||||
import java.net.URI;
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.security.MessageDigest;
|
||||
import java.security.NoSuchAlgorithmException;
|
||||
import java.time.Duration;
|
||||
import java.util.Collections;
|
||||
import java.util.HashSet;
|
||||
import java.util.Set;
|
||||
import java.util.concurrent.Callable;
|
||||
import java.util.concurrent.CancellationException;
|
||||
import java.util.concurrent.CompletableFuture;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
import java.util.concurrent.ConcurrentMap;
|
||||
import java.util.concurrent.locks.ReentrantReadWriteLock;
|
||||
import java.util.function.Function;
|
||||
|
||||
/**
|
||||
* Shared Milvus client pool and collection state.
|
||||
*/
|
||||
public class MilvusClientManager implements AutoCloseable {
|
||||
|
||||
private static final String POOL_KEY = "default";
|
||||
private static final RetryConfig SINGLE_ATTEMPT_RETRY_CONFIG = RetryConfig.builder()
|
||||
.maxRetryTimes(1)
|
||||
.retryOnRateLimit(false)
|
||||
.maxRetryTimeoutMs(0L)
|
||||
.build();
|
||||
|
||||
private final ReentrantReadWriteLock lifecycleLock = new ReentrantReadWriteLock();
|
||||
private final Set<String> initializedCollections =
|
||||
Collections.synchronizedSet(new HashSet<String>());
|
||||
private final Set<String> loadedCollections =
|
||||
Collections.synchronizedSet(new HashSet<String>());
|
||||
private final ConcurrentMap<String, CollectionLoadTicket> collectionLoads =
|
||||
new ConcurrentHashMap<String, CollectionLoadTicket>();
|
||||
private final Set<Context.CancellableContext> activeContexts =
|
||||
ConcurrentHashMap.newKeySet();
|
||||
private final ConcurrentMap<Thread, Integer> activeOperations =
|
||||
new ConcurrentHashMap<Thread, Integer>();
|
||||
private volatile ManagedMilvusClientV2Pool pool;
|
||||
private volatile String poolFingerprint;
|
||||
private volatile long poolGeneration;
|
||||
private volatile boolean acceptingOperations = true;
|
||||
private volatile boolean closed;
|
||||
|
||||
public MilvusClientManager(MilvusVectorStoreConfig config) {
|
||||
PoolSettings settings = PoolSettings.from(config);
|
||||
this.pool = createPool(settings);
|
||||
this.poolFingerprint = fingerprint(settings);
|
||||
}
|
||||
|
||||
private static ManagedMilvusClientV2Pool createPool(PoolSettings settings) {
|
||||
ConnectConfig connectConfig = buildConnectConfig(settings);
|
||||
PoolConfig poolConfig = PoolConfig.builder()
|
||||
.maxTotal(settings.poolMaxTotal())
|
||||
.maxTotalPerKey(settings.poolMaxTotalPerKey())
|
||||
.maxIdlePerKey(settings.poolMaxIdlePerKey())
|
||||
.minIdlePerKey(settings.poolMinIdlePerKey())
|
||||
.blockWhenExhausted(true)
|
||||
.maxBlockWaitDuration(Duration.ofMillis(settings.poolMaxWaitMillis()))
|
||||
.evictionPollingInterval(Duration.ofMillis(settings.poolEvictionIntervalMillis()))
|
||||
.minEvictableIdleDuration(Duration.ofMillis(settings.poolMinEvictableIdleMillis()))
|
||||
.testOnBorrow(true)
|
||||
.testOnReturn(false)
|
||||
.build();
|
||||
try {
|
||||
return new ManagedMilvusClientV2Pool(poolConfig, connectConfig);
|
||||
} catch (ReflectiveOperationException exception) {
|
||||
throw new IllegalStateException("Unable to initialize Milvus client pool", exception);
|
||||
}
|
||||
}
|
||||
|
||||
public <T> T withClient(Function<MilvusClientV2, T> operation) {
|
||||
return withClient(null, operation);
|
||||
}
|
||||
|
||||
public <T> T withClient(
|
||||
Duration maxWait,
|
||||
Function<MilvusClientV2, T> operation
|
||||
) {
|
||||
Thread operationThread = registerActiveOperation();
|
||||
lifecycleLock.readLock().lock();
|
||||
try {
|
||||
ManagedMilvusClientV2Pool currentPool = requireOpenPool();
|
||||
MilvusClientV2 client = maxWait == null
|
||||
? currentPool.getClient(POOL_KEY)
|
||||
: currentPool.getClient(POOL_KEY, maxWait);
|
||||
if (client == null) {
|
||||
throw new IllegalStateException(
|
||||
"Milvus client pool is exhausted or unavailable"
|
||||
);
|
||||
}
|
||||
Throwable operationFailure = null;
|
||||
try {
|
||||
client.retryConfig(SINGLE_ATTEMPT_RETRY_CONFIG);
|
||||
Context.CancellableContext operationContext =
|
||||
Context.current().withCancellation();
|
||||
try {
|
||||
return withRequestContext(operationContext,
|
||||
() -> operation.apply(client));
|
||||
} catch (RuntimeException | Error exception) {
|
||||
throw exception;
|
||||
} catch (Exception exception) {
|
||||
throw new IllegalStateException(
|
||||
"Milvus client operation failed", exception);
|
||||
} finally {
|
||||
operationContext.cancel(null);
|
||||
}
|
||||
} catch (RuntimeException | Error exception) {
|
||||
operationFailure = exception;
|
||||
throw exception;
|
||||
} finally {
|
||||
RuntimeException cleanupFailure = null;
|
||||
try {
|
||||
releaseClient(currentPool, client);
|
||||
} catch (RuntimeException exception) {
|
||||
cleanupFailure = exception;
|
||||
}
|
||||
if (cleanupFailure != null) {
|
||||
if (operationFailure == null) {
|
||||
throw cleanupFailure;
|
||||
}
|
||||
operationFailure.addSuppressed(cleanupFailure);
|
||||
}
|
||||
}
|
||||
} finally {
|
||||
unregisterActiveOperation(operationThread);
|
||||
lifecycleLock.readLock().unlock();
|
||||
}
|
||||
}
|
||||
|
||||
private Thread registerActiveOperation() {
|
||||
ensureAcceptingOperations();
|
||||
Thread currentThread = Thread.currentThread();
|
||||
activeOperations.merge(currentThread, 1, Integer::sum);
|
||||
if (!acceptingOperations) {
|
||||
unregisterActiveOperation(currentThread);
|
||||
ensureAcceptingOperations();
|
||||
}
|
||||
return currentThread;
|
||||
}
|
||||
|
||||
private void unregisterActiveOperation(Thread operationThread) {
|
||||
activeOperations.computeIfPresent(operationThread,
|
||||
(thread, depth) -> depth <= 1 ? null : depth - 1);
|
||||
}
|
||||
|
||||
private void releaseClient(
|
||||
ManagedMilvusClientV2Pool currentPool,
|
||||
MilvusClientV2 client
|
||||
) {
|
||||
RuntimeException readinessFailure = null;
|
||||
boolean reusable = false;
|
||||
try {
|
||||
reusable = client.clientIsReady();
|
||||
} catch (RuntimeException exception) {
|
||||
readinessFailure = exception;
|
||||
}
|
||||
try {
|
||||
if (reusable) {
|
||||
currentPool.returnClient(POOL_KEY, client);
|
||||
} else {
|
||||
discardFailedClient(currentPool, client);
|
||||
}
|
||||
} catch (RuntimeException cleanupFailure) {
|
||||
if (readinessFailure == null) {
|
||||
throw cleanupFailure;
|
||||
}
|
||||
readinessFailure.addSuppressed(cleanupFailure);
|
||||
}
|
||||
if (readinessFailure != null) {
|
||||
throw readinessFailure;
|
||||
}
|
||||
}
|
||||
|
||||
<T> T withRequestContext(
|
||||
Context.CancellableContext context,
|
||||
Callable<T> operation
|
||||
) throws Exception {
|
||||
ensureAcceptingOperations();
|
||||
activeContexts.add(context);
|
||||
if (!acceptingOperations) {
|
||||
activeContexts.remove(context);
|
||||
context.cancel(new CancellationException(
|
||||
"Milvus client pool is unavailable"));
|
||||
ensureAcceptingOperations();
|
||||
}
|
||||
try {
|
||||
return context.call(operation);
|
||||
} finally {
|
||||
activeContexts.remove(context);
|
||||
}
|
||||
}
|
||||
|
||||
private void discardFailedClient(
|
||||
ManagedMilvusClientV2Pool currentPool,
|
||||
MilvusClientV2 client
|
||||
) {
|
||||
RuntimeException cleanupFailure = null;
|
||||
try {
|
||||
currentPool.invalidateClient(POOL_KEY, client);
|
||||
} catch (RuntimeException exception) {
|
||||
cleanupFailure = exception;
|
||||
}
|
||||
if (cleanupFailure != null) {
|
||||
throw cleanupFailure;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Rebuilds the pool when connection or pool settings change.
|
||||
* Active operations are cancelled before the old pool is closed.
|
||||
*
|
||||
* @return true when a new pool was installed
|
||||
*/
|
||||
public synchronized boolean reconfigureIfNeeded(MilvusVectorStoreConfig config) {
|
||||
PoolSettings nextSettings = PoolSettings.from(config);
|
||||
String nextFingerprint = fingerprint(nextSettings);
|
||||
if (nextFingerprint.equals(poolFingerprint)) {
|
||||
return false;
|
||||
}
|
||||
acceptingOperations = false;
|
||||
cancelActiveContexts("Milvus client pool is reconfiguring");
|
||||
interruptActiveOperations();
|
||||
lifecycleLock.writeLock().lock();
|
||||
try {
|
||||
if (closed) {
|
||||
throw new IllegalStateException("Milvus client pool is closed");
|
||||
}
|
||||
if (nextFingerprint.equals(poolFingerprint)) {
|
||||
return false;
|
||||
}
|
||||
ManagedMilvusClientV2Pool replacement = createPool(nextSettings);
|
||||
ManagedMilvusClientV2Pool previous = pool;
|
||||
pool = replacement;
|
||||
poolFingerprint = nextFingerprint;
|
||||
poolGeneration++;
|
||||
initializedCollections.clear();
|
||||
loadedCollections.clear();
|
||||
failCollectionLoads("Milvus client pool was reconfigured");
|
||||
if (previous != null) {
|
||||
previous.close();
|
||||
}
|
||||
return true;
|
||||
} finally {
|
||||
lifecycleLock.writeLock().unlock();
|
||||
if (!closed) {
|
||||
acceptingOperations = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
boolean isCollectionInitialized(String collectionName) {
|
||||
return initializedCollections.contains(collectionName);
|
||||
}
|
||||
|
||||
Object initializedCollectionsLock() {
|
||||
return initializedCollections;
|
||||
}
|
||||
|
||||
void markCollectionInitialized(String collectionName) {
|
||||
initializedCollections.add(collectionName);
|
||||
}
|
||||
|
||||
boolean isCollectionLoaded(String collectionName) {
|
||||
return loadedCollections.contains(collectionName);
|
||||
}
|
||||
|
||||
void markCollectionLoaded(String collectionName) {
|
||||
loadedCollections.add(collectionName);
|
||||
}
|
||||
|
||||
void markCollectionUnloaded(String collectionName) {
|
||||
loadedCollections.remove(collectionName);
|
||||
}
|
||||
|
||||
CollectionLoadTicket beginCollectionLoad(String collectionName) {
|
||||
ensureAcceptingOperations();
|
||||
lifecycleLock.readLock().lock();
|
||||
try {
|
||||
requireOpenPool();
|
||||
CollectionLoadTicket candidate = new CollectionLoadTicket(
|
||||
collectionName,
|
||||
poolGeneration,
|
||||
new CompletableFuture<Void>(),
|
||||
true
|
||||
);
|
||||
CollectionLoadTicket existing = collectionLoads.putIfAbsent(
|
||||
collectionName, candidate);
|
||||
return existing == null ? candidate : existing.asFollower();
|
||||
} finally {
|
||||
lifecycleLock.readLock().unlock();
|
||||
}
|
||||
}
|
||||
|
||||
void completeCollectionLoad(CollectionLoadTicket ticket) {
|
||||
lifecycleLock.readLock().lock();
|
||||
try {
|
||||
requireOpenPool();
|
||||
if (ticket.generation != poolGeneration) {
|
||||
throw new IllegalStateException(
|
||||
"Milvus client pool changed while loading collection: "
|
||||
+ ticket.collectionName
|
||||
);
|
||||
}
|
||||
loadedCollections.add(ticket.collectionName);
|
||||
ticket.completion.complete(null);
|
||||
} finally {
|
||||
lifecycleLock.readLock().unlock();
|
||||
}
|
||||
}
|
||||
|
||||
void failCollectionLoad(
|
||||
CollectionLoadTicket ticket,
|
||||
Throwable failure,
|
||||
boolean retryableForFollowers
|
||||
) {
|
||||
if (retryableForFollowers && ticket.leader) {
|
||||
collectionLoads.remove(ticket.collectionName, ticket);
|
||||
}
|
||||
Throwable sharedFailure = retryableForFollowers
|
||||
? new RetryableCollectionLoadException(failure)
|
||||
: failure;
|
||||
ticket.completion.completeExceptionally(sharedFailure);
|
||||
}
|
||||
|
||||
void endCollectionLoad(CollectionLoadTicket ticket) {
|
||||
if (ticket.leader) {
|
||||
collectionLoads.remove(ticket.collectionName, ticket);
|
||||
}
|
||||
}
|
||||
|
||||
public int getActiveClientCount() {
|
||||
lifecycleLock.readLock().lock();
|
||||
try {
|
||||
return requireOpenPool().getTotalActiveClientNumber();
|
||||
} finally {
|
||||
lifecycleLock.readLock().unlock();
|
||||
}
|
||||
}
|
||||
|
||||
public int getIdleClientCount() {
|
||||
lifecycleLock.readLock().lock();
|
||||
try {
|
||||
return requireOpenPool().getTotalIdleClientNumber();
|
||||
} finally {
|
||||
lifecycleLock.readLock().unlock();
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
public synchronized void close() {
|
||||
if (closed) {
|
||||
return;
|
||||
}
|
||||
acceptingOperations = false;
|
||||
closed = true;
|
||||
cancelActiveContexts("Milvus client pool is closing");
|
||||
interruptActiveOperations();
|
||||
lifecycleLock.writeLock().lock();
|
||||
try {
|
||||
initializedCollections.clear();
|
||||
loadedCollections.clear();
|
||||
poolGeneration++;
|
||||
failCollectionLoads("Milvus client pool was closed");
|
||||
ManagedMilvusClientV2Pool currentPool = pool;
|
||||
pool = null;
|
||||
poolFingerprint = null;
|
||||
if (currentPool != null) {
|
||||
currentPool.close();
|
||||
}
|
||||
} finally {
|
||||
lifecycleLock.writeLock().unlock();
|
||||
}
|
||||
}
|
||||
|
||||
private ManagedMilvusClientV2Pool requireOpenPool() {
|
||||
ManagedMilvusClientV2Pool currentPool = pool;
|
||||
if (closed || currentPool == null) {
|
||||
throw new IllegalStateException("Milvus client pool is closed");
|
||||
}
|
||||
return currentPool;
|
||||
}
|
||||
|
||||
boolean isClosed() {
|
||||
return closed;
|
||||
}
|
||||
|
||||
private void ensureAcceptingOperations() {
|
||||
if (!acceptingOperations) {
|
||||
throw new IllegalStateException(closed
|
||||
? "Milvus client pool is closed"
|
||||
: "Milvus client pool is reconfiguring");
|
||||
}
|
||||
}
|
||||
|
||||
private void failCollectionLoads(String message) {
|
||||
IllegalStateException failure = new IllegalStateException(message);
|
||||
for (CollectionLoadTicket ticket : collectionLoads.values()) {
|
||||
ticket.completion.completeExceptionally(failure);
|
||||
}
|
||||
collectionLoads.clear();
|
||||
}
|
||||
|
||||
private void cancelActiveContexts(String message) {
|
||||
for (Context.CancellableContext context : activeContexts) {
|
||||
context.cancel(new CancellationException(message));
|
||||
}
|
||||
}
|
||||
|
||||
private void interruptActiveOperations() {
|
||||
Thread currentThread = Thread.currentThread();
|
||||
for (Thread operationThread : activeOperations.keySet()) {
|
||||
if (operationThread != currentThread) {
|
||||
operationThread.interrupt();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private static String fingerprint(PoolSettings settings) {
|
||||
String value = String.join("\u0000",
|
||||
String.valueOf(settings.uri()),
|
||||
String.valueOf(settings.databaseName()),
|
||||
String.valueOf(settings.token()),
|
||||
String.valueOf(settings.username()),
|
||||
String.valueOf(settings.password()),
|
||||
String.valueOf(settings.poolMaxTotal()),
|
||||
String.valueOf(settings.poolMaxTotalPerKey()),
|
||||
String.valueOf(settings.poolMaxIdlePerKey()),
|
||||
String.valueOf(settings.poolMinIdlePerKey()),
|
||||
String.valueOf(settings.poolMaxWaitMillis()),
|
||||
String.valueOf(settings.poolEvictionIntervalMillis()),
|
||||
String.valueOf(settings.poolMinEvictableIdleMillis())
|
||||
);
|
||||
try {
|
||||
byte[] digest = MessageDigest.getInstance("SHA-256")
|
||||
.digest(value.getBytes(StandardCharsets.UTF_8));
|
||||
StringBuilder result = new StringBuilder(digest.length * 2);
|
||||
for (byte item : digest) {
|
||||
result.append(String.format("%02x", item & 0xff));
|
||||
}
|
||||
return result.toString();
|
||||
} catch (NoSuchAlgorithmException exception) {
|
||||
throw new IllegalStateException("SHA-256 is unavailable", exception);
|
||||
}
|
||||
}
|
||||
|
||||
static ConnectConfig buildConnectConfig(MilvusVectorStoreConfig config) {
|
||||
return buildConnectConfig(PoolSettings.from(config));
|
||||
}
|
||||
|
||||
private static ConnectConfig buildConnectConfig(PoolSettings settings) {
|
||||
String uri = normalizeAndValidateUri(settings.uri());
|
||||
String databaseName = StringUtil.hasText(settings.databaseName())
|
||||
? settings.databaseName().trim()
|
||||
: "default";
|
||||
ConnectConfig.ConnectConfigBuilder<?, ?> builder = ConnectConfig.builder()
|
||||
.uri(uri)
|
||||
.dbName(databaseName);
|
||||
if (StringUtil.hasText(settings.token())) {
|
||||
builder.token(settings.token().trim());
|
||||
}
|
||||
if (StringUtil.hasText(settings.username()) && StringUtil.hasText(settings.password())) {
|
||||
builder.username(settings.username().trim());
|
||||
builder.password(settings.password().trim());
|
||||
}
|
||||
return builder.build();
|
||||
}
|
||||
|
||||
private record PoolSettings(
|
||||
String uri,
|
||||
String databaseName,
|
||||
String token,
|
||||
String username,
|
||||
String password,
|
||||
int poolMaxTotal,
|
||||
int poolMaxTotalPerKey,
|
||||
int poolMaxIdlePerKey,
|
||||
int poolMinIdlePerKey,
|
||||
long poolMaxWaitMillis,
|
||||
long poolEvictionIntervalMillis,
|
||||
long poolMinEvictableIdleMillis
|
||||
) {
|
||||
private static PoolSettings from(MilvusVectorStoreConfig config) {
|
||||
return new PoolSettings(
|
||||
config.getUri(),
|
||||
config.getDatabaseName(),
|
||||
config.getToken(),
|
||||
config.getUsername(),
|
||||
config.getPassword(),
|
||||
config.getPoolMaxTotal(),
|
||||
config.getPoolMaxTotalPerKey(),
|
||||
config.getPoolMaxIdlePerKey(),
|
||||
config.getPoolMinIdlePerKey(),
|
||||
config.getPoolMaxWaitMillis(),
|
||||
config.getPoolEvictionIntervalMillis(),
|
||||
config.getPoolMinEvictableIdleMillis()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
private static final class ManagedMilvusClientV2Pool extends MilvusClientV2Pool {
|
||||
|
||||
private ManagedMilvusClientV2Pool(
|
||||
PoolConfig poolConfig,
|
||||
ConnectConfig connectConfig
|
||||
) throws ClassNotFoundException, NoSuchMethodException {
|
||||
super(poolConfig, connectConfig);
|
||||
}
|
||||
|
||||
private void invalidateClient(String key, MilvusClientV2 client) {
|
||||
try {
|
||||
clientPool.invalidateObject(key, client);
|
||||
} catch (Exception exception) {
|
||||
throw new IllegalStateException("Unable to invalidate Milvus client", exception);
|
||||
}
|
||||
}
|
||||
|
||||
private MilvusClientV2 getClient(String key, Duration maxWait) {
|
||||
if (maxWait == null || maxWait.isZero() || maxWait.isNegative()) {
|
||||
throw new IllegalArgumentException("maxWait must be greater than zero");
|
||||
}
|
||||
try {
|
||||
long waitMillis = Math.max(1L, maxWait.toMillis());
|
||||
return clientPool.borrowObject(key, waitMillis);
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException(
|
||||
"Interrupted while waiting for a Milvus client", exception);
|
||||
} catch (Exception exception) {
|
||||
throw new IllegalStateException(
|
||||
"Unable to borrow a Milvus client", exception);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
static final class CollectionLoadTicket {
|
||||
|
||||
private final String collectionName;
|
||||
private final long generation;
|
||||
private final CompletableFuture<Void> completion;
|
||||
private final boolean leader;
|
||||
|
||||
private CollectionLoadTicket(
|
||||
String collectionName,
|
||||
long generation,
|
||||
CompletableFuture<Void> completion,
|
||||
boolean leader
|
||||
) {
|
||||
this.collectionName = collectionName;
|
||||
this.generation = generation;
|
||||
this.completion = completion;
|
||||
this.leader = leader;
|
||||
}
|
||||
|
||||
boolean isLeader() {
|
||||
return leader;
|
||||
}
|
||||
|
||||
CompletableFuture<Void> completion() {
|
||||
return completion;
|
||||
}
|
||||
|
||||
private CollectionLoadTicket asFollower() {
|
||||
return new CollectionLoadTicket(
|
||||
collectionName, generation, completion, false);
|
||||
}
|
||||
}
|
||||
|
||||
static final class RetryableCollectionLoadException
|
||||
extends RuntimeException {
|
||||
|
||||
private RetryableCollectionLoadException(Throwable cause) {
|
||||
super("The collection load leader exhausted its local budget", cause);
|
||||
}
|
||||
}
|
||||
|
||||
static String normalizeAndValidateUri(String uri) {
|
||||
if (StringUtil.noText(uri)) {
|
||||
throw new IllegalArgumentException(
|
||||
"Milvus uri is required. Example: http://127.0.0.1:19530"
|
||||
);
|
||||
}
|
||||
String normalized = uri.trim();
|
||||
if (!normalized.contains("://")) {
|
||||
normalized = "http://" + normalized;
|
||||
}
|
||||
try {
|
||||
URI parsed = URI.create(normalized);
|
||||
if (StringUtil.noText(parsed.getHost()) || parsed.getPort() <= 0) {
|
||||
throw new IllegalArgumentException("Invalid Milvus uri: " + uri);
|
||||
}
|
||||
} catch (IllegalArgumentException exception) {
|
||||
throw new IllegalArgumentException(
|
||||
"Invalid Milvus uri: " + uri + ". Example: http://127.0.0.1:19530",
|
||||
exception
|
||||
);
|
||||
}
|
||||
return normalized;
|
||||
}
|
||||
}
|
||||
@@ -15,34 +15,45 @@
|
||||
*/
|
||||
package com.easyagents.store.milvus;
|
||||
|
||||
import com.alibaba.fastjson.JSON;
|
||||
import com.alibaba.fastjson.JSONObject;
|
||||
import com.google.gson.Gson;
|
||||
import com.google.gson.JsonObject;
|
||||
import com.easyagents.core.document.Document;
|
||||
import com.easyagents.core.store.DocumentStore;
|
||||
import com.easyagents.core.store.SearchWrapper;
|
||||
import com.easyagents.core.store.StoreOptions;
|
||||
import com.easyagents.core.store.StoreResult;
|
||||
import com.easyagents.core.store.StoreTimeoutException;
|
||||
import com.easyagents.core.util.CollectionUtil;
|
||||
import com.easyagents.core.util.Maps;
|
||||
import com.easyagents.core.util.StringUtil;
|
||||
import io.milvus.v2.client.ConnectConfig;
|
||||
import io.grpc.Context;
|
||||
import io.grpc.Status;
|
||||
import io.milvus.v2.client.MilvusClientV2;
|
||||
import io.milvus.v2.common.ConsistencyLevel;
|
||||
import io.milvus.v2.common.DataType;
|
||||
import io.milvus.v2.common.IndexParam;
|
||||
import io.milvus.v2.exception.MilvusClientException;
|
||||
import io.milvus.v2.service.collection.request.CreateCollectionReq;
|
||||
import io.milvus.v2.service.collection.request.GetLoadStateReq;
|
||||
import io.milvus.v2.service.collection.request.HasCollectionReq;
|
||||
import io.milvus.v2.service.collection.request.LoadCollectionReq;
|
||||
import io.milvus.v2.service.vector.request.*;
|
||||
import io.milvus.v2.service.vector.request.data.FloatVec;
|
||||
import io.milvus.v2.service.vector.response.QueryResp;
|
||||
import io.milvus.v2.service.vector.response.SearchResp;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
|
||||
import java.net.URI;
|
||||
import java.time.Duration;
|
||||
import java.util.*;
|
||||
import java.util.concurrent.CancellationException;
|
||||
import java.util.concurrent.CompletableFuture;
|
||||
import java.util.concurrent.ExecutionException;
|
||||
import java.util.concurrent.Executors;
|
||||
import java.util.concurrent.ScheduledExecutorService;
|
||||
import java.util.concurrent.ThreadFactory;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.TimeoutException;
|
||||
import java.util.function.Function;
|
||||
|
||||
/**
|
||||
* Milvus vector store based on Milvus Java SDK v2.
|
||||
@@ -52,63 +63,74 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
private static final Logger LOG = LoggerFactory.getLogger(MilvusVectorStore.class);
|
||||
private static final long LOAD_TIMEOUT_MS = 30_000L;
|
||||
private static final long LOAD_POLL_INTERVAL_MS = 200L;
|
||||
private static final long DEADLINE_SAFETY_MARGIN_MS = 200L;
|
||||
private static final ScheduledExecutorService DEADLINE_SCHEDULER =
|
||||
Executors.newSingleThreadScheduledExecutor(new ThreadFactory() {
|
||||
@Override
|
||||
public Thread newThread(Runnable runnable) {
|
||||
Thread thread = new Thread(runnable, "milvus-deadline");
|
||||
thread.setDaemon(true);
|
||||
return thread;
|
||||
}
|
||||
});
|
||||
|
||||
private static final String FIELD_ID = "id";
|
||||
private static final String FIELD_CONTENT = "content";
|
||||
private static final String FIELD_METADATA = "metadata";
|
||||
private static final String FIELD_VECTOR = "vector";
|
||||
|
||||
private final MilvusClientV2 client;
|
||||
private static final Gson GSON = new Gson();
|
||||
|
||||
private final MilvusClientManager clientManager;
|
||||
private final MilvusVectorStoreConfig config;
|
||||
private final String defaultCollectionName;
|
||||
private final Set<String> initializedCollections = Collections.synchronizedSet(new HashSet<String>());
|
||||
private final Set<String> loadedCollections = Collections.synchronizedSet(new HashSet<String>());
|
||||
private final boolean ownsClientManager;
|
||||
private volatile MilvusClientV2 compatibilityClient;
|
||||
private volatile boolean closed;
|
||||
|
||||
public MilvusVectorStore(MilvusVectorStoreConfig config) {
|
||||
this.config = config;
|
||||
this.defaultCollectionName = config.getDefaultCollectionName();
|
||||
String uri = normalizeAndValidateUri(config.getUri());
|
||||
String dbName = StringUtil.hasText(config.getDatabaseName()) ? config.getDatabaseName().trim() : "default";
|
||||
|
||||
ConnectConfig.ConnectConfigBuilder<?, ?> builder = ConnectConfig.builder()
|
||||
.uri(uri)
|
||||
.dbName(dbName);
|
||||
|
||||
if (StringUtil.hasText(config.getToken())) {
|
||||
builder.token(config.getToken().trim());
|
||||
}
|
||||
|
||||
if (StringUtil.hasText(config.getUsername()) && StringUtil.hasText(config.getPassword())) {
|
||||
builder.username(config.getUsername().trim());
|
||||
builder.password(config.getPassword().trim());
|
||||
}
|
||||
|
||||
ConnectConfig connectConfig = builder.build();
|
||||
this.client = new MilvusClientV2(connectConfig);
|
||||
this(config, createOwnedClientManager(config), true);
|
||||
}
|
||||
|
||||
private String normalizeAndValidateUri(String uri) {
|
||||
if (StringUtil.noText(uri)) {
|
||||
throw new IllegalArgumentException("Milvus uri is required. Example: http://127.0.0.1:19530");
|
||||
}
|
||||
public MilvusVectorStore(
|
||||
MilvusVectorStoreConfig config,
|
||||
MilvusClientManager clientManager
|
||||
) {
|
||||
this(config, clientManager, false);
|
||||
}
|
||||
|
||||
String normalized = uri.trim();
|
||||
if (!normalized.contains("://")) {
|
||||
normalized = "http://" + normalized;
|
||||
}
|
||||
private MilvusVectorStore(
|
||||
MilvusVectorStoreConfig config,
|
||||
MilvusClientManager clientManager,
|
||||
boolean ownsClientManager
|
||||
) {
|
||||
validateConfig(config);
|
||||
this.config = config;
|
||||
this.defaultCollectionName = config.getDefaultCollectionName();
|
||||
this.clientManager = Objects.requireNonNull(clientManager, "clientManager");
|
||||
this.ownsClientManager = ownsClientManager;
|
||||
}
|
||||
|
||||
URI parsed;
|
||||
try {
|
||||
parsed = URI.create(normalized);
|
||||
} catch (Exception e) {
|
||||
throw new IllegalArgumentException("Invalid Milvus uri: " + uri + ". Example: http://127.0.0.1:19530", e);
|
||||
}
|
||||
private static MilvusClientManager createOwnedClientManager(
|
||||
MilvusVectorStoreConfig config
|
||||
) {
|
||||
validateConfig(config);
|
||||
return new MilvusClientManager(config);
|
||||
}
|
||||
|
||||
if (StringUtil.noText(parsed.getHost()) || parsed.getPort() <= 0) {
|
||||
throw new IllegalArgumentException("Invalid Milvus uri: " + uri + ". Example: http://127.0.0.1:19530");
|
||||
private static void validateConfig(MilvusVectorStoreConfig config) {
|
||||
Objects.requireNonNull(config, "config");
|
||||
if (config.getSearchTimeoutMillis() <= DEADLINE_SAFETY_MARGIN_MS) {
|
||||
throw new IllegalArgumentException(
|
||||
"Milvus searchTimeoutMillis must be greater than "
|
||||
+ DEADLINE_SAFETY_MARGIN_MS
|
||||
);
|
||||
}
|
||||
if (config.getPoolMaxWaitMillis() <= 0L) {
|
||||
throw new IllegalArgumentException(
|
||||
"Milvus poolMaxWaitMillis must be greater than zero"
|
||||
);
|
||||
}
|
||||
|
||||
return normalized;
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -121,22 +143,25 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
throw new IllegalStateException("CollectionName is null or blank. please config the \"defaultCollectionName\" or store with designative collectionName.");
|
||||
}
|
||||
|
||||
int dimension = getDimension(documents);
|
||||
ensureCollectionExists(collectionName, dimension);
|
||||
|
||||
try {
|
||||
InsertReq.InsertReqBuilder<?, ?> builder = InsertReq.builder();
|
||||
if (StringUtil.hasText(options.getPartitionName())) {
|
||||
builder.partitionName(options.getPartitionName());
|
||||
}
|
||||
InsertReq insertReq = builder
|
||||
.collectionName(collectionName)
|
||||
.data(toMilvusDocuments(documents))
|
||||
.build();
|
||||
client.insert(insertReq);
|
||||
int dimension = getDimension(documents);
|
||||
clientManager.withClient(client -> {
|
||||
ensureCollectionExists(client, collectionName, dimension);
|
||||
InsertReq.InsertReqBuilder<?, ?> builder = InsertReq.builder();
|
||||
if (StringUtil.hasText(options.getPartitionName())) {
|
||||
builder.partitionName(options.getPartitionName());
|
||||
}
|
||||
client.insert(builder
|
||||
.collectionName(collectionName)
|
||||
.data(toMilvusDocuments(documents))
|
||||
.build());
|
||||
return null;
|
||||
});
|
||||
return StoreResult.successWithIds(documents);
|
||||
} catch (MilvusClientException e) {
|
||||
return StoreResult.fail();
|
||||
} catch (RuntimeException e) {
|
||||
LOG.error("Milvus insert failed. collection={}, message={}",
|
||||
collectionName, e.getMessage(), e);
|
||||
return StoreResult.fail(e.getMessage());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -159,7 +184,10 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
.collectionName(collectionName)
|
||||
.ids(MilvusPrimaryKeySupport.normalize(ids))
|
||||
.build();
|
||||
client.delete(deleteReq);
|
||||
clientManager.withClient(client -> {
|
||||
client.delete(deleteReq);
|
||||
return null;
|
||||
});
|
||||
return StoreResult.success();
|
||||
} catch (Exception e) {
|
||||
LOG.error("Milvus delete failed. collection={}, message={}",
|
||||
@@ -178,19 +206,22 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
throw new IllegalStateException("CollectionName is null or blank. please config the \"defaultCollectionName\" or store with designative collectionName.");
|
||||
}
|
||||
|
||||
int dimension = getDimension(documents);
|
||||
ensureCollectionExists(collectionName, dimension);
|
||||
|
||||
try {
|
||||
UpsertReq upsertReq = UpsertReq.builder()
|
||||
.collectionName(collectionName)
|
||||
.partitionName(options.getPartitionName())
|
||||
.data(toMilvusDocuments(documents))
|
||||
.build();
|
||||
client.upsert(upsertReq);
|
||||
int dimension = getDimension(documents);
|
||||
clientManager.withClient(client -> {
|
||||
ensureCollectionExists(client, collectionName, dimension);
|
||||
client.upsert(UpsertReq.builder()
|
||||
.collectionName(collectionName)
|
||||
.partitionName(options.getPartitionName())
|
||||
.data(toMilvusDocuments(documents))
|
||||
.build());
|
||||
return null;
|
||||
});
|
||||
return StoreResult.successWithIds(documents);
|
||||
} catch (Exception e) {
|
||||
return StoreResult.fail();
|
||||
LOG.error("Milvus upsert failed. collection={}, message={}",
|
||||
collectionName, e.getMessage(), e);
|
||||
return StoreResult.fail(e.getMessage());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -200,56 +231,104 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
if (StringUtil.noText(collectionName)) {
|
||||
throw new IllegalStateException("CollectionName is null or blank. please config the \"defaultCollectionName\" or store with designative collectionName.");
|
||||
}
|
||||
ensureCollectionLoaded(collectionName);
|
||||
long timeoutMillis = resolveSearchTimeoutMillis(options);
|
||||
long rpcBudgetMillis = timeoutMillis - DEADLINE_SAFETY_MARGIN_MS;
|
||||
long deadlineNanos = deadlineAfterMillis(rpcBudgetMillis);
|
||||
Context.CancellableContext context = Context.current().withDeadlineAfter(
|
||||
rpcBudgetMillis, TimeUnit.MILLISECONDS, DEADLINE_SCHEDULER);
|
||||
try {
|
||||
return clientManager.withRequestContext(context, () ->
|
||||
searchWithinDeadline(
|
||||
wrapper, options, collectionName, deadlineNanos));
|
||||
} catch (RuntimeException exception) {
|
||||
if (!(exception instanceof StoreTimeoutException)
|
||||
&& deadlineExpired(deadlineNanos, exception)) {
|
||||
throw timeoutException(collectionName, exception);
|
||||
}
|
||||
throw exception;
|
||||
} catch (Exception exception) {
|
||||
throw new IllegalStateException("Milvus search failed", exception);
|
||||
} finally {
|
||||
context.cancel(null);
|
||||
}
|
||||
}
|
||||
|
||||
private List<Document> searchWithinDeadline(
|
||||
SearchWrapper wrapper,
|
||||
StoreOptions options,
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
String operation = wrapper.getVector() == null
|
||||
|| wrapper.getVector().length == 0
|
||||
? "query"
|
||||
: "search";
|
||||
ensureCollectionLoaded(collectionName, deadlineNanos);
|
||||
try {
|
||||
return searchOnce(wrapper, options, collectionName, deadlineNanos);
|
||||
} catch (RuntimeException exception) {
|
||||
if (!isCollectionNotLoaded(exception)) {
|
||||
throw propagateSearchFailure(
|
||||
operation, collectionName, exception);
|
||||
}
|
||||
clientManager.markCollectionUnloaded(collectionName);
|
||||
try {
|
||||
ensureCollectionLoaded(collectionName, deadlineNanos);
|
||||
return searchOnce(
|
||||
wrapper, options, collectionName, deadlineNanos);
|
||||
} catch (RuntimeException retryException) {
|
||||
retryException.addSuppressed(exception);
|
||||
throw propagateSearchFailure(
|
||||
operation, collectionName, retryException);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private List<Document> searchOnce(
|
||||
SearchWrapper wrapper,
|
||||
StoreOptions options,
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
if (wrapper.getVector() == null || wrapper.getVector().length == 0) {
|
||||
return queryByCondition(wrapper, options, collectionName);
|
||||
return queryByCondition(
|
||||
wrapper, options, collectionName, deadlineNanos);
|
||||
}
|
||||
return searchByVector(wrapper, options, collectionName);
|
||||
return searchByVector(wrapper, options, collectionName, deadlineNanos);
|
||||
}
|
||||
|
||||
private List<Document> searchByVector(SearchWrapper wrapper, StoreOptions options, String collectionName) {
|
||||
private List<Document> searchByVector(
|
||||
SearchWrapper wrapper,
|
||||
StoreOptions options,
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
SearchReq searchReq = buildSearchReq(wrapper, options, collectionName);
|
||||
try {
|
||||
SearchResp resp = client.search(searchReq);
|
||||
return parseSearchResults(resp, wrapper.getMinScore());
|
||||
} catch (Exception e) {
|
||||
if (isCollectionNotLoaded(e)) {
|
||||
loadedCollections.remove(collectionName);
|
||||
try {
|
||||
ensureCollectionLoaded(collectionName);
|
||||
SearchResp retryResp = client.search(searchReq);
|
||||
return parseSearchResults(retryResp, wrapper.getMinScore());
|
||||
} catch (Exception retryException) {
|
||||
LOG.warn("Milvus search retry failed after load. collection={}, message={}", collectionName, retryException.getMessage());
|
||||
return Collections.emptyList();
|
||||
}
|
||||
}
|
||||
LOG.warn("Milvus search failed. collection={}, message={}", collectionName, e.getMessage());
|
||||
return Collections.emptyList();
|
||||
}
|
||||
SearchResp resp = withClientBeforeDeadline(
|
||||
deadlineNanos, client -> client.search(searchReq));
|
||||
return parseSearchResults(resp, wrapper.getMinScore());
|
||||
}
|
||||
|
||||
private List<Document> queryByCondition(SearchWrapper wrapper, StoreOptions options, String collectionName) {
|
||||
private List<Document> queryByCondition(
|
||||
SearchWrapper wrapper,
|
||||
StoreOptions options,
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
QueryReq queryReq = buildQueryReq(wrapper, options, collectionName);
|
||||
try {
|
||||
QueryResp resp = client.query(queryReq);
|
||||
return parseQueryResults(resp);
|
||||
} catch (Exception e) {
|
||||
if (isCollectionNotLoaded(e)) {
|
||||
loadedCollections.remove(collectionName);
|
||||
try {
|
||||
ensureCollectionLoaded(collectionName);
|
||||
QueryResp retryResp = client.query(queryReq);
|
||||
return parseQueryResults(retryResp);
|
||||
} catch (Exception retryException) {
|
||||
LOG.warn("Milvus query retry failed after load. collection={}, message={}", collectionName, retryException.getMessage());
|
||||
return Collections.emptyList();
|
||||
}
|
||||
}
|
||||
LOG.warn("Milvus query failed. collection={}, message={}", collectionName, e.getMessage());
|
||||
return Collections.emptyList();
|
||||
}
|
||||
QueryResp resp = withClientBeforeDeadline(
|
||||
deadlineNanos, client -> client.query(queryReq));
|
||||
return parseQueryResults(resp);
|
||||
}
|
||||
|
||||
private RuntimeException propagateSearchFailure(
|
||||
String operation,
|
||||
String collectionName,
|
||||
RuntimeException exception
|
||||
) {
|
||||
LOG.error("Milvus {} failed. collection={}, message={}",
|
||||
operation, collectionName, exception.getMessage(), exception);
|
||||
return exception;
|
||||
}
|
||||
|
||||
private SearchReq buildSearchReq(SearchWrapper wrapper, StoreOptions options, String collectionName) {
|
||||
@@ -259,7 +338,7 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
.outputFields(getOutputFields(wrapper))
|
||||
.topK(wrapper.getMaxResults())
|
||||
.annsField(FIELD_VECTOR)
|
||||
.data(Collections.singletonList(toFloatList(wrapper.getVector())))
|
||||
.data(Collections.singletonList(new FloatVec(wrapper.getVector())))
|
||||
.searchParams(Maps.of("ef", 64));
|
||||
|
||||
if (CollectionUtil.hasItems(options.getPartitionNamesOrEmpty())) {
|
||||
@@ -305,11 +384,7 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
continue;
|
||||
}
|
||||
document.setId(result.getId());
|
||||
Float distance = result.getDistance();
|
||||
if (distance != null) {
|
||||
double score = (distance + 1.0d) / 2.0d;
|
||||
document.setScore(score);
|
||||
}
|
||||
document.setScore(normalizeScore(result.getScore()));
|
||||
if (minScore == null || document.getScore() == null || document.getScore() >= minScore) {
|
||||
documents.add(document);
|
||||
}
|
||||
@@ -318,6 +393,10 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
return documents;
|
||||
}
|
||||
|
||||
static Double normalizeScore(Float rawScore) {
|
||||
return rawScore == null ? null : (rawScore + 1.0d) / 2.0d;
|
||||
}
|
||||
|
||||
private List<Document> parseQueryResults(QueryResp resp) {
|
||||
List<QueryResp.QueryResult> results = resp.getQueryResults();
|
||||
if (CollectionUtil.noItems(results)) {
|
||||
@@ -360,22 +439,28 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
document.addMetadata(metadata);
|
||||
} else if (metadataObj != null) {
|
||||
@SuppressWarnings("unchecked")
|
||||
Map<String, Object> metadata = JSON.parseObject(JSON.toJSONString(metadataObj), Map.class);
|
||||
Map<String, Object> metadata = GSON.fromJson(
|
||||
GSON.toJsonTree(metadataObj),
|
||||
Map.class
|
||||
);
|
||||
document.addMetadata(metadata);
|
||||
}
|
||||
|
||||
return document;
|
||||
}
|
||||
|
||||
private List<JSONObject> toMilvusDocuments(List<Document> documents) {
|
||||
List<JSONObject> rows = new ArrayList<JSONObject>(documents.size());
|
||||
List<JsonObject> toMilvusDocuments(List<Document> documents) {
|
||||
List<JsonObject> rows = new ArrayList<JsonObject>(documents.size());
|
||||
for (Document doc : documents) {
|
||||
JSONObject row = new JSONObject();
|
||||
row.put(FIELD_ID, String.valueOf(doc.getId()));
|
||||
row.put(FIELD_CONTENT, doc.getContent());
|
||||
row.put(FIELD_VECTOR, toFloatList(doc.getVector()));
|
||||
JsonObject row = new JsonObject();
|
||||
row.addProperty(FIELD_ID, String.valueOf(doc.getId()));
|
||||
row.addProperty(FIELD_CONTENT, doc.getContent());
|
||||
row.add(FIELD_VECTOR, GSON.toJsonTree(toFloatList(doc.getVector())));
|
||||
Map<String, Object> metadatas = doc.getMetadataMap();
|
||||
row.put(FIELD_METADATA, metadatas == null ? new JSONObject() : new JSONObject(metadatas));
|
||||
row.add(
|
||||
FIELD_METADATA,
|
||||
metadatas == null ? new JsonObject() : GSON.toJsonTree(metadatas)
|
||||
);
|
||||
rows.add(row);
|
||||
}
|
||||
return rows;
|
||||
@@ -413,67 +498,175 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
throw new IllegalStateException("Unable to determine vector dimension for Milvus collection.");
|
||||
}
|
||||
|
||||
private void ensureCollectionExists(String collectionName, int dimension) {
|
||||
if (initializedCollections.contains(collectionName)) {
|
||||
private void ensureCollectionExists(
|
||||
MilvusClientV2 client,
|
||||
String collectionName,
|
||||
int dimension
|
||||
) {
|
||||
if (clientManager.isCollectionInitialized(collectionName)) {
|
||||
return;
|
||||
}
|
||||
synchronized (initializedCollections) {
|
||||
if (initializedCollections.contains(collectionName)) {
|
||||
synchronized (clientManager.initializedCollectionsLock()) {
|
||||
if (clientManager.isCollectionInitialized(collectionName)) {
|
||||
return;
|
||||
}
|
||||
Boolean exists = client.hasCollection(HasCollectionReq.builder().collectionName(collectionName).build());
|
||||
if (Boolean.TRUE.equals(exists)) {
|
||||
initializedCollections.add(collectionName);
|
||||
clientManager.markCollectionInitialized(collectionName);
|
||||
return;
|
||||
}
|
||||
if (!config.isAutoCreateCollection()) {
|
||||
throw new IllegalStateException("Milvus collection not found and autoCreateCollection is disabled: " + collectionName);
|
||||
}
|
||||
createCollection(collectionName, dimension);
|
||||
initializedCollections.add(collectionName);
|
||||
createCollection(client, collectionName, dimension);
|
||||
clientManager.markCollectionInitialized(collectionName);
|
||||
}
|
||||
}
|
||||
|
||||
private void ensureCollectionLoaded(String collectionName) {
|
||||
if (loadedCollections.contains(collectionName)) {
|
||||
return;
|
||||
}
|
||||
synchronized (loadedCollections) {
|
||||
if (loadedCollections.contains(collectionName)) {
|
||||
return;
|
||||
private void ensureCollectionLoaded(
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
while (!clientManager.isCollectionLoaded(collectionName)) {
|
||||
MilvusClientManager.CollectionLoadTicket ticket =
|
||||
clientManager.beginCollectionLoad(collectionName);
|
||||
if (!ticket.isLeader()) {
|
||||
if (awaitCollectionLoad(
|
||||
ticket, collectionName, deadlineNanos)) {
|
||||
return;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
boolean loaded = false;
|
||||
try {
|
||||
loaded = Boolean.TRUE.equals(client.getLoadState(GetLoadStateReq.builder().collectionName(collectionName).build()));
|
||||
} catch (Exception e) {
|
||||
LOG.warn("Milvus getLoadState failed. collection={}, message={}", collectionName, e.getMessage());
|
||||
if (clientManager.isCollectionLoaded(collectionName)) {
|
||||
clientManager.completeCollectionLoad(ticket);
|
||||
return;
|
||||
}
|
||||
withClientBeforeDeadline(deadlineNanos, client -> {
|
||||
boolean loaded = Boolean.TRUE.equals(client.getLoadState(
|
||||
GetLoadStateReq.builder()
|
||||
.collectionName(collectionName)
|
||||
.build()
|
||||
));
|
||||
if (!loaded) {
|
||||
client.loadCollection(LoadCollectionReq.builder()
|
||||
.collectionName(collectionName)
|
||||
.async(false)
|
||||
.build());
|
||||
waitForCollectionLoaded(
|
||||
client, collectionName, deadlineNanos);
|
||||
}
|
||||
return null;
|
||||
});
|
||||
clientManager.completeCollectionLoad(ticket);
|
||||
return;
|
||||
} catch (RuntimeException | Error failure) {
|
||||
clientManager.failCollectionLoad(
|
||||
ticket, failure, isLeaderLocalAbort(failure));
|
||||
throw failure;
|
||||
} finally {
|
||||
clientManager.endCollectionLoad(ticket);
|
||||
}
|
||||
|
||||
if (!loaded) {
|
||||
client.loadCollection(LoadCollectionReq.builder().collectionName(collectionName).build());
|
||||
waitForCollectionLoaded(collectionName);
|
||||
}
|
||||
loadedCollections.add(collectionName);
|
||||
}
|
||||
}
|
||||
|
||||
private void waitForCollectionLoaded(String collectionName) {
|
||||
long deadline = System.currentTimeMillis() + LOAD_TIMEOUT_MS;
|
||||
while (System.currentTimeMillis() < deadline) {
|
||||
private boolean isLeaderLocalAbort(Throwable failure) {
|
||||
if (clientManager.isClosed()) {
|
||||
return false;
|
||||
}
|
||||
if (failure instanceof StoreTimeoutException
|
||||
|| Thread.currentThread().isInterrupted()
|
||||
|| Context.current().isCancelled()) {
|
||||
return true;
|
||||
}
|
||||
Throwable current = failure;
|
||||
while (current != null) {
|
||||
if (current instanceof InterruptedException) {
|
||||
return true;
|
||||
}
|
||||
current = current.getCause();
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
private boolean awaitCollectionLoad(
|
||||
MilvusClientManager.CollectionLoadTicket ticket,
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
Context currentContext = Context.current();
|
||||
CompletableFuture<Void> cancelled = new CompletableFuture<Void>();
|
||||
Context.CancellationListener cancellationListener = context ->
|
||||
cancelled.completeExceptionally(new CancellationException(
|
||||
"Milvus search was cancelled"));
|
||||
currentContext.addListener(cancellationListener, Runnable::run);
|
||||
try {
|
||||
CompletableFuture.anyOf(ticket.completion(), cancelled).get(
|
||||
remainingNanos(deadlineNanos, collectionName),
|
||||
TimeUnit.NANOSECONDS
|
||||
);
|
||||
return true;
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException(
|
||||
"Interrupted while loading Milvus collection: "
|
||||
+ collectionName,
|
||||
exception
|
||||
);
|
||||
} catch (TimeoutException exception) {
|
||||
throw timeoutException(collectionName, exception);
|
||||
} catch (ExecutionException exception) {
|
||||
Throwable cause = exception.getCause();
|
||||
if (cause instanceof MilvusClientManager
|
||||
.RetryableCollectionLoadException) {
|
||||
remainingNanos(deadlineNanos, collectionName);
|
||||
return false;
|
||||
}
|
||||
if (cause instanceof RuntimeException runtimeException) {
|
||||
throw runtimeException;
|
||||
}
|
||||
if (cause instanceof Error error) {
|
||||
throw error;
|
||||
}
|
||||
throw new IllegalStateException(
|
||||
"Unable to load Milvus collection: " + collectionName,
|
||||
cause
|
||||
);
|
||||
} catch (CancellationException exception) {
|
||||
throw new IllegalStateException(
|
||||
"Milvus collection load was cancelled: " + collectionName,
|
||||
exception
|
||||
);
|
||||
} finally {
|
||||
currentContext.removeListener(cancellationListener);
|
||||
}
|
||||
}
|
||||
|
||||
private void waitForCollectionLoaded(
|
||||
MilvusClientV2 client,
|
||||
String collectionName,
|
||||
long deadlineNanos
|
||||
) {
|
||||
while (true) {
|
||||
long remainingNanos = remainingNanos(
|
||||
deadlineNanos, collectionName);
|
||||
if (Boolean.TRUE.equals(client.getLoadState(GetLoadStateReq.builder().collectionName(collectionName).build()))) {
|
||||
return;
|
||||
}
|
||||
try {
|
||||
Thread.sleep(LOAD_POLL_INTERVAL_MS);
|
||||
long sleepMillis = Math.min(
|
||||
LOAD_POLL_INTERVAL_MS,
|
||||
Math.max(1L, TimeUnit.NANOSECONDS.toMillis(remainingNanos))
|
||||
);
|
||||
Thread.sleep(sleepMillis);
|
||||
} catch (InterruptedException e) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException("Interrupted while loading Milvus collection: " + collectionName, e);
|
||||
}
|
||||
}
|
||||
throw new IllegalStateException("Timeout waiting for Milvus collection loaded: " + collectionName);
|
||||
}
|
||||
|
||||
private boolean isCollectionNotLoaded(Exception e) {
|
||||
private boolean isCollectionNotLoaded(Throwable e) {
|
||||
Throwable current = e;
|
||||
while (current != null) {
|
||||
String message = current.getMessage();
|
||||
@@ -485,7 +678,91 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
return false;
|
||||
}
|
||||
|
||||
private void createCollection(String collectionName, int dimension) {
|
||||
private <T> T withClientBeforeDeadline(
|
||||
long deadlineNanos,
|
||||
Function<MilvusClientV2, T> operation
|
||||
) {
|
||||
long remainingNanos = remainingNanos(deadlineNanos, null);
|
||||
long poolWaitNanos = TimeUnit.MILLISECONDS.toNanos(
|
||||
config.getPoolMaxWaitMillis());
|
||||
try {
|
||||
return clientManager.withClient(
|
||||
Duration.ofNanos(Math.min(remainingNanos, poolWaitNanos)),
|
||||
operation
|
||||
);
|
||||
} catch (RuntimeException exception) {
|
||||
if (deadlineExpired(deadlineNanos, exception)) {
|
||||
throw timeoutException(null, exception);
|
||||
}
|
||||
throw exception;
|
||||
}
|
||||
}
|
||||
|
||||
private long resolveSearchTimeoutMillis(StoreOptions options) {
|
||||
long timeoutMillis = config.getSearchTimeoutMillis();
|
||||
Long requestedTimeoutMillis = options.getTimeoutMillis();
|
||||
if (requestedTimeoutMillis != null) {
|
||||
timeoutMillis = Math.min(timeoutMillis, requestedTimeoutMillis);
|
||||
}
|
||||
if (timeoutMillis <= DEADLINE_SAFETY_MARGIN_MS) {
|
||||
throw new StoreTimeoutException(
|
||||
"Insufficient time remaining for Milvus search"
|
||||
);
|
||||
}
|
||||
return timeoutMillis;
|
||||
}
|
||||
|
||||
private static long deadlineAfterMillis(long timeoutMillis) {
|
||||
long now = System.nanoTime();
|
||||
long timeoutNanos = TimeUnit.MILLISECONDS.toNanos(timeoutMillis);
|
||||
if (now > Long.MAX_VALUE - timeoutNanos) {
|
||||
return Long.MAX_VALUE;
|
||||
}
|
||||
return now + timeoutNanos;
|
||||
}
|
||||
|
||||
private static long remainingNanos(
|
||||
long deadlineNanos,
|
||||
String collectionName
|
||||
) {
|
||||
if (deadlineNanos == Long.MAX_VALUE) {
|
||||
return Long.MAX_VALUE;
|
||||
}
|
||||
long remaining = deadlineNanos - System.nanoTime();
|
||||
if (remaining <= 0L) {
|
||||
throw timeoutException(collectionName, null);
|
||||
}
|
||||
return remaining;
|
||||
}
|
||||
|
||||
private static boolean deadlineExpired(
|
||||
long deadlineNanos,
|
||||
Throwable failure
|
||||
) {
|
||||
if (deadlineNanos != Long.MAX_VALUE
|
||||
&& System.nanoTime() >= deadlineNanos) {
|
||||
return true;
|
||||
}
|
||||
Throwable cancellationCause = Context.current().cancellationCause();
|
||||
return cancellationCause instanceof TimeoutException
|
||||
|| Status.fromThrowable(failure).getCode()
|
||||
== Status.Code.DEADLINE_EXCEEDED;
|
||||
}
|
||||
|
||||
private static StoreTimeoutException timeoutException(
|
||||
String collectionName,
|
||||
Throwable cause
|
||||
) {
|
||||
String suffix = StringUtil.hasText(collectionName)
|
||||
? ": " + collectionName
|
||||
: "";
|
||||
return new StoreTimeoutException(
|
||||
"Timeout waiting for Milvus search" + suffix,
|
||||
cause
|
||||
);
|
||||
}
|
||||
|
||||
private void createCollection(MilvusClientV2 client, String collectionName, int dimension) {
|
||||
List<CreateCollectionReq.FieldSchema> fieldSchemaList = new ArrayList<CreateCollectionReq.FieldSchema>();
|
||||
fieldSchemaList.add(CreateCollectionReq.FieldSchema.builder()
|
||||
.name(FIELD_ID)
|
||||
@@ -531,31 +808,83 @@ public class MilvusVectorStore extends DocumentStore implements AutoCloseable {
|
||||
.indexParams(indexParams)
|
||||
.build();
|
||||
client.createCollection(createCollectionReq);
|
||||
ensureCollectionLoaded(collectionName);
|
||||
ensureCollectionLoadedForWrite(client, collectionName);
|
||||
}
|
||||
|
||||
public MilvusClientV2 getClient() {
|
||||
return client;
|
||||
private void ensureCollectionLoadedForWrite(
|
||||
MilvusClientV2 client,
|
||||
String collectionName
|
||||
) {
|
||||
if (clientManager.isCollectionLoaded(collectionName)) {
|
||||
return;
|
||||
}
|
||||
boolean loaded = Boolean.TRUE.equals(client.getLoadState(
|
||||
GetLoadStateReq.builder().collectionName(collectionName).build()));
|
||||
if (!loaded) {
|
||||
client.loadCollection(LoadCollectionReq.builder()
|
||||
.collectionName(collectionName)
|
||||
.async(false)
|
||||
.build());
|
||||
waitForCollectionLoaded(
|
||||
client,
|
||||
collectionName,
|
||||
deadlineAfterMillis(LOAD_TIMEOUT_MS)
|
||||
);
|
||||
}
|
||||
clientManager.markCollectionLoaded(collectionName);
|
||||
}
|
||||
|
||||
public boolean checkAvailable() {
|
||||
try {
|
||||
return client.hasCollection(HasCollectionReq.builder()
|
||||
.collectionName("__milvus_boot_probe__")
|
||||
.build()) != null;
|
||||
return clientManager.withClient(client -> client.hasCollection(
|
||||
HasCollectionReq.builder()
|
||||
.collectionName("__milvus_boot_probe__")
|
||||
.build()
|
||||
)) != null;
|
||||
} catch (Exception e) {
|
||||
LOG.warn("Milvus availability check failed. message={}", e.getMessage());
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Returns a compatibility client for integrations that used the pre-pool API.
|
||||
* Prefer store operations so pooled lifecycle management remains automatic.
|
||||
*/
|
||||
@Deprecated
|
||||
public synchronized MilvusClientV2 getClient() {
|
||||
if (closed) {
|
||||
throw new IllegalStateException("Milvus vector store is closed");
|
||||
}
|
||||
if (compatibilityClient == null) {
|
||||
compatibilityClient = new MilvusClientV2(
|
||||
MilvusClientManager.buildConnectConfig(config)
|
||||
);
|
||||
}
|
||||
return compatibilityClient;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void close() {
|
||||
try {
|
||||
client.close(1L);
|
||||
} catch (InterruptedException e) {
|
||||
Thread.currentThread().interrupt();
|
||||
LOG.warn("Interrupted while closing Milvus client. uri={}", config.getUri(), e);
|
||||
MilvusClientV2 legacyClient;
|
||||
synchronized (this) {
|
||||
if (closed) {
|
||||
return;
|
||||
}
|
||||
closed = true;
|
||||
legacyClient = compatibilityClient;
|
||||
compatibilityClient = null;
|
||||
}
|
||||
if (legacyClient != null) {
|
||||
try {
|
||||
legacyClient.close(1L);
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
LOG.warn("Interrupted while closing compatibility Milvus client", exception);
|
||||
}
|
||||
}
|
||||
if (ownsClientManager) {
|
||||
clientManager.close();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -30,6 +30,14 @@ public class MilvusVectorStoreConfig implements DocumentStoreConfig {
|
||||
private String password;
|
||||
private String defaultCollectionName;
|
||||
private boolean autoCreateCollection = true;
|
||||
private int poolMaxTotal = 8;
|
||||
private int poolMaxTotalPerKey = 8;
|
||||
private int poolMaxIdlePerKey = 4;
|
||||
private int poolMinIdlePerKey = 1;
|
||||
private long poolMaxWaitMillis = 3_000L;
|
||||
private long poolEvictionIntervalMillis = 60_000L;
|
||||
private long poolMinEvictableIdleMillis = 300_000L;
|
||||
private long searchTimeoutMillis = 10_000L;
|
||||
|
||||
public String getUri() {
|
||||
return uri;
|
||||
@@ -87,6 +95,70 @@ public class MilvusVectorStoreConfig implements DocumentStoreConfig {
|
||||
this.autoCreateCollection = autoCreateCollection;
|
||||
}
|
||||
|
||||
public int getPoolMaxTotal() {
|
||||
return poolMaxTotal;
|
||||
}
|
||||
|
||||
public void setPoolMaxTotal(int poolMaxTotal) {
|
||||
this.poolMaxTotal = poolMaxTotal;
|
||||
}
|
||||
|
||||
public int getPoolMaxTotalPerKey() {
|
||||
return poolMaxTotalPerKey;
|
||||
}
|
||||
|
||||
public void setPoolMaxTotalPerKey(int poolMaxTotalPerKey) {
|
||||
this.poolMaxTotalPerKey = poolMaxTotalPerKey;
|
||||
}
|
||||
|
||||
public int getPoolMaxIdlePerKey() {
|
||||
return poolMaxIdlePerKey;
|
||||
}
|
||||
|
||||
public void setPoolMaxIdlePerKey(int poolMaxIdlePerKey) {
|
||||
this.poolMaxIdlePerKey = poolMaxIdlePerKey;
|
||||
}
|
||||
|
||||
public int getPoolMinIdlePerKey() {
|
||||
return poolMinIdlePerKey;
|
||||
}
|
||||
|
||||
public void setPoolMinIdlePerKey(int poolMinIdlePerKey) {
|
||||
this.poolMinIdlePerKey = poolMinIdlePerKey;
|
||||
}
|
||||
|
||||
public long getPoolMaxWaitMillis() {
|
||||
return poolMaxWaitMillis;
|
||||
}
|
||||
|
||||
public void setPoolMaxWaitMillis(long poolMaxWaitMillis) {
|
||||
this.poolMaxWaitMillis = poolMaxWaitMillis;
|
||||
}
|
||||
|
||||
public long getPoolEvictionIntervalMillis() {
|
||||
return poolEvictionIntervalMillis;
|
||||
}
|
||||
|
||||
public void setPoolEvictionIntervalMillis(long poolEvictionIntervalMillis) {
|
||||
this.poolEvictionIntervalMillis = poolEvictionIntervalMillis;
|
||||
}
|
||||
|
||||
public long getPoolMinEvictableIdleMillis() {
|
||||
return poolMinEvictableIdleMillis;
|
||||
}
|
||||
|
||||
public void setPoolMinEvictableIdleMillis(long poolMinEvictableIdleMillis) {
|
||||
this.poolMinEvictableIdleMillis = poolMinEvictableIdleMillis;
|
||||
}
|
||||
|
||||
public long getSearchTimeoutMillis() {
|
||||
return searchTimeoutMillis;
|
||||
}
|
||||
|
||||
public void setSearchTimeoutMillis(long searchTimeoutMillis) {
|
||||
this.searchTimeoutMillis = searchTimeoutMillis;
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean checkAvailable() {
|
||||
return StringUtil.hasText(this.uri);
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
package com.easyagents.store.milvus;
|
||||
|
||||
import com.easyagents.core.document.Document;
|
||||
import com.google.gson.JsonObject;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/**
|
||||
* Milvus SDK 2.3.11 数据适配回归测试。
|
||||
*/
|
||||
public class MilvusVectorStoreCompatibilityTest {
|
||||
|
||||
@Test
|
||||
public void shouldConvertRowsToGsonWithoutLosingMetadata() {
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
config.setUri("http://127.0.0.1:19530");
|
||||
config.setDefaultCollectionName("test");
|
||||
MilvusVectorStore store = new MilvusVectorStore(config);
|
||||
try {
|
||||
Document document = Document.of("正文");
|
||||
document.setId("chunk-1");
|
||||
document.setVector(new float[] { 0.25F, 0.75F });
|
||||
document.addMetadata(Map.of("knowledgeId", "knowledge-1"));
|
||||
|
||||
List<JsonObject> rows = store.toMilvusDocuments(List.of(document));
|
||||
|
||||
Assert.assertEquals(1, rows.size());
|
||||
Assert.assertEquals("chunk-1", rows.get(0).get("id").getAsString());
|
||||
Assert.assertEquals("正文", rows.get(0).get("content").getAsString());
|
||||
Assert.assertEquals(2, rows.get(0).getAsJsonArray("vector").size());
|
||||
Assert.assertEquals(
|
||||
"knowledge-1",
|
||||
rows.get(0).getAsJsonObject("metadata").get("knowledgeId").getAsString()
|
||||
);
|
||||
} finally {
|
||||
store.close();
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldNormalizeUriWithoutExposingCredentialsInPoolKey() {
|
||||
Assert.assertEquals(
|
||||
"http://127.0.0.1:19530",
|
||||
MilvusClientManager.normalizeAndValidateUri("127.0.0.1:19530")
|
||||
);
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRebuildPoolOnlyWhenConnectionSettingsChange() {
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
config.setUri("http://127.0.0.1:19530");
|
||||
MilvusClientManager manager = new MilvusClientManager(config);
|
||||
try {
|
||||
Assert.assertFalse(manager.reconfigureIfNeeded(config));
|
||||
|
||||
config.setPoolMaxTotal(9);
|
||||
|
||||
Assert.assertTrue(manager.reconfigureIfNeeded(config));
|
||||
Assert.assertFalse(manager.reconfigureIfNeeded(config));
|
||||
} finally {
|
||||
manager.close();
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldPreserveCosineScoreNormalization() {
|
||||
Assert.assertEquals(Double.valueOf(1.0D), MilvusVectorStore.normalizeScore(1.0F));
|
||||
Assert.assertEquals(Double.valueOf(0.5D), MilvusVectorStore.normalizeScore(0.0F));
|
||||
Assert.assertEquals(Double.valueOf(0.0D), MilvusVectorStore.normalizeScore(-1.0F));
|
||||
Assert.assertNull(MilvusVectorStore.normalizeScore(null));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void shouldRejectCompatibilityClientAfterStoreCloses() {
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
config.setUri("http://127.0.0.1:19530");
|
||||
MilvusVectorStore store = new MilvusVectorStore(config);
|
||||
|
||||
store.close();
|
||||
|
||||
try {
|
||||
store.getClient();
|
||||
Assert.fail("A closed store must not recreate a compatibility client");
|
||||
} catch (IllegalStateException expected) {
|
||||
Assert.assertEquals(
|
||||
"Milvus vector store is closed", expected.getMessage());
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -34,4 +34,24 @@ public class MilvusVectorStoreConfigTest {
|
||||
config.setPassword("Milvus");
|
||||
Assert.assertTrue(config.checkAvailable());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testPoolDefaultsAreBounded() {
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
Assert.assertEquals(8, config.getPoolMaxTotal());
|
||||
Assert.assertEquals(8, config.getPoolMaxTotalPerKey());
|
||||
Assert.assertEquals(4, config.getPoolMaxIdlePerKey());
|
||||
Assert.assertEquals(1, config.getPoolMinIdlePerKey());
|
||||
Assert.assertEquals(3_000L, config.getPoolMaxWaitMillis());
|
||||
Assert.assertEquals(300_000L, config.getPoolMinEvictableIdleMillis());
|
||||
Assert.assertEquals(10_000L, config.getSearchTimeoutMillis());
|
||||
}
|
||||
|
||||
@Test(expected = IllegalArgumentException.class)
|
||||
public void testSearchTimeoutMustLeaveCleanupMargin() {
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
config.setUri("http://127.0.0.1:19530");
|
||||
config.setSearchTimeoutMillis(200L);
|
||||
new MilvusVectorStore(config);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,935 @@
|
||||
package com.easyagents.store.milvus;
|
||||
|
||||
import com.easyagents.core.document.Document;
|
||||
import com.easyagents.core.store.SearchWrapper;
|
||||
import com.easyagents.core.store.StoreOptions;
|
||||
import com.easyagents.core.store.StoreTimeoutException;
|
||||
import io.grpc.Context;
|
||||
import io.grpc.Server;
|
||||
import io.grpc.netty.shaded.io.grpc.netty.NettyServerBuilder;
|
||||
import io.grpc.stub.ServerCallStreamObserver;
|
||||
import io.grpc.stub.StreamObserver;
|
||||
import io.milvus.grpc.CheckHealthRequest;
|
||||
import io.milvus.grpc.CheckHealthResponse;
|
||||
import io.milvus.grpc.ConnectRequest;
|
||||
import io.milvus.grpc.ConnectResponse;
|
||||
import io.milvus.grpc.CollectionSchema;
|
||||
import io.milvus.grpc.DataType;
|
||||
import io.milvus.grpc.DescribeCollectionRequest;
|
||||
import io.milvus.grpc.DescribeCollectionResponse;
|
||||
import io.milvus.grpc.ErrorCode;
|
||||
import io.milvus.grpc.FieldSchema;
|
||||
import io.milvus.grpc.GetLoadStateRequest;
|
||||
import io.milvus.grpc.GetLoadStateResponse;
|
||||
import io.milvus.grpc.ListDatabasesRequest;
|
||||
import io.milvus.grpc.ListDatabasesResponse;
|
||||
import io.milvus.grpc.LoadCollectionRequest;
|
||||
import io.milvus.grpc.LoadState;
|
||||
import io.milvus.grpc.MilvusServiceGrpc;
|
||||
import io.milvus.grpc.QueryRequest;
|
||||
import io.milvus.grpc.QueryResults;
|
||||
import io.milvus.v2.client.MilvusClientV2;
|
||||
import io.milvus.v2.service.vector.request.QueryReq;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Set;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
import java.util.concurrent.CountDownLatch;
|
||||
import java.util.concurrent.ExecutionException;
|
||||
import java.util.concurrent.ExecutorService;
|
||||
import java.util.concurrent.Executors;
|
||||
import java.util.concurrent.Future;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.TimeoutException;
|
||||
import java.util.concurrent.atomic.AtomicBoolean;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
import java.util.function.BooleanSupplier;
|
||||
|
||||
public class MilvusVectorStoreGrpcTest {
|
||||
|
||||
private static final long AWAIT_SECONDS = 5L;
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldUseOneSdkAttemptAndKeepClientReusable() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 2_000L)) {
|
||||
fixture.manager.markCollectionLoaded("docs");
|
||||
Assert.assertEquals(
|
||||
0L,
|
||||
MilvusClientManager.buildConnectConfig(fixture.config).getRpcDeadlineMs()
|
||||
);
|
||||
|
||||
server.failQueriesWithUnavailable();
|
||||
assertSearchFails(fixture.store, "docs");
|
||||
|
||||
Assert.assertEquals(1, server.queryCalls.get());
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(1, fixture.manager.getIdleClientCount());
|
||||
|
||||
server.succeedQueries();
|
||||
Assert.assertTrue(search(fixture.store, "docs").isEmpty());
|
||||
Assert.assertEquals(2, server.queryCalls.get());
|
||||
Assert.assertEquals(1, server.connectCalls.get());
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldReleaseAndReuseClientAfterContextDeadline() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 1_000L)) {
|
||||
fixture.manager.markCollectionLoaded("docs");
|
||||
fixture.manager.withClient(client -> client);
|
||||
QueryBlock block = server.blockQueries();
|
||||
ExecutorService executor = Executors.newSingleThreadExecutor();
|
||||
try {
|
||||
Future<List<Document>> search = executor.submit(() ->
|
||||
search(fixture.store, "docs"));
|
||||
|
||||
block.awaitEntered();
|
||||
block.awaitCancelled();
|
||||
Throwable failure = futureFailure(search);
|
||||
Assert.assertTrue(failure.toString(),
|
||||
failure instanceof StoreTimeoutException);
|
||||
|
||||
Assert.assertTrue(server.querySawDeadline.get());
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(1, fixture.manager.getIdleClientCount());
|
||||
|
||||
server.succeedQueries();
|
||||
Assert.assertTrue(search(fixture.store, "docs").isEmpty());
|
||||
Assert.assertEquals(1, server.connectCalls.get());
|
||||
} finally {
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldReleaseAndReuseClientAfterThreadInterrupt() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 5_000L)) {
|
||||
fixture.manager.markCollectionLoaded("docs");
|
||||
fixture.manager.withClient(client -> client);
|
||||
QueryBlock block = server.blockQueries();
|
||||
CountDownLatch finished = new CountDownLatch(1);
|
||||
AtomicReference<Throwable> failure = new AtomicReference<>();
|
||||
AtomicBoolean interrupted = new AtomicBoolean();
|
||||
Thread searchThread = new Thread(() -> {
|
||||
try {
|
||||
search(fixture.store, "docs");
|
||||
} catch (Throwable exception) {
|
||||
failure.set(exception);
|
||||
} finally {
|
||||
interrupted.set(Thread.currentThread().isInterrupted());
|
||||
finished.countDown();
|
||||
}
|
||||
}, "milvus-interrupt-test");
|
||||
searchThread.start();
|
||||
|
||||
block.awaitEntered();
|
||||
searchThread.interrupt();
|
||||
Assert.assertTrue(finished.await(AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
block.awaitCancelled();
|
||||
|
||||
Assert.assertNotNull(failure.get());
|
||||
Assert.assertTrue(interrupted.get());
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(1, fixture.manager.getIdleClientCount());
|
||||
|
||||
server.succeedQueries();
|
||||
Assert.assertTrue(search(fixture.store, "docs").isEmpty());
|
||||
Assert.assertEquals(1, server.connectCalls.get());
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldKeepPoolAfterOrdinaryBusinessFailure() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 2_000L)) {
|
||||
MilvusClientV2 first = fixture.manager.withClient(client -> client);
|
||||
|
||||
try {
|
||||
fixture.manager.withClient(client -> {
|
||||
throw new IllegalArgumentException("synthetic business failure");
|
||||
});
|
||||
Assert.fail("The business failure must be propagated");
|
||||
} catch (IllegalArgumentException expected) {
|
||||
Assert.assertEquals("synthetic business failure", expected.getMessage());
|
||||
}
|
||||
|
||||
MilvusClientV2 second = fixture.manager.withClient(client -> client);
|
||||
Assert.assertSame(first, second);
|
||||
Assert.assertEquals(1, server.connectCalls.get());
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(1, fixture.manager.getIdleClientCount());
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldAttachContextToEveryClientOperation() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 2_000L)) {
|
||||
Context callerContext = Context.current();
|
||||
Context operationContext = fixture.manager.withClient(client ->
|
||||
Context.current());
|
||||
|
||||
Assert.assertNotSame(callerContext, operationContext);
|
||||
Assert.assertTrue(operationContext.isCancelled());
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldCapPoolWaitByRemainingSearchDeadline() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 5_000L)) {
|
||||
fixture.manager.markCollectionLoaded("docs");
|
||||
ExecutorService executor = Executors.newSingleThreadExecutor();
|
||||
CountDownLatch borrowed = new CountDownLatch(1);
|
||||
CountDownLatch release = new CountDownLatch(1);
|
||||
try {
|
||||
Future<?> holder = executor.submit(() ->
|
||||
fixture.manager.withClient(client -> {
|
||||
borrowed.countDown();
|
||||
try {
|
||||
Assert.assertTrue(release.await(
|
||||
AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException(exception);
|
||||
}
|
||||
return null;
|
||||
}));
|
||||
Assert.assertTrue(borrowed.await(
|
||||
AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
|
||||
StoreOptions options = StoreOptions.ofCollectionName("docs");
|
||||
options.setTimeoutMillis(500L);
|
||||
long startedAt = System.nanoTime();
|
||||
try {
|
||||
search(fixture.store, options);
|
||||
Assert.fail("Pool wait must respect the remaining deadline");
|
||||
} catch (RuntimeException expected) {
|
||||
Assert.assertTrue(expected.toString(),
|
||||
expected instanceof StoreTimeoutException);
|
||||
long elapsedMillis = TimeUnit.NANOSECONDS.toMillis(
|
||||
System.nanoTime() - startedAt);
|
||||
Assert.assertTrue("elapsedMillis=" + elapsedMillis,
|
||||
elapsedMillis < 800L);
|
||||
}
|
||||
|
||||
release.countDown();
|
||||
holder.get(AWAIT_SECONDS, TimeUnit.SECONDS);
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
} finally {
|
||||
release.countDown();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldInvalidateOnlyClosedClient() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 2, 2_000L)) {
|
||||
ExecutorService executor = Executors.newSingleThreadExecutor();
|
||||
CountDownLatch firstBorrowed = new CountDownLatch(1);
|
||||
CountDownLatch releaseFirst = new CountDownLatch(1);
|
||||
try {
|
||||
Future<?> holder = executor.submit(() ->
|
||||
fixture.manager.withClient(client -> {
|
||||
firstBorrowed.countDown();
|
||||
try {
|
||||
Assert.assertTrue(releaseFirst.await(
|
||||
AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException(exception);
|
||||
}
|
||||
return null;
|
||||
}));
|
||||
Assert.assertTrue(firstBorrowed.await(
|
||||
AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
fixture.manager.withClient(client -> client);
|
||||
releaseFirst.countDown();
|
||||
holder.get(AWAIT_SECONDS, TimeUnit.SECONDS);
|
||||
Assert.assertEquals(2, fixture.manager.getIdleClientCount());
|
||||
|
||||
AtomicReference<MilvusClientV2> closed = new AtomicReference<>();
|
||||
try {
|
||||
fixture.manager.withClient(client -> {
|
||||
closed.set(client);
|
||||
client.close();
|
||||
throw new IllegalStateException("synthetic closed client");
|
||||
});
|
||||
Assert.fail("The closed-client failure must be propagated");
|
||||
} catch (IllegalStateException expected) {
|
||||
Assert.assertEquals(
|
||||
"synthetic closed client", expected.getMessage());
|
||||
}
|
||||
|
||||
Assert.assertEquals(1, fixture.manager.getIdleClientCount());
|
||||
MilvusClientV2 remaining = fixture.manager.withClient(
|
||||
client -> client);
|
||||
Assert.assertNotSame(closed.get(), remaining);
|
||||
Assert.assertEquals(2, server.connectCalls.get());
|
||||
} finally {
|
||||
releaseFirst.countDown();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldLoadSameCollectionOnlyOnce() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 2, 3_000L)) {
|
||||
LoadGate gate = server.blockLoad("shared");
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
Future<List<Document>> first = executor.submit(() ->
|
||||
search(fixture.store, "shared"));
|
||||
gate.awaitEntered();
|
||||
|
||||
Future<List<Document>> second = executor.submit(() ->
|
||||
search(fixture.store, "shared"));
|
||||
MilvusClientManager.CollectionLoadTicket follower =
|
||||
fixture.manager.beginCollectionLoad("shared");
|
||||
Assert.assertFalse(follower.isLeader());
|
||||
awaitCondition(() -> follower.completion().getNumberOfDependents() > 0);
|
||||
|
||||
Assert.assertEquals(1, server.loadCalls("shared"));
|
||||
Assert.assertEquals(1, fixture.manager.getActiveClientCount());
|
||||
gate.release();
|
||||
|
||||
Assert.assertTrue(first.get(AWAIT_SECONDS, TimeUnit.SECONDS).isEmpty());
|
||||
Assert.assertTrue(second.get(AWAIT_SECONDS, TimeUnit.SECONDS).isEmpty());
|
||||
Assert.assertEquals(1, server.loadCalls("shared"));
|
||||
Assert.assertTrue(fixture.manager.isCollectionLoaded("shared"));
|
||||
} finally {
|
||||
gate.release();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldRemoveLeaderTicketAfterLateCacheHit() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer()) {
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
config.setUri(server.uri());
|
||||
config.setDefaultCollectionName("docs");
|
||||
config.setPoolMinIdlePerKey(0);
|
||||
config.setSearchTimeoutMillis(2_000L);
|
||||
RacingMilvusClientManager manager =
|
||||
new RacingMilvusClientManager(config, "late-hit");
|
||||
MilvusVectorStore store = new MilvusVectorStore(config, manager);
|
||||
try {
|
||||
Assert.assertTrue(search(store, "late-hit").isEmpty());
|
||||
manager.markCollectionUnloaded("late-hit");
|
||||
|
||||
Assert.assertTrue(search(store, "late-hit").isEmpty());
|
||||
|
||||
Assert.assertEquals(1, server.loadCalls("late-hit"));
|
||||
Assert.assertTrue(manager.isCollectionLoaded("late-hit"));
|
||||
} finally {
|
||||
store.close();
|
||||
manager.close();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldLoadDifferentCollectionsInParallel() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 2, 3_000L)) {
|
||||
LoadGate firstGate = server.blockLoad("first");
|
||||
LoadGate secondGate = server.blockLoad("second");
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
Future<List<Document>> first = executor.submit(() ->
|
||||
search(fixture.store, "first"));
|
||||
Future<List<Document>> second = executor.submit(() ->
|
||||
search(fixture.store, "second"));
|
||||
|
||||
firstGate.awaitEntered();
|
||||
secondGate.awaitEntered();
|
||||
Assert.assertEquals(2, server.activeLoads.get());
|
||||
Assert.assertEquals(2, server.maxConcurrentLoads.get());
|
||||
Assert.assertEquals(2, fixture.manager.getActiveClientCount());
|
||||
|
||||
firstGate.release();
|
||||
secondGate.release();
|
||||
Assert.assertTrue(first.get(AWAIT_SECONDS, TimeUnit.SECONDS).isEmpty());
|
||||
Assert.assertTrue(second.get(AWAIT_SECONDS, TimeUnit.SECONDS).isEmpty());
|
||||
} finally {
|
||||
firstGate.release();
|
||||
secondGate.release();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldReelectFollowerAfterLoadLeaderTimesOut() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 3_000L)) {
|
||||
LoadGate gate = server.blockLoad("reelect");
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
StoreOptions shortBudget =
|
||||
StoreOptions.ofCollectionName("reelect");
|
||||
shortBudget.setTimeoutMillis(500L);
|
||||
Future<List<Document>> first = executor.submit(() ->
|
||||
search(fixture.store, shortBudget));
|
||||
gate.awaitEntered();
|
||||
|
||||
Future<List<Document>> second = executor.submit(() ->
|
||||
search(fixture.store, "reelect"));
|
||||
awaitCondition(() -> server.loadCalls("reelect") == 2);
|
||||
gate.release();
|
||||
|
||||
Throwable firstFailure = futureFailure(first);
|
||||
Assert.assertTrue(firstFailure.toString(),
|
||||
firstFailure instanceof StoreTimeoutException);
|
||||
Assert.assertTrue(second.get(
|
||||
AWAIT_SECONDS, TimeUnit.SECONDS).isEmpty());
|
||||
Assert.assertEquals(2, server.loadCalls("reelect"));
|
||||
Assert.assertTrue(
|
||||
fixture.manager.isCollectionLoaded("reelect"));
|
||||
} finally {
|
||||
gate.release();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldCancelWaitingFollowerWithoutBorrowingClient() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 5_000L)) {
|
||||
MilvusClientManager.CollectionLoadTicket leader =
|
||||
fixture.manager.beginCollectionLoad("waiting");
|
||||
CountDownLatch finished = new CountDownLatch(1);
|
||||
AtomicReference<Throwable> failure = new AtomicReference<>();
|
||||
Thread follower = new Thread(() -> {
|
||||
try {
|
||||
search(fixture.store, "waiting");
|
||||
} catch (Throwable exception) {
|
||||
failure.set(exception);
|
||||
} finally {
|
||||
finished.countDown();
|
||||
}
|
||||
}, "milvus-load-follower-test");
|
||||
try {
|
||||
follower.start();
|
||||
awaitCondition(() -> leader.completion().getNumberOfDependents() > 0);
|
||||
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(0, server.connectCalls.get());
|
||||
follower.interrupt();
|
||||
Assert.assertTrue(finished.await(AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
|
||||
Assert.assertNotNull(failure.get());
|
||||
Assert.assertFalse(leader.completion().isDone());
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(0, server.connectCalls.get());
|
||||
} finally {
|
||||
follower.interrupt();
|
||||
fixture.manager.failCollectionLoad(
|
||||
leader, new IllegalStateException("test cleanup"), false);
|
||||
fixture.manager.endCollectionLoad(leader);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldRecoverAfterCollectionLoadFailure() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 2_000L)) {
|
||||
server.failNextLoad("recoverable");
|
||||
LoadGate gate = server.blockLoad("recoverable");
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
Future<List<Document>> first = executor.submit(() ->
|
||||
search(fixture.store, "recoverable"));
|
||||
gate.awaitEntered();
|
||||
Future<List<Document>> second = executor.submit(() ->
|
||||
search(fixture.store, "recoverable"));
|
||||
MilvusClientManager.CollectionLoadTicket follower =
|
||||
fixture.manager.beginCollectionLoad("recoverable");
|
||||
Assert.assertFalse(follower.isLeader());
|
||||
awaitCondition(() ->
|
||||
follower.completion().getNumberOfDependents() > 0);
|
||||
gate.release();
|
||||
|
||||
assertFutureFails(first);
|
||||
assertFutureFails(second);
|
||||
} finally {
|
||||
gate.release();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
|
||||
Assert.assertEquals(1, server.loadCalls("recoverable"));
|
||||
Assert.assertFalse(fixture.manager.isCollectionLoaded("recoverable"));
|
||||
Assert.assertEquals(0, fixture.manager.getActiveClientCount());
|
||||
Assert.assertEquals(1, fixture.manager.getIdleClientCount());
|
||||
|
||||
Assert.assertTrue(search(fixture.store, "recoverable").isEmpty());
|
||||
Assert.assertEquals(2, server.loadCalls("recoverable"));
|
||||
Assert.assertTrue(fixture.manager.isCollectionLoaded("recoverable"));
|
||||
Assert.assertEquals(1, server.connectCalls.get());
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldCancelActiveRpcWhenManagerCloses() throws Exception {
|
||||
FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 5_000L);
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
fixture.manager.markCollectionLoaded("docs");
|
||||
QueryBlock block = server.blockQueries();
|
||||
Future<?> operation = executor.submit(() ->
|
||||
query(fixture.manager));
|
||||
block.awaitEntered();
|
||||
|
||||
Future<?> close = executor.submit(fixture.manager::close);
|
||||
block.awaitCancelled();
|
||||
close.get(AWAIT_SECONDS, TimeUnit.SECONDS);
|
||||
assertFutureFails(operation);
|
||||
|
||||
try {
|
||||
fixture.manager.withClient(client -> null);
|
||||
Assert.fail("A closed manager must reject client borrows");
|
||||
} catch (IllegalStateException expected) {
|
||||
Assert.assertEquals("Milvus client pool is closed", expected.getMessage());
|
||||
}
|
||||
} finally {
|
||||
executor.shutdownNow();
|
||||
fixture.close();
|
||||
server.close();
|
||||
}
|
||||
}
|
||||
|
||||
@Test(timeout = 10_000L)
|
||||
public void shouldCancelActiveRpcWhenManagerReconfigures() throws Exception {
|
||||
try (FakeMilvusServer server = new FakeMilvusServer();
|
||||
Fixture fixture = new Fixture(server, 1, 5_000L)) {
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
try {
|
||||
QueryBlock block = server.blockQueries();
|
||||
Future<?> operation = executor.submit(() ->
|
||||
query(fixture.manager));
|
||||
block.awaitEntered();
|
||||
|
||||
fixture.config.setPoolMaxTotal(2);
|
||||
Future<Boolean> reconfigure = executor.submit(() ->
|
||||
fixture.manager.reconfigureIfNeeded(fixture.config));
|
||||
block.awaitCancelled();
|
||||
|
||||
Assert.assertTrue(reconfigure.get(
|
||||
AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
assertFutureFails(operation);
|
||||
|
||||
server.succeedQueries();
|
||||
query(fixture.manager);
|
||||
} finally {
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private static Object query(MilvusClientManager manager) {
|
||||
return manager.withClient(client -> client.query(
|
||||
QueryReq.builder()
|
||||
.collectionName("docs")
|
||||
.filter("id == \"synthetic-id\"")
|
||||
.outputFields(List.of("id"))
|
||||
.build()
|
||||
));
|
||||
}
|
||||
|
||||
private static List<Document> search(MilvusVectorStore store, String collection) {
|
||||
return search(store, StoreOptions.ofCollectionName(collection));
|
||||
}
|
||||
|
||||
private static List<Document> search(
|
||||
MilvusVectorStore store,
|
||||
StoreOptions options
|
||||
) {
|
||||
SearchWrapper wrapper = new SearchWrapper();
|
||||
wrapper.setWithVector(false);
|
||||
wrapper.eq("id", "synthetic-id");
|
||||
return store.search(wrapper, options);
|
||||
}
|
||||
|
||||
private static void assertSearchFails(MilvusVectorStore store, String collection) {
|
||||
try {
|
||||
search(store, collection);
|
||||
Assert.fail("The synthetic Milvus failure must be propagated");
|
||||
} catch (RuntimeException expected) {
|
||||
Assert.assertNotNull(expected);
|
||||
}
|
||||
}
|
||||
|
||||
private static void assertFutureFails(Future<?> future)
|
||||
throws InterruptedException, TimeoutException {
|
||||
futureFailure(future);
|
||||
}
|
||||
|
||||
private static Throwable futureFailure(Future<?> future)
|
||||
throws InterruptedException, TimeoutException {
|
||||
try {
|
||||
future.get(AWAIT_SECONDS, TimeUnit.SECONDS);
|
||||
Assert.fail("The synthetic Milvus failure must be propagated");
|
||||
} catch (ExecutionException expected) {
|
||||
Assert.assertNotNull(expected.getCause());
|
||||
return expected.getCause();
|
||||
}
|
||||
throw new AssertionError("Expected future to fail");
|
||||
}
|
||||
|
||||
private static void awaitCondition(BooleanSupplier condition)
|
||||
throws InterruptedException, TimeoutException {
|
||||
long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(AWAIT_SECONDS);
|
||||
while (!condition.getAsBoolean()) {
|
||||
if (System.nanoTime() >= deadline) {
|
||||
throw new TimeoutException("Timed out waiting for test condition");
|
||||
}
|
||||
if (Thread.interrupted()) {
|
||||
throw new InterruptedException();
|
||||
}
|
||||
Thread.onSpinWait();
|
||||
}
|
||||
}
|
||||
|
||||
private static final class Fixture implements AutoCloseable {
|
||||
|
||||
private final MilvusVectorStoreConfig config;
|
||||
private final MilvusClientManager manager;
|
||||
private final MilvusVectorStore store;
|
||||
|
||||
private Fixture(FakeMilvusServer server, int poolSize, long searchTimeoutMillis) {
|
||||
config = new MilvusVectorStoreConfig();
|
||||
config.setUri(server.uri());
|
||||
config.setDefaultCollectionName("docs");
|
||||
config.setPoolMaxTotal(poolSize);
|
||||
config.setPoolMaxTotalPerKey(poolSize);
|
||||
config.setPoolMaxIdlePerKey(poolSize);
|
||||
config.setPoolMinIdlePerKey(0);
|
||||
config.setPoolMaxWaitMillis(1_000L);
|
||||
config.setSearchTimeoutMillis(searchTimeoutMillis);
|
||||
manager = new MilvusClientManager(config);
|
||||
store = new MilvusVectorStore(config, manager);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void close() {
|
||||
store.close();
|
||||
manager.close();
|
||||
}
|
||||
}
|
||||
|
||||
private static final class RacingMilvusClientManager
|
||||
extends MilvusClientManager {
|
||||
|
||||
private final String collectionName;
|
||||
private final AtomicInteger observations = new AtomicInteger();
|
||||
|
||||
private RacingMilvusClientManager(
|
||||
MilvusVectorStoreConfig config,
|
||||
String collectionName
|
||||
) {
|
||||
super(config);
|
||||
this.collectionName = collectionName;
|
||||
}
|
||||
|
||||
@Override
|
||||
boolean isCollectionLoaded(String requestedCollectionName) {
|
||||
if (collectionName.equals(requestedCollectionName)) {
|
||||
int observation = observations.getAndIncrement();
|
||||
if (observation == 0) {
|
||||
return false;
|
||||
}
|
||||
if (observation == 1) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return super.isCollectionLoaded(requestedCollectionName);
|
||||
}
|
||||
}
|
||||
|
||||
private static final class QueryBlock {
|
||||
|
||||
private final CountDownLatch entered = new CountDownLatch(1);
|
||||
private final CountDownLatch cancelled = new CountDownLatch(1);
|
||||
|
||||
private void awaitEntered() throws InterruptedException {
|
||||
Assert.assertTrue(entered.await(AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
}
|
||||
|
||||
private void awaitCancelled() throws InterruptedException {
|
||||
Assert.assertTrue(cancelled.await(AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
}
|
||||
}
|
||||
|
||||
private static final class LoadGate {
|
||||
|
||||
private final CountDownLatch entered = new CountDownLatch(1);
|
||||
private final CountDownLatch release = new CountDownLatch(1);
|
||||
|
||||
private void awaitEntered() throws InterruptedException {
|
||||
Assert.assertTrue(entered.await(AWAIT_SECONDS, TimeUnit.SECONDS));
|
||||
}
|
||||
|
||||
private void release() {
|
||||
release.countDown();
|
||||
}
|
||||
}
|
||||
|
||||
private static final class FakeMilvusServer implements AutoCloseable {
|
||||
|
||||
private static final io.milvus.grpc.Status SUCCESS =
|
||||
io.milvus.grpc.Status.newBuilder()
|
||||
.setErrorCode(ErrorCode.Success)
|
||||
.setCode(0)
|
||||
.build();
|
||||
|
||||
private final AtomicInteger connectCalls = new AtomicInteger();
|
||||
private final AtomicInteger queryCalls = new AtomicInteger();
|
||||
private final AtomicInteger activeLoads = new AtomicInteger();
|
||||
private final AtomicInteger maxConcurrentLoads = new AtomicInteger();
|
||||
private final AtomicBoolean querySawDeadline = new AtomicBoolean();
|
||||
private final ConcurrentHashMap<String, AtomicInteger> loadCalls =
|
||||
new ConcurrentHashMap<>();
|
||||
private final ConcurrentHashMap<String, LoadGate> loadGates =
|
||||
new ConcurrentHashMap<>();
|
||||
private final Set<String> loadedCollections = ConcurrentHashMap.newKeySet();
|
||||
private final Set<String> failNextLoads = ConcurrentHashMap.newKeySet();
|
||||
private final ExecutorService rpcExecutor = Executors.newCachedThreadPool();
|
||||
private final Server server;
|
||||
private volatile QueryAction queryAction = QueryAction.SUCCESS;
|
||||
private volatile QueryBlock queryBlock;
|
||||
|
||||
private FakeMilvusServer() throws IOException {
|
||||
server = NettyServerBuilder.forPort(0)
|
||||
.executor(rpcExecutor)
|
||||
.addService(new Service())
|
||||
.build()
|
||||
.start();
|
||||
}
|
||||
|
||||
private String uri() {
|
||||
return "http://127.0.0.1:" + server.getPort();
|
||||
}
|
||||
|
||||
private void succeedQueries() {
|
||||
queryAction = QueryAction.SUCCESS;
|
||||
queryBlock = null;
|
||||
}
|
||||
|
||||
private void failQueriesWithUnavailable() {
|
||||
queryAction = QueryAction.UNAVAILABLE;
|
||||
queryBlock = null;
|
||||
}
|
||||
|
||||
private QueryBlock blockQueries() {
|
||||
QueryBlock block = new QueryBlock();
|
||||
queryBlock = block;
|
||||
queryAction = QueryAction.BLOCK;
|
||||
return block;
|
||||
}
|
||||
|
||||
private LoadGate blockLoad(String collectionName) {
|
||||
LoadGate gate = new LoadGate();
|
||||
loadGates.put(collectionName, gate);
|
||||
return gate;
|
||||
}
|
||||
|
||||
private void failNextLoad(String collectionName) {
|
||||
failNextLoads.add(collectionName);
|
||||
}
|
||||
|
||||
private int loadCalls(String collectionName) {
|
||||
AtomicInteger calls = loadCalls.get(collectionName);
|
||||
return calls == null ? 0 : calls.get();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void close() {
|
||||
for (LoadGate gate : new ArrayList<>(loadGates.values())) {
|
||||
gate.release();
|
||||
}
|
||||
server.shutdownNow();
|
||||
try {
|
||||
server.awaitTermination(AWAIT_SECONDS, TimeUnit.SECONDS);
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
} finally {
|
||||
rpcExecutor.shutdownNow();
|
||||
}
|
||||
}
|
||||
|
||||
private enum QueryAction {
|
||||
SUCCESS,
|
||||
UNAVAILABLE,
|
||||
BLOCK
|
||||
}
|
||||
|
||||
private final class Service extends MilvusServiceGrpc.MilvusServiceImplBase {
|
||||
|
||||
@Override
|
||||
public void connect(
|
||||
ConnectRequest request,
|
||||
StreamObserver<ConnectResponse> observer
|
||||
) {
|
||||
connectCalls.incrementAndGet();
|
||||
observer.onNext(ConnectResponse.newBuilder()
|
||||
.setStatus(SUCCESS)
|
||||
.setIdentifier(1L)
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void listDatabases(
|
||||
ListDatabasesRequest request,
|
||||
StreamObserver<ListDatabasesResponse> observer
|
||||
) {
|
||||
observer.onNext(ListDatabasesResponse.newBuilder()
|
||||
.setStatus(SUCCESS)
|
||||
.addDbNames("default")
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void checkHealth(
|
||||
CheckHealthRequest request,
|
||||
StreamObserver<CheckHealthResponse> observer
|
||||
) {
|
||||
observer.onNext(CheckHealthResponse.newBuilder()
|
||||
.setStatus(SUCCESS)
|
||||
.setIsHealthy(true)
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void getLoadState(
|
||||
GetLoadStateRequest request,
|
||||
StreamObserver<GetLoadStateResponse> observer
|
||||
) {
|
||||
LoadState state = loadedCollections.contains(request.getCollectionName())
|
||||
? LoadState.LoadStateLoaded
|
||||
: LoadState.LoadStateNotLoad;
|
||||
observer.onNext(GetLoadStateResponse.newBuilder()
|
||||
.setStatus(SUCCESS)
|
||||
.setState(state)
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void describeCollection(
|
||||
DescribeCollectionRequest request,
|
||||
StreamObserver<DescribeCollectionResponse> observer
|
||||
) {
|
||||
CollectionSchema schema = CollectionSchema.newBuilder()
|
||||
.setName(request.getCollectionName())
|
||||
.addFields(FieldSchema.newBuilder()
|
||||
.setName("id")
|
||||
.setIsPrimaryKey(true)
|
||||
.setDataType(DataType.VarChar)
|
||||
.build())
|
||||
.build();
|
||||
observer.onNext(DescribeCollectionResponse.newBuilder()
|
||||
.setStatus(SUCCESS)
|
||||
.setCollectionName(request.getCollectionName())
|
||||
.setSchema(schema)
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void loadCollection(
|
||||
LoadCollectionRequest request,
|
||||
StreamObserver<io.milvus.grpc.Status> observer
|
||||
) {
|
||||
String collectionName = request.getCollectionName();
|
||||
loadCalls.computeIfAbsent(
|
||||
collectionName, ignored -> new AtomicInteger()).incrementAndGet();
|
||||
LoadGate gate = loadGates.get(collectionName);
|
||||
if (gate != null) {
|
||||
int active = activeLoads.incrementAndGet();
|
||||
maxConcurrentLoads.accumulateAndGet(active, Math::max);
|
||||
gate.entered.countDown();
|
||||
try {
|
||||
if (!gate.release.await(AWAIT_SECONDS, TimeUnit.SECONDS)) {
|
||||
observer.onError(io.grpc.Status.DEADLINE_EXCEEDED
|
||||
.withDescription("test load gate timed out")
|
||||
.asRuntimeException());
|
||||
return;
|
||||
}
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
observer.onError(io.grpc.Status.CANCELLED
|
||||
.withCause(exception)
|
||||
.asRuntimeException());
|
||||
return;
|
||||
} finally {
|
||||
activeLoads.decrementAndGet();
|
||||
}
|
||||
}
|
||||
if (failNextLoads.remove(collectionName)) {
|
||||
observer.onNext(io.milvus.grpc.Status.newBuilder()
|
||||
.setErrorCode(ErrorCode.UnexpectedError)
|
||||
.setCode(1)
|
||||
.setReason("synthetic load failure")
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
return;
|
||||
}
|
||||
loadedCollections.add(collectionName);
|
||||
observer.onNext(SUCCESS);
|
||||
observer.onCompleted();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void query(
|
||||
QueryRequest request,
|
||||
StreamObserver<QueryResults> observer
|
||||
) {
|
||||
queryCalls.incrementAndGet();
|
||||
querySawDeadline.compareAndSet(
|
||||
false, Context.current().getDeadline() != null);
|
||||
QueryAction action = queryAction;
|
||||
if (action == QueryAction.UNAVAILABLE) {
|
||||
observer.onError(io.grpc.Status.UNAVAILABLE
|
||||
.withDescription("synthetic query failure")
|
||||
.asRuntimeException());
|
||||
return;
|
||||
}
|
||||
if (action == QueryAction.BLOCK) {
|
||||
QueryBlock block = queryBlock;
|
||||
@SuppressWarnings("unchecked")
|
||||
ServerCallStreamObserver<QueryResults> serverObserver =
|
||||
(ServerCallStreamObserver<QueryResults>) observer;
|
||||
serverObserver.setOnCancelHandler(block.cancelled::countDown);
|
||||
block.entered.countDown();
|
||||
return;
|
||||
}
|
||||
observer.onNext(QueryResults.newBuilder()
|
||||
.setStatus(SUCCESS)
|
||||
.setCollectionName(request.getCollectionName())
|
||||
.build());
|
||||
observer.onCompleted();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,190 @@
|
||||
package com.easyagents.store.milvus;
|
||||
|
||||
import com.easyagents.core.document.Document;
|
||||
import com.easyagents.core.store.SearchWrapper;
|
||||
import com.easyagents.core.store.StoreOptions;
|
||||
import io.milvus.v2.service.collection.request.DropCollectionReq;
|
||||
import org.junit.Assert;
|
||||
import org.junit.Assume;
|
||||
import org.junit.Test;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.UUID;
|
||||
import java.util.concurrent.CountDownLatch;
|
||||
import java.util.concurrent.ExecutionException;
|
||||
import java.util.concurrent.ExecutorService;
|
||||
import java.util.concurrent.Executors;
|
||||
import java.util.concurrent.Future;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
|
||||
/**
|
||||
* Opt-in compatibility smoke tests for a real Milvus instance.
|
||||
*/
|
||||
public class MilvusVectorStoreIntegrationTest {
|
||||
|
||||
@Test
|
||||
public void shouldCrudAgainstRealMilvusAndRecoverPooledClients() throws Exception {
|
||||
String uri = System.getenv("MILVUS_TEST_URI");
|
||||
Assume.assumeTrue("MILVUS_TEST_URI is not configured", uri != null && !uri.isBlank());
|
||||
String collectionName = "easy_agents_sdk_2311_" + UUID.randomUUID().toString().replace("-", "");
|
||||
MilvusVectorStoreConfig config = new MilvusVectorStoreConfig();
|
||||
config.setUri(uri);
|
||||
config.setDefaultCollectionName(collectionName);
|
||||
config.setPoolMaxTotal(1);
|
||||
config.setPoolMaxTotalPerKey(1);
|
||||
config.setPoolMaxIdlePerKey(1);
|
||||
config.setPoolMinIdlePerKey(0);
|
||||
config.setPoolMaxWaitMillis(250L);
|
||||
MilvusClientManager manager = new MilvusClientManager(config);
|
||||
MilvusVectorStore store = new MilvusVectorStore(config, manager);
|
||||
StoreOptions options = StoreOptions.ofCollectionName(collectionName);
|
||||
try {
|
||||
Document first = document("chunk-1", "first", 1.0F, 0.0F);
|
||||
Document second = document("chunk-2", "second", 0.0F, 1.0F);
|
||||
Assert.assertTrue(store.store(List.of(first, second), options).isSuccess());
|
||||
|
||||
SearchWrapper nearest = new SearchWrapper();
|
||||
nearest.setVector(new float[] { 1.0F, 0.0F });
|
||||
nearest.setMaxResults(1);
|
||||
List<Document> initial = store.search(nearest, options);
|
||||
Assert.assertEquals(1, initial.size());
|
||||
Assert.assertEquals("chunk-1", String.valueOf(initial.get(0).getId()));
|
||||
|
||||
Document updated = document("chunk-1", "updated", 1.0F, 0.0F);
|
||||
Assert.assertTrue(store.update(List.of(updated), options).isSuccess());
|
||||
Assert.assertEquals("updated", store.search(nearest, options).get(0).getContent());
|
||||
|
||||
Assert.assertTrue(store.delete(List.of("chunk-2"), options).isSuccess());
|
||||
SearchWrapper deleted = new SearchWrapper();
|
||||
deleted.setWithVector(false);
|
||||
deleted.eq("id", "chunk-2");
|
||||
Assert.assertTrue(store.search(deleted, options).isEmpty());
|
||||
|
||||
assertQueryFailureIsNotReportedAsEmpty(store, options);
|
||||
assertPoolExhaustionIsBounded(manager);
|
||||
assertBusinessFailurePreservesClient(manager);
|
||||
Assert.assertTrue(store.checkAvailable());
|
||||
} finally {
|
||||
try {
|
||||
manager.withClient(client -> {
|
||||
client.dropCollection(DropCollectionReq.builder()
|
||||
.collectionName(collectionName)
|
||||
.build());
|
||||
return null;
|
||||
});
|
||||
} finally {
|
||||
store.close();
|
||||
manager.close();
|
||||
}
|
||||
}
|
||||
try {
|
||||
manager.getActiveClientCount();
|
||||
Assert.fail("A closed pool must reject further use");
|
||||
} catch (IllegalStateException expected) {
|
||||
Assert.assertEquals("Milvus client pool is closed", expected.getMessage());
|
||||
}
|
||||
}
|
||||
|
||||
private static void assertQueryFailureIsNotReportedAsEmpty(
|
||||
MilvusVectorStore store,
|
||||
StoreOptions options
|
||||
) {
|
||||
SearchWrapper invalid = new SearchWrapper();
|
||||
invalid.setWithVector(false);
|
||||
invalid.eq("id", "chunk-1");
|
||||
StoreOptions invalidOptions = StoreOptions.ofCollectionName(
|
||||
options.getCollectionName()
|
||||
).partitionName("__missing_partition__");
|
||||
try {
|
||||
store.search(invalid, invalidOptions);
|
||||
Assert.fail("A Milvus query failure must not be reported as an empty result");
|
||||
} catch (RuntimeException expected) {
|
||||
Assert.assertNotNull(expected.getMessage());
|
||||
}
|
||||
}
|
||||
|
||||
private static void assertBusinessFailurePreservesClient(MilvusClientManager manager)
|
||||
throws InterruptedException, ExecutionException {
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
CountDownLatch borrowed = new CountDownLatch(1);
|
||||
CountDownLatch waiterStarted = new CountDownLatch(1);
|
||||
CountDownLatch fail = new CountDownLatch(1);
|
||||
AtomicReference<Object> failedClient = new AtomicReference<>();
|
||||
try {
|
||||
Future<?> failing = executor.submit(() -> {
|
||||
try {
|
||||
manager.withClient(client -> {
|
||||
failedClient.set(client);
|
||||
borrowed.countDown();
|
||||
try {
|
||||
if (!fail.await(2, TimeUnit.SECONDS)) {
|
||||
throw new IllegalStateException("Timed out waiting to fail client");
|
||||
}
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException(exception);
|
||||
}
|
||||
throw new IllegalStateException("synthetic RPC failure");
|
||||
});
|
||||
Assert.fail("The synthetic client failure must be propagated");
|
||||
} catch (IllegalStateException expected) {
|
||||
Assert.assertEquals("synthetic RPC failure", expected.getMessage());
|
||||
}
|
||||
});
|
||||
Assert.assertTrue(borrowed.await(2, TimeUnit.SECONDS));
|
||||
Future<Object> waiting = executor.submit(() -> {
|
||||
waiterStarted.countDown();
|
||||
return manager.withClient(client -> client);
|
||||
});
|
||||
Assert.assertTrue(waiterStarted.await(2, TimeUnit.SECONDS));
|
||||
fail.countDown();
|
||||
failing.get();
|
||||
Assert.assertSame(failedClient.get(), waiting.get());
|
||||
} finally {
|
||||
fail.countDown();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
|
||||
private static void assertPoolExhaustionIsBounded(MilvusClientManager manager)
|
||||
throws InterruptedException, ExecutionException {
|
||||
ExecutorService executor = Executors.newFixedThreadPool(2);
|
||||
CountDownLatch borrowed = new CountDownLatch(1);
|
||||
CountDownLatch release = new CountDownLatch(1);
|
||||
try {
|
||||
Future<?> holder = executor.submit(() -> manager.withClient(client -> {
|
||||
borrowed.countDown();
|
||||
try {
|
||||
if (!release.await(2, TimeUnit.SECONDS)) {
|
||||
throw new IllegalStateException("Timed out waiting to release pooled client");
|
||||
}
|
||||
} catch (InterruptedException exception) {
|
||||
Thread.currentThread().interrupt();
|
||||
throw new IllegalStateException(exception);
|
||||
}
|
||||
return null;
|
||||
}));
|
||||
Assert.assertTrue(borrowed.await(2, TimeUnit.SECONDS));
|
||||
Future<?> waiter = executor.submit(() -> manager.withClient(client -> null));
|
||||
try {
|
||||
waiter.get();
|
||||
Assert.fail("Pool exhaustion must fail after the configured wait");
|
||||
} catch (ExecutionException expected) {
|
||||
Assert.assertNotNull(expected.getCause());
|
||||
}
|
||||
release.countDown();
|
||||
holder.get();
|
||||
} finally {
|
||||
release.countDown();
|
||||
executor.shutdownNow();
|
||||
}
|
||||
}
|
||||
|
||||
private static Document document(String id, String content, float first, float second) {
|
||||
Document document = Document.of(content);
|
||||
document.setId(id);
|
||||
document.setVector(new float[] { first, second });
|
||||
return document;
|
||||
}
|
||||
}
|
||||
7
pom.xml
7
pom.xml
@@ -49,6 +49,7 @@
|
||||
<agentscope.version>1.0.12</agentscope.version>
|
||||
<snakeyaml.version>2.6</snakeyaml.version>
|
||||
<commons-compress.version>1.28.0</commons-compress.version>
|
||||
<milvus.version>2.3.11</milvus.version>
|
||||
</properties>
|
||||
|
||||
|
||||
@@ -129,6 +130,12 @@
|
||||
<version>${commons-compress.version}</version>
|
||||
</dependency>
|
||||
|
||||
<dependency>
|
||||
<groupId>io.milvus</groupId>
|
||||
<artifactId>milvus-sdk-java</artifactId>
|
||||
<version>${milvus.version}</version>
|
||||
</dependency>
|
||||
|
||||
<!--easy-agents dependency management-->
|
||||
<dependency>
|
||||
<groupId>com.easyagents</groupId>
|
||||
|
||||
Reference in New Issue
Block a user