Compare commits

6 Commits

Author SHA1 Message Date
ea67d0519f feat: M28 支持工作流多入边汇聚模式 2026-08-31 15:56:05 +08:00
2d26494775 feat: 增强 Agentic RAG 主动检索引导
- 统一组合用户、知识库与异步工具系统提示词

- 补充知识库调用策略与运行时回归测试
2026-08-29 17:01:44 +08:00
651b292d83 feat: 支持一库一工具 Agentic RAG
- 将知识库声明与 Retriever 绑定为独立 Registration 并注册到 AgentScope Toolkit

- 统一检索分数、最终文档事件与长答案引用语义

- 补齐历史会话、工具类加载器和知识库调用状态回归测试
2026-08-29 15:44:43 +08:00
67b69669c4 fix: 保留模型异常后的会话上下文
- 保存模型失败与无正文取消路径中的 AgentScope 会话

- 补充失败和推理中断场景的上下文恢复测试
2026-08-26 22:56:29 +08:00
6192df4bd1 Merge pull request '发布 v1.1.0' (#2) from develop into main
Reviewed-on: #2
2026-08-20 11:35:40 +08:00
bb37d9d708 Merge pull request 'fix: 修复增加 FAQ 报错的 bug' (#1) from hotfix/embedding_error into main
Reviewed-on: #1
2026-06-14 16:19:44 +08:00
25 changed files with 2241 additions and 471 deletions

View File

@@ -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;
}
/**

View File

@@ -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);
}
/**

View File

@@ -2,291 +2,361 @@ 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);
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.");
}
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;
}
/**
* 创建带事件 sink 的聚合 Knowledge
* 将知识库工具注册到现有 AgentScope Toolkit
*
* @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 context 运行时上下文
* @param toolSpecs 知识库工具定义
* @param toolkit AgentScope Toolkit
* @param toolAdapter 中立工具适配器
* @param approvalCoordinator 工具审批协调器
* @param turnContextHolder 当前运行轮次上下文持有器
* @return 聚合 Knowledge未配置知识库时返回 null
*/
public Knowledge createAggregateKnowledge(AgentRuntimeExecutionContext request,
public void registerTools(AgentRuntimeExecutionContext context,
List<AgentToolSpec> toolSpecs,
Toolkit toolkit,
AgentScopeToolAdapter toolAdapter,
AgentToolApprovalCoordinator approvalCoordinator,
AgentRuntimeTurnContextHolder turnContextHolder) {
if (request.getAgentDefinition().getKnowledgeSpecs().isEmpty()) {
return null;
Objects.requireNonNull(toolkit, "toolkit");
Objects.requireNonNull(toolAdapter, "toolAdapter");
if (toolSpecs == null || toolSpecs.isEmpty()) {
return;
}
return new AggregateKnowledge(request, turnContextHolder);
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));
}
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;
}
/**
* 将运行时文档转换为 AgentScope 文档
* 按知识库 ID 建立运行时绑定索引并拒绝重复绑定
*
* @param documents 运行时文档
* @return AgentScope 文档
* @param registrations 知识库运行时绑定
* @return 以知识库 ID 为键的绑定索引
*/
public List<Document> toDocuments(List<AgentKnowledgeDocument> documents) {
List<Document> converted = new ArrayList<>();
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));
}
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;
}
/**
* 获取 AgentScope 要求的非空文档 ID。
*
* @param document 知识文档
* @return 非空文档 ID
*/
private String safeDocumentId(AgentKnowledgeDocument document) {
if (document.getDocumentId() != null && !document.getDocumentId().isBlank()) {
return document.getDocumentId();
}
if (document.getChunkId() != null && !document.getChunkId().isBlank()) {
return document.getChunkId();
}
return "knowledge-document";
}
/**
* 获取 AgentScope 要求的非空分片 ID。
*
* @param document 知识文档
* @return 非空分片 ID
*/
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";
}
/**
* 获取 AgentScope 要求的非空文档内容。
*
* @param document 知识文档
* @return 文档内容
*/
private String safeContent(AgentKnowledgeDocument document) {
return document.getContent() == null ? "" : document.getContent();
}
/**
* 将检索调用分发到多个知识源的聚合 Knowledge 实现。
*/
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<>());
if (document == null || !passesThreshold(document, knowledgeSpec.getScoreThreshold())) {
continue;
}
emitKnowledgeRetrievalEvent(query, spec, retrievalRequest, result.getDocuments());
for (AgentKnowledgeDocument document : result.getDocuments()) {
preserveKnowledgeMetadata(spec, document);
allDocuments.add(document);
preserveKnowledgeMetadata(knowledgeSpec, document);
finalDocuments.add(document);
}
}
allDocuments.sort(Comparator.comparing(
finalDocuments.sort(Comparator.comparing(
AgentKnowledgeDocument::getScore,
Comparator.nullsLast(Comparator.reverseOrder())
));
if (allDocuments.size() > globalLimit) {
allDocuments = new ArrayList<>(allDocuments.subList(0, globalLimit));
Comparator.nullsLast(Comparator.reverseOrder())));
int limit = Math.max(knowledgeSpec.getLimit(), 1);
if (finalDocuments.size() > limit) {
return new ArrayList<>(finalDocuments.subList(0, limit));
}
return toDocuments(allDocuments);
return finalDocuments;
}
/**
* 在单条文档上保留知识库级元数据
* 判断文档最终分数是否达到绑定阈值
*
* @param spec 知识库声明
* @param document 文档
* @param document 检索文档
* @param scoreThreshold 分数阈值
* @return 达到阈值时为 true
*/
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());
private boolean passesThreshold(AgentKnowledgeDocument document, double scoreThreshold) {
if (scoreThreshold <= 0D) {
return true;
}
return document.getScore() != null && document.getScore() >= scoreThreshold;
}
/**
* 将知识库归属信息合并到文档元数据中。
*
* @param knowledgeSpec 知识库声明
* @param document 检索文档
*/
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", spec.getKnowledgeId());
document.getMetadata().putIfAbsent("knowledgeName", spec.getName());
document.getMetadata().putIfAbsent("knowledgeId", knowledgeSpec.getKnowledgeId());
document.getMetadata().putIfAbsent("knowledgeName", knowledgeSpec.getName());
}
/**
* 发射知识库检索旁路事件,供聊天界面展示检索过程
* 创建与模型最终证据一致的知识库检索事件
*
* <p>知识库检索本身属于 AgentScope RAG 主线路,返回的 Document 会继续进入
* AgentScope 的上下文注入流程;这里发出的 {@code KNOWLEDGE_RETRIEVAL}
* 只是旁路告知调用方,不会回写 memory也不会参与模型消息序列。</p>
*
* @param query 查询
* @param spec 知识库声明
* @param toolContext 工具执行上下文
* @param knowledgeSpec 知识库声明
* @param retrievalRequest 检索请求
* @param documents 检索文档
* @param documents 最终文档
* @return 检索旁路事件
*/
private void emitKnowledgeRetrievalEvent(String query,
AgentKnowledgeSpec spec,
private AgentRuntimeEvent retrievalEvent(AgentToolContext toolContext,
AgentKnowledgeSpec knowledgeSpec,
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"));
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 == null ? 0 : documents.size());
event.getPayload().put("documentCount", 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);
event.getMetadata().put("toolName", AgentKnowledgeToolNames.build(knowledgeSpec.getRuntimeName()));
event.getMetadata().put("knowledgeRuntimeName", knowledgeSpec.getRuntimeName());
return event;
}
/**
* 构建用于事件展示的命中片段,保留前端引注需要的原始 chunk 内容
* 将最终文档转换为 UI 和完成消息引用可消费的稳定摘要
*
* @param documents 检索文档
* @return 命中片段列表
* @param documents 最终文档
* @return 文档摘要列表
*/
private List<Map<String, Object>> documentSummaries(List<AgentKnowledgeDocument> documents) {
List<Map<String, Object>> summaries = new ArrayList<>();
@@ -306,5 +376,78 @@ public class AgentScopeKnowledgeAdapter {
}
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);
}
}

View File

@@ -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));
}
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) {

View File

@@ -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,6 +422,14 @@ public class AgentScopeToolAdapter {
* @return 工具结果块
*/
private ToolResultBlock invokeTool(ToolCallParam param, Map<String, Object> input) {
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);
@@ -411,6 +439,11 @@ public class AgentScopeToolAdapter {
emit(toolResultEvent(block));
}
return block;
} finally {
if (switchClassLoader) {
currentThread.setContextClassLoader(originalClassLoader);
}
}
}
/**

View File

@@ -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()

View File

@@ -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) {

View File

@@ -1,10 +0,0 @@
package com.easyagents.agent.runtime.knowledge;
/**
* 知识库检索策略。
*/
public enum AgentKnowledgePolicy {
AGENTIC,
GENERIC,
DISABLED
}

View File

@@ -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;
}
}

View File

@@ -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;
}
/**
* 获取限制数量。
*

View File

@@ -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);
}
}
}

View File

@@ -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();
}
/**
* 计算答案与候选片段之间的文本支撑分。
*

View File

@@ -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;
}
}

View File

@@ -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) {
}
}

View File

@@ -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 ->
Assert.assertTrue("knowledgeInvoked=" + knowledgeInvoked.get() + ", events="
+ events.stream().map(AgentRuntimeEvent::getEventType).toList(),
events.stream().anyMatch(event ->
event.getEventType() == AgentRuntimeEventType.KNOWLEDGE_RETRIEVAL));
return;
}
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());
}
/**
* 创建每次模型调用仅返回下一条预设响应的运行时。
*

View File

@@ -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();
}
}

View File

@@ -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())));
}
/**
* 空答案或空候选不应返回引用。
*/

View File

@@ -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;
}
}

View File

@@ -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)

View File

@@ -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<>();

View File

@@ -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'");
}
}

View File

@@ -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);

View File

@@ -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;
}
}

View File

@@ -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;

View File

@@ -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);