diff --git a/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/AgentResumeRequest.java b/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/AgentResumeRequest.java index b0e7494..e042bfc 100644 --- a/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/AgentResumeRequest.java +++ b/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/AgentResumeRequest.java @@ -35,7 +35,9 @@ public class AgentResumeRequest { * *

该字段仅供服务端集成层使用。普通调用方不应设置该标记;设置后 runtime 会跳过 * 当前进程内 {@code AgentToolApprovalCoordinator} 的 token 存在性校验,用于服务重启或跨节点后 - * 从 AgentScope session 中继续 pending tool。

+ * 从 AgentScope session 中继续 pending tool。批准请求必须在 metadata 中提供 + * {@code toolCallId/toolName/toolInput},多个调用使用 {@code approvedToolCalls} 列表, + * 以便 runtime 将持久化审批结果绑定到实际工具调用。

*/ private boolean trusted; diff --git a/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/agentscope/AgentScopeReActRuntime.java b/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/agentscope/AgentScopeReActRuntime.java index 7530e33..1dcbdef 100644 --- a/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/agentscope/AgentScopeReActRuntime.java +++ b/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/agentscope/AgentScopeReActRuntime.java @@ -9,7 +9,9 @@ import com.easyagents.agent.runtime.event.observer.AgentRuntimeErrorObserver; import com.easyagents.agent.runtime.event.observer.ReasoningLifecycleObserver; import com.easyagents.agent.runtime.event.observer.SkillExecutionObserver; import com.easyagents.agent.runtime.event.observer.ToolExecutionObserver; +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; @@ -153,6 +155,9 @@ public class AgentScopeReActRuntime implements AgentRuntime { */ @Override public void close() { + if (approvalCoordinator != null) { + approvalCoordinator.cancelAll("Agent runtime has been closed."); + } closeMcpClients(); initialized.set(false); } @@ -188,17 +193,27 @@ public class AgentScopeReActRuntime implements AgentRuntime { return Flux.error(new AgentRuntimeException("Agent runtime is already streaming.")); } AgentRuntimeExecutionContext executionContext = createResumeExecutionContext(request); + AgentToolApprovalResolution resolution = null; try { - if (!request.isTrusted()) { - approvalCoordinator.consume(request); + if (request.isTrusted()) { + approvalCoordinator.authorizeTrustedExecution(request); + } else { + resolution = approvalCoordinator.resolve(request); } } catch (RuntimeException error) { running.set(false); throw error; } - // 审批拒绝 - if (!request.isApproved()) { - executionContext.setCancelReason(request.getRejectReason()); + if (resolution != null + && resolution.getStatus() == AgentToolApprovalResolution.Status.WAITING) { + return waitingForRemainingApprovals(executionContext, resolution); + } + if (!request.isApproved() + || resolution != null + && (resolution.getStatus() == AgentToolApprovalResolution.Status.REJECTED + || resolution.getStatus() == AgentToolApprovalResolution.Status.EXPIRED)) { + String cancelReason = resolution == null ? request.getRejectReason() : resolution.getReason(); + executionContext.setCancelReason(cancelReason); return Flux.defer(() -> { saveSession(); return Flux.just(started(executionContext), cancelled(executionContext)); @@ -240,8 +255,6 @@ public class AgentScopeReActRuntime implements AgentRuntime { AtomicReference finalMessage = new AtomicReference<>(); // HITL 暂停事件。被设置后,本轮以 SUSPENDED 挂起而不是 COMPLETED 结束。 AtomicReference suspendedEvent = new AtomicReference<>(); - // 本轮 HITL 待审批项来自旁路交互事件,最终会合并进 SUSPENDED 挂起事件。 - List> pendingApprovals = new CopyOnWriteArrayList<>(); // 知识库引注。 Map knowledgeReferences = new LinkedHashMap<>(); // 流式输出归一化,防止出现累计快照的重复输出。 @@ -250,8 +263,7 @@ public class AgentScopeReActRuntime implements AgentRuntime { AtomicBoolean cancelled = new AtomicBoolean(false); // 旁线路监察事件流式输出。 Flux sideEventFlux = sideEvents.asFlux() - .doOnNext(event -> updateKnowledgeReferences(knowledgeReferences, event)) - .doOnNext(event -> updatePendingApprovals(pendingApprovals, event)); + .doOnNext(event -> updateKnowledgeReferences(knowledgeReferences, event)); // 主线路 agent 交互。resume 场景会传入空列表,让 AgentScope 从 pending tool 继续执行。 Flux mainEventFlux = agent.stream(inputSupplier.get(), streamOptions()) .timeout(executionContext.getAgentDefinition().getExecutionOptions().getTimeout()) @@ -271,9 +283,8 @@ public class AgentScopeReActRuntime implements AgentRuntime { .concatWith(Flux.defer(() -> { AgentRuntimeEvent suspended = suspendedEvent.get(); if (suspended != null) { - // 触发 hitl 审批事件,暂时挂起。 - suspended.getPayload().put("pendingApprovals", pendingApprovals); - return Flux.just(suspended); + // SUSPENDED 已在主线路中输出,结束阶段不再重复发送。 + return Flux.empty(); } return Flux.just(completed(executionContext, finalText.toString(), finalMessage.get(), knowledgeReferences)); @@ -542,9 +553,10 @@ public class AgentScopeReActRuntime implements AgentRuntime { if (sourceEvent.getMessage() != null) { event.setMessage(messageAdapter.toAgentMessage(sourceEvent.getMessage())); } - event.getPayload().put("reason", context.getMetadata().getOrDefault("hitlSuspendReason", "TOOL_APPROVAL_REQUIRED")); - Object pendingApprovals = context.getMetadata().get("hitlPendingApprovals"); - event.getPayload().put("pendingApprovals", pendingApprovals instanceof List list ? list : List.of()); + event.getPayload().put("reason", "TOOL_APPROVAL_REQUIRED"); + event.getPayload().put("pendingApprovals", approvalCoordinator.pendingStates(context.getSessionId()).stream() + .map(this::pendingApprovalPayload) + .toList()); event.getMetadata().put("source", "AGENTSCOPE_STREAM"); event.getMetadata().put("generateReason", sourceEvent.getMessage() == null ? GenerateReason.REASONING_STOP_REQUESTED.name() @@ -552,6 +564,46 @@ public class AgentScopeReActRuntime implements AgentRuntime { return event; } + /** + * 在同一审批批次仍有未决工具时返回挂起事件,并保持 AgentScope pending tools 不执行。 + * + * @param context 本轮恢复上下文 + * @param resolution 审批批次决议 + * @return 开始与挂起事件流 + */ + private Flux waitingForRemainingApprovals(AgentRuntimeExecutionContext context, + AgentToolApprovalResolution resolution) { + AgentRuntimeEvent suspended = base(context, AgentRuntimeEventType.SUSPENDED); + suspended.getPayload().put("reason", "TOOL_APPROVAL_REQUIRED"); + suspended.getPayload().put("pendingApprovals", resolution.getRemainingStates().stream() + .map(this::pendingApprovalPayload) + .toList()); + suspended.getMetadata().put("source", "APPROVAL_COORDINATOR"); + suspended.getMetadata().put("approvalStatus", AgentToolApprovalResolution.Status.WAITING.name()); + return Flux.just(started(context), suspended) + .doOnNext(event -> context.getConversationRecorder().record(context, event)) + .doFinally(signalType -> cleanupTurn()); + } + + /** + * 将待审批状态转换为前端可消费的稳定字段。 + * + * @param state 待审批状态 + * @return 待审批载荷 + */ + private Map pendingApprovalPayload(AgentPendingState state) { + Map payload = new LinkedHashMap<>(); + payload.put("resumeToken", state.getResumeToken().getValue()); + payload.put("toolCallId", state.getToolCallId()); + payload.put("toolName", state.getToolName()); + payload.put("toolInput", state.getToolInput()); + payload.put("input", state.getToolInput()); + payload.put("approvalPrompt", state.getApprovalPrompt()); + payload.put("approvalMetadata", state.getMetadata()); + payload.put("expiresAt", state.getExpiresAt() == null ? null : state.getExpiresAt().toString()); + return payload; + } + /** * 生成开始事件。 * @@ -742,6 +794,7 @@ public class AgentScopeReActRuntime implements AgentRuntime { * 清理本轮状态。 */ private void cleanupTurn() { + approvalCoordinator.clearExecutionAuthorizations(); turnContextHolder.clear(); running.set(false); } @@ -822,26 +875,6 @@ public class AgentScopeReActRuntime implements AgentRuntime { } } - /** - * 从工具审批旁路事件中收集本轮待审批项。 - * - * @param pendingApprovals 待审批项集合 - * @param event 运行时事件 - */ - private void updatePendingApprovals(List> pendingApprovals, AgentRuntimeEvent event) { - if (event == null || event.getEventType() != AgentRuntimeEventType.TOOL_APPROVAL_REQUIRED) { - return; - } - Map approval = new LinkedHashMap<>(); - approval.put("resumeToken", event.getPayload().get("resumeToken")); - approval.put("toolCallId", event.getPayload().get("toolCallId")); - approval.put("toolName", event.getPayload().get("toolName")); - approval.put("toolInput", event.getPayload().get("toolInput")); - approval.put("expiresAt", event.getPayload().get("expiresAt")); - approval.put("approvalPrompt", event.getPayload().get("approvalPrompt")); - pendingApprovals.add(approval); - } - /** * 从知识库旁路事件中收集本轮候选引用。 * diff --git a/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/event/interceptor/ToolHitlInterceptor.java b/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/event/interceptor/ToolHitlInterceptor.java index 4012f8e..b72bbc2 100644 --- a/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/event/interceptor/ToolHitlInterceptor.java +++ b/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/event/interceptor/ToolHitlInterceptor.java @@ -1,6 +1,7 @@ package com.easyagents.agent.runtime.event.interceptor; import com.easyagents.agent.runtime.AgentRuntimeExecutionContext; +import com.easyagents.agent.runtime.AgentRuntimeException; import com.easyagents.agent.runtime.event.AgentRuntimeEvent; import com.easyagents.agent.runtime.event.AgentRuntimeEventBridge; import com.easyagents.agent.runtime.event.AgentRuntimeEventType; @@ -11,6 +12,8 @@ import com.easyagents.agent.runtime.hitl.AgentToolApprovalRequest; import com.easyagents.agent.runtime.tool.AgentToolSpec; import io.agentscope.core.hook.HookEvent; import io.agentscope.core.hook.PostReasoningEvent; +import io.agentscope.core.hook.PreActingEvent; +import io.agentscope.core.message.ContentBlock; import io.agentscope.core.message.Msg; import io.agentscope.core.message.ToolUseBlock; import reactor.core.publisher.Mono; @@ -21,25 +24,28 @@ import java.util.ArrayList; import java.util.LinkedHashMap; import java.util.List; import java.util.Map; +import java.util.Objects; +import java.util.UUID; import java.util.function.Function; import java.util.stream.Collectors; /** * 工具 HITL 主线路干预器。 * - *

本 interceptor 专门处理“工具执行前人工审批”。监听 AgentScope 原生 - * {@link PostReasoningEvent}

+ *

本 interceptor 专门处理“工具执行前人工审批”。通过 AgentScope 原生 + * {@link PostReasoningEvent} 建立审批批次,并在 {@link PreActingEvent} 消费一次性执行授权。

* - *

这里包含两类动作: + *

这里包含三类动作: *

    *
  • 主线路干预:发现待审批工具后调用 {@link PostReasoningEvent#stopAgent()}, * 让 AgentScope 返回当前带 ToolUseBlock 的消息并暂停工具执行。
  • + *
  • 执行前校验:按工具调用身份消费一次性执行授权,阻止未批准或被篡改的调用。
  • *
  • 旁路交互事件:通过 {@link AgentRuntimeEventBridge} 发出 * {@link AgentRuntimeEventType#TOOL_APPROVAL_REQUIRED},通知调用方展示审批交互。
  • *
* - *

注意:本 interceptor 不执行工具、不写入 AgentScope memory/session,也不实现恢复。 - * 后续 resume 流程应基于 AgentScope pending tool 状态继续调用 agent stream/call。

+ *

注意:本 interceptor 不执行工具。后续 resume 流程应基于 AgentScope pending tool + * 状态继续调用 agent stream/call,实际工具执行仍由 AgentScope Toolkit 完成。

*/ public class ToolHitlInterceptor implements AgentRuntimeInterceptor { @@ -76,6 +82,8 @@ public class ToolHitlInterceptor implements AgentRuntimeInterceptor { public Mono intercept(T event) { if (event instanceof PostReasoningEvent postReasoningEvent) { interceptPostReasoning(postReasoningEvent); + } else if (event instanceof PreActingEvent preActingEvent) { + interceptPreActing(preActingEvent); } return Mono.just(event); } @@ -93,47 +101,191 @@ public class ToolHitlInterceptor implements AgentRuntimeInterceptor { return 50; } + /** + * 归一化待审批调用、创建审批批次并暂停 AgentScope。 + * + * @param event 推理完成事件 + */ private void interceptPostReasoning(PostReasoningEvent event) { - Msg reasoningMessage = event.getReasoningMessage(); + Msg reasoningMessage = normalizeApprovalToolUses(event.getReasoningMessage()); if (reasoningMessage == null) { return; } + if (reasoningMessage != event.getReasoningMessage()) { + event.setReasoningMessage(reasoningMessage); + } List approvalRequiredTools = approvalRequiredTools(reasoningMessage); if (approvalRequiredTools.isEmpty()) { return; } - List> pendingApprovals = new ArrayList<>(); + String approvalBatchId = approvalBatchId(reasoningMessage); for (ToolUseBlock toolUse : approvalRequiredTools) { AgentToolSpec toolSpec = toolSpecs.get(toolUse.getName()); - AgentPendingState pendingState = registerPendingState(toolSpec, toolUse); + AgentPendingState pendingState = registerPendingState(toolSpec, toolUse, approvalBatchId); + if (pendingState.getEventId() != null && !pendingState.getEventId().isBlank()) { + continue; + } AgentRuntimeEvent approvalEvent = toolApprovalRequiredEvent(toolSpec, toolUse, pendingState); pendingState.setEventId(approvalEvent.getEventId()); - pendingApprovals.add(pendingApprovalPayload(pendingState, toolUse)); eventBridge.emit(approvalEvent); } - AgentRuntimeExecutionContext context = eventBridge.executionContext(); - if (context != null) { - context.getMetadata().put("hitlSuspended", true); - context.getMetadata().put("hitlSuspendReason", "TOOL_APPROVAL_REQUIRED"); - context.getMetadata().put("hitlPendingApprovals", pendingApprovals); - } event.stopAgent(); } + /** + * 在工具实际执行前消费与调用身份绑定的一次性授权。 + * + * @param event 工具执行前事件 + */ + private void interceptPreActing(PreActingEvent event) { + ToolUseBlock toolUse = event.getToolUse(); + AgentToolSpec toolSpec = toolUse == null ? null : toolSpecs.get(toolUse.getName()); + if (toolSpec == null || !toolSpec.isApprovalRequired()) { + return; + } + // 执行授权与 toolCallId、工具名称及入参同时绑定,并且只能消费一次。 + approvalCoordinator.consumeExecutionAuthorization( + toolUse.getId(), + toolUse.getName(), + toolUse.getInput()); + } + + /** + * 为缺少ID的审批调用补充稳定ID,并按 toolCallId 去除同轮重放调用。 + * + * @param reasoningMessage 原始推理消息 + * @return 归一化后的推理消息 + */ + private Msg normalizeApprovalToolUses(Msg reasoningMessage) { + if (reasoningMessage == null || reasoningMessage.getContent() == null + || reasoningMessage.getContent().isEmpty()) { + return reasoningMessage; + } + List normalizedContent = new ArrayList<>(reasoningMessage.getContent().size()); + Map seenApprovalCalls = new LinkedHashMap<>(); + boolean changed = false; + for (int contentIndex = 0; contentIndex < reasoningMessage.getContent().size(); contentIndex++) { + ContentBlock block = reasoningMessage.getContent().get(contentIndex); + if (!(block instanceof ToolUseBlock toolUse) || !isApprovalRequired(toolUse)) { + normalizedContent.add(block); + continue; + } + ToolUseBlock normalizedToolUse = toolUse; + if (toolUse.getId() == null || toolUse.getId().isBlank()) { + normalizedToolUse = copyWithId(toolUse, stableToolCallId(reasoningMessage, contentIndex)); + changed = true; + } + ToolCallSignature signature = new ToolCallSignature( + normalizedToolUse.getName(), normalizedToolUse.getInput()); + ToolCallSignature existing = seenApprovalCalls.putIfAbsent(normalizedToolUse.getId(), signature); + if (existing != null) { + if (!existing.equals(signature)) { + throw new AgentRuntimeException( + "Duplicate toolCallId is bound to a different tool call: " + normalizedToolUse.getId()); + } + // 相同 toolCallId 表示同一协议调用被重复返回,只保留第一次出现。 + changed = true; + continue; + } + normalizedContent.add(normalizedToolUse); + } + if (!changed) { + return reasoningMessage; + } + return Msg.builder() + .id(reasoningMessage.getId()) + .name(reasoningMessage.getName()) + .role(reasoningMessage.getRole()) + .content(normalizedContent) + .metadata(reasoningMessage.getMetadata()) + .timestamp(reasoningMessage.getTimestamp()) + .build(); + } + + /** + * 为缺少调用ID的工具生成跨同一推理消息重放稳定的调用ID。 + * + * @param reasoningMessage 推理消息 + * @param contentIndex 工具块在消息内容中的位置 + * @return 稳定工具调用ID + */ + private String stableToolCallId(Msg reasoningMessage, int contentIndex) { + String messageId = reasoningMessage.getId(); + if (messageId == null || messageId.isBlank()) { + return "hitl-" + UUID.randomUUID(); + } + return "hitl-" + messageId + "-" + contentIndex; + } + + /** + * 为同一推理消息生成稳定审批批次ID。 + * + * @param reasoningMessage 推理消息 + * @return 审批批次ID + */ + private String approvalBatchId(Msg reasoningMessage) { + String messageId = reasoningMessage == null ? null : reasoningMessage.getId(); + if (messageId == null || messageId.isBlank()) { + return "hitl-batch-" + UUID.randomUUID(); + } + return "hitl-batch-" + messageId; + } + + /** + * 判断工具调用是否要求人工审批。 + * + * @param toolUse 工具调用 + * @return 要求审批时为 true + */ + private boolean isApprovalRequired(ToolUseBlock toolUse) { + AgentToolSpec toolSpec = toolUse == null ? null : toolSpecs.get(toolUse.getName()); + return toolSpec != null && toolSpec.isApprovalRequired(); + } + + /** + * 复制工具调用并替换调用ID。 + * + * @param toolUse 原始工具调用 + * @param toolCallId 新工具调用ID + * @return 新工具调用 + */ + private ToolUseBlock copyWithId(ToolUseBlock toolUse, String toolCallId) { + return ToolUseBlock.builder() + .id(toolCallId) + .name(toolUse.getName()) + .input(toolUse.getInput()) + .content(toolUse.getContent()) + .metadata(toolUse.getMetadata()) + .build(); + } + + /** + * 获取推理消息中要求审批的工具调用。 + * + * @param reasoningMessage 推理消息 + * @return 待审批工具调用 + */ private List approvalRequiredTools(Msg reasoningMessage) { List toolUses = reasoningMessage.getContentBlocks(ToolUseBlock.class); if (toolUses == null || toolUses.isEmpty()) { return List.of(); } return toolUses.stream() - .filter(toolUse -> { - AgentToolSpec toolSpec = toolUse == null ? null : toolSpecs.get(toolUse.getName()); - return toolSpec != null && toolSpec.isApprovalRequired(); - }) + .filter(this::isApprovalRequired) .toList(); } - private AgentPendingState registerPendingState(AgentToolSpec toolSpec, ToolUseBlock toolUse) { + /** + * 注册审批批次中的待审批状态。 + * + * @param toolSpec 工具声明 + * @param toolUse 工具调用 + * @param approvalBatchId 审批批次ID + * @return 待审批状态 + */ + private AgentPendingState registerPendingState(AgentToolSpec toolSpec, + ToolUseBlock toolUse, + String approvalBatchId) { AgentRuntimeExecutionContext context = eventBridge.executionContext(); AgentToolApprovalRequest approvalRequest = toolSpec.getApprovalRequest(); Duration timeout = approvalRequest == null || approvalRequest.getTimeout() == null @@ -156,9 +308,18 @@ public class ToolHitlInterceptor implements AgentRuntimeInterceptor { approvalPrompt(approvalRequest), toolUse.getInput(), metadata, - Instant.now().plus(timeout)); + Instant.now().plus(timeout), + approvalBatchId); } + /** + * 构建工具审批请求事件。 + * + * @param toolSpec 工具声明 + * @param toolUse 工具调用 + * @param pendingState 待审批状态 + * @return 审批请求事件 + */ private AgentRuntimeEvent toolApprovalRequiredEvent(AgentToolSpec toolSpec, ToolUseBlock toolUse, AgentPendingState pendingState) { @@ -179,6 +340,13 @@ public class ToolHitlInterceptor implements AgentRuntimeInterceptor { return event; } + /** + * 构建待审批工具的稳定事件载荷。 + * + * @param pendingState 待审批状态 + * @param toolUse 工具调用 + * @return 待审批载荷 + */ private Map pendingApprovalPayload(AgentPendingState pendingState, ToolUseBlock toolUse) { Map payload = new LinkedHashMap<>(); payload.put("resumeToken", pendingState.getResumeToken().getValue()); @@ -191,6 +359,12 @@ public class ToolHitlInterceptor implements AgentRuntimeInterceptor { return payload; } + /** + * 将工具展示元数据补充到审批事件载荷。 + * + * @param payload 审批事件载荷 + * @param toolSpec 工具声明 + */ private void enrichToolPayload(Map payload, AgentToolSpec toolSpec) { if (toolSpec == null || toolSpec.getMetadata() == null || toolSpec.getMetadata().isEmpty()) { return; @@ -203,12 +377,25 @@ public class ToolHitlInterceptor implements AgentRuntimeInterceptor { putIfPresent(payload, metadata, "mcpTitle"); } + /** + * 在元数据包含指定字段时复制到事件载荷。 + * + * @param payload 事件载荷 + * @param metadata 工具元数据 + * @param key 字段名 + */ private void putIfPresent(Map payload, Map metadata, String key) { if (metadata.containsKey(key)) { payload.put(key, metadata.get(key)); } } + /** + * 获取审批提示文案。 + * + * @param approvalRequest 审批配置 + * @return 审批提示文案 + */ private String approvalPrompt(AgentToolApprovalRequest approvalRequest) { if (approvalRequest != null && approvalRequest.getApprovalPrompt() != null @@ -217,4 +404,51 @@ public class ToolHitlInterceptor implements AgentRuntimeInterceptor { } return "是否批准执行该工具?"; } + + /** + * 工具名称与输入组成的调用身份校验值。 + */ + private static final class ToolCallSignature { + private final String toolName; + private final Map toolInput; + + /** + * 创建工具调用身份校验值。 + * + * @param toolName 工具名称 + * @param toolInput 工具入参 + */ + private ToolCallSignature(String toolName, Map toolInput) { + this.toolName = toolName; + this.toolInput = toolInput == null ? Map.of() : new LinkedHashMap<>(toolInput); + } + + /** + * 比较工具调用语义是否一致。 + * + * @param object 待比较对象 + * @return 语义一致时为 true + */ + @Override + public boolean equals(Object object) { + if (this == object) { + return true; + } + if (!(object instanceof ToolCallSignature that)) { + return false; + } + return Objects.equals(toolName, that.toolName) + && Objects.equals(toolInput, that.toolInput); + } + + /** + * 计算工具调用语义哈希。 + * + * @return 哈希值 + */ + @Override + public int hashCode() { + return Objects.hash(toolName, toolInput); + } + } } diff --git a/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/hitl/AgentToolApprovalCoordinator.java b/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/hitl/AgentToolApprovalCoordinator.java index 340b527..db1ca3e 100644 --- a/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/hitl/AgentToolApprovalCoordinator.java +++ b/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/hitl/AgentToolApprovalCoordinator.java @@ -5,18 +5,30 @@ import com.easyagents.agent.runtime.AgentRuntimeException; import reactor.core.publisher.Mono; import java.time.Instant; +import java.util.ArrayList; +import java.util.LinkedHashMap; +import java.util.LinkedHashSet; +import java.util.List; import java.util.Map; import java.util.Objects; +import java.util.Set; import java.util.concurrent.CompletableFuture; -import java.util.concurrent.ConcurrentHashMap; /** * 工具审批协调器。 */ public class AgentToolApprovalCoordinator { + /** 是否启用内存审批协调。 */ private final boolean enabled; - private final Map approvals = new ConcurrentHashMap<>(); + /** 恢复令牌到待审批项的索引。 */ + private final Map approvals = new LinkedHashMap<>(); + /** 审批批次ID到批次状态的索引。 */ + private final Map approvalBatches = new LinkedHashMap<>(); + /** 工具调用ID到恢复令牌的唯一索引。 */ + private final Map tokensByToolCallId = new LinkedHashMap<>(); + /** 工具调用ID到一次性执行授权的索引。 */ + private final Map executionAuthorizations = new LinkedHashMap<>(); /** * 创建已启用的协调器。 @@ -66,6 +78,54 @@ public class AgentToolApprovalCoordinator { Map toolInput, Map metadata, Instant expiresAt) { + return register(sessionId, agentId, toolCallId, toolName, approvalPrompt, toolInput, + metadata, expiresAt, null); + } + + /** + * 注册一个属于指定审批批次的待审批请求。 + * + *

同一批次内的全部工具调用均批准后,协调器才会签发执行授权。任意一项拒绝或 + * 过期都会关闭整个批次,避免未批准调用跟随已批准调用一起恢复执行。

+ * + * @param sessionId 会话ID + * @param agentId 智能体ID + * @param toolCallId 工具调用ID + * @param toolName 工具名称 + * @param approvalPrompt 审批文案 + * @param toolInput 工具入参 + * @param metadata 元数据 + * @param expiresAt 过期时间 + * @param approvalBatchId 审批批次ID;为空时创建单调用批次 + * @return 待审批状态 + */ + public synchronized AgentPendingState register(String sessionId, + String agentId, + String toolCallId, + String toolName, + String approvalPrompt, + Map toolInput, + Map metadata, + Instant expiresAt, + String approvalBatchId) { + if (enabled && (toolCallId == null || toolCallId.isBlank())) { + throw new AgentRuntimeException("Approval-required tool call must include toolCallId."); + } + if (enabled && toolCallId != null && !toolCallId.isBlank()) { + String existingToken = tokensByToolCallId.get(toolCallId); + PendingApproval existing = existingToken == null ? null : approvals.get(existingToken); + if (existing != null && !isExpired(existing.state)) { + if (!Objects.equals(existing.state.getToolName(), toolName) + || !Objects.equals(existing.state.getToolInput(), toolInput)) { + throw new AgentRuntimeException( + "Duplicate toolCallId is bound to a different tool call: " + toolCallId); + } + return existing.state; + } + if (existing != null) { + closeBatch(existing.batchId, "审批请求已过期。"); + } + } AgentPendingState state = new AgentPendingState(); state.setSessionId(sessionId); state.setAgentId(agentId); @@ -73,13 +133,26 @@ public class AgentToolApprovalCoordinator { state.setToolName(toolName); state.setApprovalPrompt(approvalPrompt); state.setToolInput(toolInput); - state.setMetadata(metadata); state.setExpiresAt(expiresAt); + String token = state.getResumeToken().getValue(); + String effectiveBatchId = approvalBatchId == null || approvalBatchId.isBlank() + ? token + : approvalBatchId; + Map effectiveMetadata = metadata == null + ? new LinkedHashMap<>() + : new LinkedHashMap<>(metadata); + effectiveMetadata.put("approvalBatchId", effectiveBatchId); + state.setMetadata(effectiveMetadata); if (enabled) { - String token = state.getResumeToken().getValue(); - PendingApproval pendingApproval = new PendingApproval(state, new CompletableFuture<>()); + PendingApproval pendingApproval = new PendingApproval( + state, effectiveBatchId, new CompletableFuture<>()); approvals.put(token, pendingApproval); - pendingApproval.future.whenComplete((response, error) -> approvals.remove(token)); + ApprovalBatch batch = approvalBatches.computeIfAbsent(effectiveBatchId, ApprovalBatch::new); + batch.tokens.add(token); + batch.members.put(token, pendingApproval); + if (toolCallId != null && !toolCallId.isBlank()) { + tokensByToolCallId.put(toolCallId, token); + } } return state; } @@ -90,7 +163,7 @@ public class AgentToolApprovalCoordinator { * @param resumeToken 恢复令牌 * @return 恢复请求 */ - public Mono await(AgentResumeToken resumeToken) { + public synchronized Mono await(AgentResumeToken resumeToken) { if (!enabled) { AgentResumeRequest request = new AgentResumeRequest(); request.setResumeToken(resumeToken); @@ -104,35 +177,211 @@ public class AgentToolApprovalCoordinator { if (pendingApproval == null) { return Mono.error(new AgentToolApprovalRejectedException("审批请求已失效。")); } + if (isExpired(pendingApproval.state)) { + closeBatch(pendingApproval.batchId, "审批请求已过期。"); + return Mono.error(new AgentToolApprovalRejectedException("审批请求已过期。")); + } return Mono.fromFuture(pendingApproval.future); } + /** + * 处理恢复请求并生成批次级审批决议。 + * + * @param request 恢复请求 + * @return 审批决议 + */ + public synchronized AgentToolApprovalResolution resolve(AgentResumeRequest request) { + validateResumeRequest(request); + if (!enabled) { + AgentPendingState state = new AgentPendingState(); + state.setResumeToken(request.getResumeToken()); + return new AgentToolApprovalResolution( + AgentToolApprovalResolution.Status.READY, state, List.of(), null); + } + String token = request.getResumeToken().getValue(); + PendingApproval pendingApproval = approvals.get(token); + if (pendingApproval == null || pendingApproval.decision != ApprovalDecision.PENDING) { + throw new AgentRuntimeException("Agent resume token is invalid, expired, or already consumed."); + } + if (isExpired(pendingApproval.state)) { + closeBatch(pendingApproval.batchId, "审批请求已过期。"); + return new AgentToolApprovalResolution( + AgentToolApprovalResolution.Status.EXPIRED, + pendingApproval.state, + List.of(), + "审批请求已过期。"); + } + + pendingApproval.decision = request.isApproved() + ? ApprovalDecision.APPROVED + : ApprovalDecision.REJECTED; + approvals.remove(token); + removeToolCallIndex(pendingApproval.state, token); + pendingApproval.future.complete(request); + + if (!request.isApproved()) { + String reason = request.getRejectReason() == null || request.getRejectReason().isBlank() + ? "工具执行已被拒绝。" + : request.getRejectReason(); + closeBatch(pendingApproval.batchId, reason); + return new AgentToolApprovalResolution( + AgentToolApprovalResolution.Status.REJECTED, + pendingApproval.state, + List.of(), + reason); + } + + ApprovalBatch batch = approvalBatches.get(pendingApproval.batchId); + if (batch == null) { + authorize(pendingApproval); + return new AgentToolApprovalResolution( + AgentToolApprovalResolution.Status.READY, + pendingApproval.state, + List.of(), + null); + } + List remainingStates = pendingStatesInBatch(batch); + AgentPendingState expiredState = remainingStates.stream() + .filter(this::isExpired) + .findFirst() + .orElse(null); + if (expiredState != null) { + closeBatch(batch.batchId, "审批请求已过期。"); + return new AgentToolApprovalResolution( + AgentToolApprovalResolution.Status.EXPIRED, + pendingApproval.state, + List.of(), + "审批请求已过期。"); + } + if (!remainingStates.isEmpty()) { + return new AgentToolApprovalResolution( + AgentToolApprovalResolution.Status.WAITING, + pendingApproval.state, + remainingStates, + null); + } + + for (String batchToken : batch.tokens) { + PendingApproval member = batch.members.get(batchToken); + if (member != null && member.decision == ApprovalDecision.APPROVED) { + authorize(member); + } + } + approvalBatches.remove(batch.batchId); + return new AgentToolApprovalResolution( + AgentToolApprovalResolution.Status.READY, + pendingApproval.state, + List.of(), + null); + } + /** * 消费恢复请求对应的待审批状态。 * - *

该方法用于有状态 runtime 的 HITL resume。第一版 pending state 仅保存在 - * 当前进程内存中,因此消费成功后会立即移除 token,避免重复恢复。

+ *

该兼容入口供工具适配器内部的 await/consume 流程使用。批次中仍有未决调用时 + * 会拒绝提前消费;单调用批准后会移除无需经过 PreActing 的执行凭证。

* * @param request 恢复请求 * @return 待审批状态 */ - public AgentPendingState consume(AgentResumeRequest request) { - if (request == null || request.getResumeToken() == null - || request.getResumeToken().getValue() == null - || request.getResumeToken().getValue().isBlank()) { - throw new AgentRuntimeException("Agent resume token is required."); + public synchronized AgentPendingState consume(AgentResumeRequest request) { + AgentToolApprovalResolution resolution = resolve(request); + if (resolution.getStatus() == AgentToolApprovalResolution.Status.WAITING) { + throw new AgentRuntimeException("Approval batch still has pending tool calls."); } + AgentPendingState state = resolution.getResolvedState(); + if (resolution.getStatus() == AgentToolApprovalResolution.Status.READY + && state != null + && state.getToolCallId() != null) { + // 兼容工具适配器内部 await/consume 流程,该流程会直接调用工具,不经过 PreActing 二次校验。 + executionAuthorizations.remove(state.getToolCallId()); + } + return state; + } + + /** + * 根据服务端持久化审批结果签发受信任的一次性执行授权。 + * + *

恢复元数据应提供单个 {@code toolCallId/toolName/toolInput},或通过 + * {@code approvedToolCalls} 提供上述字段组成的列表。该入口仅供已经完成持久化令牌 + * 校验和一次性消费的服务端集成层使用。

+ * + * @param request 受信任恢复请求 + */ + public synchronized void authorizeTrustedExecution(AgentResumeRequest request) { + validateResumeRequest(request); + if (!enabled || !request.isApproved()) { + return; + } + Map metadata = request.getMetadata() == null + ? Map.of() + : request.getMetadata(); + Object approvedToolCalls = metadata.get("approvedToolCalls"); + Map trustedAuthorizations = new LinkedHashMap<>(); + int authorizationCount = 0; + if (approvedToolCalls instanceof List calls) { + for (Object call : calls) { + if (call instanceof Map callMap) { + authorizeTrustedCall(callMap, trustedAuthorizations); + authorizationCount++; + } + } + } else if (metadata.containsKey("toolCallId")) { + authorizeTrustedCall(metadata, trustedAuthorizations); + authorizationCount++; + } + if (authorizationCount == 0) { + throw new AgentRuntimeException( + "Trusted resume metadata must include approved toolCallId, toolName, and toolInput."); + } + executionAuthorizations.putAll(trustedAuthorizations); + } + + /** + * 消费指定工具调用的一次性执行授权。 + * + * @param toolCallId 工具调用ID + * @param toolName 工具名称 + * @param toolInput 工具入参 + */ + public synchronized void consumeExecutionAuthorization(String toolCallId, + String toolName, + Map toolInput) { if (!enabled) { - AgentPendingState state = new AgentPendingState(); - state.setResumeToken(request.getResumeToken()); - return state; + return; } - PendingApproval pendingApproval = approvals.remove(request.getResumeToken().getValue()); - if (pendingApproval == null) { - throw new AgentRuntimeException("Agent resume token is invalid or expired."); + if (toolCallId == null || toolCallId.isBlank()) { + throw new AgentToolApprovalRejectedException("待执行工具缺少 toolCallId,无法校验审批结果。"); } - pendingApproval.future.complete(request); - return pendingApproval.state; + ExecutionAuthorization authorization = executionAuthorizations.remove(toolCallId); + if (authorization == null) { + throw new AgentToolApprovalRejectedException("工具调用未获得批准或批准已被消费。"); + } + if (!Objects.equals(authorization.toolName, toolName) + || !Objects.equals(authorization.toolInput, toolInput)) { + throw new AgentToolApprovalRejectedException("工具调用与已批准内容不一致。"); + } + } + + /** + * 清理尚未消费的一次性执行授权。 + */ + public synchronized void clearExecutionAuthorizations() { + executionAuthorizations.clear(); + } + + /** + * 获取指定会话当前仍待处理的审批状态。 + * + * @param sessionId 会话ID + * @return 待审批状态快照 + */ + public synchronized List pendingStates(String sessionId) { + return approvals.values().stream() + .filter(pending -> pending.decision == ApprovalDecision.PENDING) + .map(pending -> pending.state) + .filter(state -> sessionId == null || Objects.equals(sessionId, state.getSessionId())) + .toList(); } /** @@ -140,18 +389,16 @@ public class AgentToolApprovalCoordinator { * * @param reason 取消原因 */ - public void cancelAll(String reason) { + public synchronized void cancelAll(String reason) { if (!enabled) { return; } - for (PendingApproval pendingApproval : approvals.values()) { - AgentResumeRequest request = new AgentResumeRequest(); - request.setResumeToken(pendingApproval.state.getResumeToken()); - request.setApproved(false); - request.setRejectReason(reason); - pendingApproval.future.complete(request); + for (String batchId : new ArrayList<>(approvalBatches.keySet())) { + closeBatch(batchId, reason); } approvals.clear(); + tokensByToolCallId.clear(); + executionAuthorizations.clear(); } /** @@ -163,13 +410,232 @@ public class AgentToolApprovalCoordinator { return enabled; } + /** + * 校验恢复请求中的令牌字段。 + * + * @param request 恢复请求 + */ + private void validateResumeRequest(AgentResumeRequest request) { + if (request == null || request.getResumeToken() == null + || request.getResumeToken().getValue() == null + || request.getResumeToken().getValue().isBlank()) { + throw new AgentRuntimeException("Agent resume token is required."); + } + } + + /** + * 判断待审批状态是否已经过期。 + * + * @param state 待审批状态 + * @return 已过期时为 true + */ + private boolean isExpired(AgentPendingState state) { + return state != null + && state.getExpiresAt() != null + && !state.getExpiresAt().isAfter(Instant.now()); + } + + /** + * 获取批次内尚未决策的审批状态。 + * + * @param batch 审批批次 + * @return 未决审批状态 + */ + private List pendingStatesInBatch(ApprovalBatch batch) { + List states = new ArrayList<>(); + for (String token : batch.tokens) { + PendingApproval member = batch.members.get(token); + if (member != null && member.decision == ApprovalDecision.PENDING) { + states.add(member.state); + } + } + return states; + } + + /** + * 为已批准状态签发一次性执行授权。 + * + * @param pendingApproval 已批准状态 + */ + private void authorize(PendingApproval pendingApproval) { + AgentPendingState state = pendingApproval.state; + if (state.getToolCallId() == null || state.getToolCallId().isBlank()) { + throw new AgentRuntimeException("Approved tool call is missing toolCallId."); + } + executionAuthorizations.put(state.getToolCallId(), new ExecutionAuthorization( + state.getToolName(), + state.getToolInput())); + } + + /** + * 为服务端持久化审批结果签发一次性执行授权。 + * + * @param callMap 已批准调用元数据 + * @param trustedAuthorizations 本次恢复待签发的临时授权集合 + */ + private void authorizeTrustedCall(Map callMap, + Map trustedAuthorizations) { + String toolCallId = stringValue(callMap.get("toolCallId")); + String toolName = stringValue(callMap.get("toolName")); + if (toolCallId == null || toolName == null) { + throw new AgentRuntimeException( + "Trusted resume metadata must include non-empty toolCallId and toolName."); + } + Map toolInput = stringKeyMap(callMap.get("toolInput")); + ExecutionAuthorization authorization = new ExecutionAuthorization(toolName, toolInput); + ExecutionAuthorization previous = trustedAuthorizations.put(toolCallId, authorization); + if (previous != null + && (!Objects.equals(previous.toolName, toolName) + || !Objects.equals(previous.toolInput, toolInput))) { + throw new AgentRuntimeException( + "Trusted resume contains conflicting tool calls for toolCallId: " + toolCallId); + } + } + + /** + * 将值转换为非空字符串。 + * + * @param value 原始值 + * @return 非空字符串;无法转换时返回 null + */ + private String stringValue(Object value) { + if (value == null || String.valueOf(value).isBlank()) { + return null; + } + return String.valueOf(value); + } + + /** + * 将任意键 Map 转换为字符串键 Map。 + * + * @param value 原始值 + * @return 字符串键 Map + */ + private Map stringKeyMap(Object value) { + if (value == null) { + return Map.of(); + } + if (!(value instanceof Map source)) { + throw new AgentRuntimeException("Trusted resume toolInput must be a map."); + } + Map result = new LinkedHashMap<>(); + for (Map.Entry entry : source.entrySet()) { + if (entry.getKey() != null) { + result.put(String.valueOf(entry.getKey()), entry.getValue()); + } + } + return result; + } + + /** + * 移除工具调用ID到审批令牌的索引。 + * + * @param state 待审批状态 + * @param token 审批令牌 + */ + private void removeToolCallIndex(AgentPendingState state, String token) { + if (state.getToolCallId() != null && !state.getToolCallId().isBlank()) { + tokensByToolCallId.remove(state.getToolCallId(), token); + } + } + + /** + * 关闭审批批次并拒绝尚未决策的审批项。 + * + * @param batchId 审批批次ID + * @param reason 关闭原因 + */ + private void closeBatch(String batchId, String reason) { + ApprovalBatch batch = approvalBatches.remove(batchId); + if (batch == null) { + return; + } + for (String token : batch.tokens) { + PendingApproval member = batch.members.get(token); + if (member == null) { + continue; + } + approvals.remove(token); + removeToolCallIndex(member.state, token); + if (member.decision == ApprovalDecision.PENDING) { + member.decision = ApprovalDecision.REJECTED; + AgentResumeRequest rejection = new AgentResumeRequest(); + rejection.setResumeToken(member.state.getResumeToken()); + rejection.setApproved(false); + rejection.setRejectReason(reason); + member.future.complete(rejection); + } + } + } + + /** + * 单个审批项的内部决策状态。 + */ + private enum ApprovalDecision { + PENDING, + APPROVED, + REJECTED + } + + /** + * 待审批项及其异步等待句柄。 + */ private static class PendingApproval { private final AgentPendingState state; + private final String batchId; private final CompletableFuture future; + private ApprovalDecision decision = ApprovalDecision.PENDING; - private PendingApproval(AgentPendingState state, CompletableFuture future) { + /** + * 创建待审批项。 + * + * @param state 待审批状态 + * @param batchId 审批批次ID + * @param future 审批响应等待句柄 + */ + private PendingApproval(AgentPendingState state, + String batchId, + CompletableFuture future) { this.state = Objects.requireNonNull(state, "state"); + this.batchId = Objects.requireNonNull(batchId, "batchId"); this.future = Objects.requireNonNull(future, "future"); } } + + /** + * 同一轮推理产生的审批批次。 + */ + private static class ApprovalBatch { + private final String batchId; + private final Set tokens = new LinkedHashSet<>(); + private final Map members = new LinkedHashMap<>(); + + /** + * 创建审批批次。 + * + * @param batchId 审批批次ID + */ + private ApprovalBatch(String batchId) { + this.batchId = Objects.requireNonNull(batchId, "batchId"); + } + } + + /** + * 已批准工具调用的一次性执行凭证。 + */ + private static class ExecutionAuthorization { + private final String toolName; + private final Map toolInput; + + /** + * 创建一次性执行授权。 + * + * @param toolName 工具名称 + * @param toolInput 工具入参 + */ + private ExecutionAuthorization(String toolName, Map toolInput) { + this.toolName = toolName; + this.toolInput = toolInput == null ? Map.of() : new LinkedHashMap<>(toolInput); + } + } } diff --git a/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/hitl/AgentToolApprovalResolution.java b/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/hitl/AgentToolApprovalResolution.java new file mode 100644 index 0000000..c67e72c --- /dev/null +++ b/easy-agents-agent-runtime/src/main/java/com/easyagents/agent/runtime/hitl/AgentToolApprovalResolution.java @@ -0,0 +1,93 @@ +package com.easyagents.agent.runtime.hitl; + +import java.util.List; + +/** + * 工具审批决议。 + */ +public final class AgentToolApprovalResolution { + + /** + * 审批决议状态。 + */ + public enum Status { + /** + * 当前审批批次仍有待处理调用。 + */ + WAITING, + + /** + * 当前审批批次已全部批准,可以恢复执行。 + */ + READY, + + /** + * 当前审批批次已被拒绝。 + */ + REJECTED, + + /** + * 当前审批批次已过期。 + */ + EXPIRED + } + + private final Status status; + private final AgentPendingState resolvedState; + private final List remainingStates; + private final String reason; + + /** + * 创建工具审批决议。 + * + * @param status 决议状态 + * @param resolvedState 本次处理的待审批状态 + * @param remainingStates 同一批次剩余的待审批状态 + * @param reason 拒绝或过期原因 + */ + public AgentToolApprovalResolution(Status status, + AgentPendingState resolvedState, + List remainingStates, + String reason) { + this.status = status; + this.resolvedState = resolvedState; + this.remainingStates = remainingStates == null ? List.of() : List.copyOf(remainingStates); + this.reason = reason; + } + + /** + * 获取决议状态。 + * + * @return 决议状态 + */ + public Status getStatus() { + return status; + } + + /** + * 获取本次处理的待审批状态。 + * + * @return 待审批状态 + */ + public AgentPendingState getResolvedState() { + return resolvedState; + } + + /** + * 获取同一批次剩余的待审批状态。 + * + * @return 剩余待审批状态 + */ + public List getRemainingStates() { + return remainingStates; + } + + /** + * 获取拒绝或过期原因。 + * + * @return 原因 + */ + public String getReason() { + return reason; + } +} diff --git a/easy-agents-agent-runtime/src/test/java/com/easyagents/agent/runtime/agentscope/AgentScopeStatefulRuntimeTest.java b/easy-agents-agent-runtime/src/test/java/com/easyagents/agent/runtime/agentscope/AgentScopeStatefulRuntimeTest.java index 50028ae..c4543e4 100644 --- a/easy-agents-agent-runtime/src/test/java/com/easyagents/agent/runtime/agentscope/AgentScopeStatefulRuntimeTest.java +++ b/easy-agents-agent-runtime/src/test/java/com/easyagents/agent/runtime/agentscope/AgentScopeStatefulRuntimeTest.java @@ -56,6 +56,7 @@ import java.util.Set; import java.util.concurrent.CompletableFuture; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; import java.util.function.BooleanSupplier; /** @@ -843,6 +844,9 @@ public class AgentScopeStatefulRuntimeTest { Assert.assertFalse(((List) suspended.getPayload().get("pendingApprovals")).isEmpty()); Assert.assertFalse(events.stream().anyMatch(event -> event.getEventType() == AgentRuntimeEventType.COMPLETED)); Assert.assertFalse(events.stream().anyMatch(event -> event.getEventType() == AgentRuntimeEventType.TOOL_RESULT)); + Assert.assertEquals(1, events.stream() + .filter(event -> event.getEventType() == AgentRuntimeEventType.SUSPENDED) + .count()); } @Test(expected = AgentRuntimeException.class) @@ -927,6 +931,267 @@ public class AgentScopeStatefulRuntimeTest { Assert.assertTrue(sessionStore.exists("session-1")); } + /** + * 验证同一轮推理包含多个审批工具时,全部批准前不会执行任何工具。 + */ + @Test + public void shouldWaitForAllToolApprovalsBeforeExecutingBatch() { + AgentInitRequest request = initRequest(); + AgentToolSpec searchSpec = new AgentToolSpec(); + searchSpec.setName("search"); + searchSpec.setDescription("search"); + searchSpec.setApprovalRequired(true); + AgentToolSpec auditSpec = new AgentToolSpec(); + auditSpec.setName("audit"); + auditSpec.setDescription("audit"); + auditSpec.setApprovalRequired(true); + request.getAgentDefinition().setToolSpecs(List.of(searchSpec, auditSpec)); + AtomicInteger invocationCount = new AtomicInteger(); + request.setToolInvokers(Map.of( + "search", (arguments, context) -> { + invocationCount.incrementAndGet(); + return AgentToolResult.success("search result"); + }, + "audit", (arguments, context) -> { + invocationCount.incrementAndGet(); + return AgentToolResult.success("audit result"); + })); + AgentScopeReActRuntime runtime = runtimeWithModel(List.of( + ChatResponse.builder() + .id("tool-call-message") + .content(List.of( + ToolUseBlock.builder() + .id("call-search") + .name("search") + .input(Map.of("q", "easyflow")) + .build(), + ToolUseBlock.builder() + .id("call-audit") + .name("audit") + .input(Map.of("scope", "current")) + .build())) + .finishReason("tool_calls") + .build(), + ChatResponse.builder() + .id("final-message") + .content(List.of(TextBlock.builder().text("done").build())) + .finishReason("stop") + .build())); + runtime.init(request); + + List initialEvents = runtime.stream( + AgentMessage.text(AgentMessageRole.USER, "use tools")) + .collectList() + .block(); + List approvals = initialEvents.stream() + .filter(event -> event.getEventType() == AgentRuntimeEventType.TOOL_APPROVAL_REQUIRED) + .toList(); + + Assert.assertEquals(2, approvals.size()); + Assert.assertEquals(1, initialEvents.stream() + .filter(event -> event.getEventType() == AgentRuntimeEventType.SUSPENDED) + .count()); + List firstResumeEvents = runtime.resume(resumeFromApproval(approvals.get(0), true)) + .collectList() + .block(); + + Assert.assertEquals(0, invocationCount.get()); + AgentRuntimeEvent waitingEvent = firstResumeEvents.stream() + .filter(event -> event.getEventType() == AgentRuntimeEventType.SUSPENDED) + .findFirst() + .orElseThrow(); + Assert.assertEquals(1, ((List) waitingEvent.getPayload().get("pendingApprovals")).size()); + Assert.assertFalse(firstResumeEvents.stream() + .anyMatch(event -> event.getEventType() == AgentRuntimeEventType.TOOL_CALL)); + + List secondResumeEvents = runtime.resume(resumeFromApproval(approvals.get(1), true)) + .collectList() + .block(); + + Assert.assertEquals(2, invocationCount.get()); + Assert.assertEquals(2, secondResumeEvents.stream() + .filter(event -> event.getEventType() == AgentRuntimeEventType.TOOL_RESULT) + .count()); + Assert.assertTrue(secondResumeEvents.stream() + .anyMatch(event -> event.getEventType() == AgentRuntimeEventType.COMPLETED)); + } + + /** + * 验证不同 toolCallId 即使名称和入参相同也会分别审批和执行。 + */ + @Test + public void shouldPreserveDistinctApprovedToolCallsWithIdenticalInput() { + AgentInitRequest request = initRequest(); + AgentToolSpec toolSpec = new AgentToolSpec(); + toolSpec.setName("search"); + toolSpec.setDescription("search"); + toolSpec.setApprovalRequired(true); + request.getAgentDefinition().setToolSpecs(List.of(toolSpec)); + AtomicInteger invocationCount = new AtomicInteger(); + request.setToolInvokers(Map.of("search", (arguments, context) -> { + invocationCount.incrementAndGet(); + return AgentToolResult.success("tool result"); + })); + Map input = Map.of("q", "easyflow"); + AgentScopeReActRuntime runtime = runtimeWithModel(List.of( + ChatResponse.builder() + .id("tool-call-message") + .content(List.of( + ToolUseBlock.builder() + .id("call-search") + .name("search") + .input(input) + .build(), + ToolUseBlock.builder() + .id("call-search-duplicate") + .name("search") + .input(input) + .build())) + .finishReason("tool_calls") + .build(), + ChatResponse.builder() + .id("final-message") + .content(List.of(TextBlock.builder().text("done").build())) + .finishReason("stop") + .build())); + runtime.init(request); + + List suspendedEvents = runtime.stream( + AgentMessage.text(AgentMessageRole.USER, "use tool")) + .collectList() + .block(); + List approvals = suspendedEvents.stream() + .filter(event -> event.getEventType() == AgentRuntimeEventType.TOOL_APPROVAL_REQUIRED) + .toList(); + + Assert.assertEquals(2, approvals.size()); + List firstResumeEvents = runtime.resume(resumeFromApproval(approvals.get(0), true)) + .collectList() + .block(); + + Assert.assertEquals(0, invocationCount.get()); + Assert.assertTrue(firstResumeEvents.stream() + .anyMatch(event -> event.getEventType() == AgentRuntimeEventType.SUSPENDED)); + List secondResumeEvents = runtime.resume(resumeFromApproval(approvals.get(1), true)) + .collectList() + .block(); + + Assert.assertEquals(2, invocationCount.get()); + Assert.assertEquals(2, secondResumeEvents.stream() + .filter(event -> event.getEventType() == AgentRuntimeEventType.TOOL_RESULT) + .count()); + Assert.assertTrue(secondResumeEvents.stream() + .anyMatch(event -> event.getEventType() == AgentRuntimeEventType.COMPLETED)); + } + + /** + * 验证相同 toolCallId 的同轮重放只审批并执行一次。 + */ + @Test + public void shouldDeduplicateRepeatedToolCallIdWithinReasoning() { + AgentInitRequest request = initRequest(); + AgentToolSpec toolSpec = new AgentToolSpec(); + toolSpec.setName("search"); + toolSpec.setDescription("search"); + toolSpec.setApprovalRequired(true); + request.getAgentDefinition().setToolSpecs(List.of(toolSpec)); + AtomicInteger invocationCount = new AtomicInteger(); + request.setToolInvokers(Map.of("search", (arguments, context) -> { + invocationCount.incrementAndGet(); + return AgentToolResult.success("tool result"); + })); + Map input = Map.of("q", "easyflow"); + AgentScopeReActRuntime runtime = runtimeWithModel(List.of( + ChatResponse.builder() + .id("tool-call-message") + .content(List.of( + ToolUseBlock.builder() + .id("call-search") + .name("search") + .input(input) + .build(), + ToolUseBlock.builder() + .id("call-search") + .name("search") + .input(input) + .build())) + .finishReason("tool_calls") + .build(), + ChatResponse.builder() + .id("final-message") + .content(List.of(TextBlock.builder().text("done").build())) + .finishReason("stop") + .build())); + runtime.init(request); + + List suspendedEvents = runtime.stream( + AgentMessage.text(AgentMessageRole.USER, "use tool")) + .collectList() + .block(); + List approvals = suspendedEvents.stream() + .filter(event -> event.getEventType() == AgentRuntimeEventType.TOOL_APPROVAL_REQUIRED) + .toList(); + + Assert.assertEquals(1, approvals.size()); + List resumeEvents = runtime.resume(resumeFromApproval(approvals.get(0), true)) + .collectList() + .block(); + + Assert.assertEquals(1, invocationCount.get()); + Assert.assertEquals(1, resumeEvents.stream() + .filter(event -> event.getEventType() == AgentRuntimeEventType.TOOL_RESULT) + .count()); + Assert.assertTrue(resumeEvents.stream() + .anyMatch(event -> event.getEventType() == AgentRuntimeEventType.COMPLETED)); + } + + /** + * 验证模型未提供 toolCallId 时审批链路仍会获得完整的调用身份。 + */ + @Test + public void shouldProvideApprovalIdentityWhenModelOmitsToolCallId() { + AgentInitRequest request = initRequest(); + AgentToolSpec toolSpec = new AgentToolSpec(); + toolSpec.setName("search"); + toolSpec.setDescription("search"); + toolSpec.setApprovalRequired(true); + request.getAgentDefinition().setToolSpecs(List.of(toolSpec)); + request.setToolInvokers(Map.of("search", (arguments, context) -> + AgentToolResult.success("tool result"))); + AgentScopeReActRuntime runtime = runtimeWithModel(List.of( + ChatResponse.builder() + .id("tool-call-message") + .content(List.of(ToolUseBlock.builder() + .name("search") + .input(Map.of("q", "easyflow")) + .build())) + .finishReason("tool_calls") + .build(), + ChatResponse.builder() + .id("final-message") + .content(List.of(TextBlock.builder().text("done").build())) + .finishReason("stop") + .build())); + runtime.init(request); + + List events = runtime.stream( + AgentMessage.text(AgentMessageRole.USER, "use tool")) + .collectList() + .block(); + AgentRuntimeEvent approval = events.stream() + .filter(event -> event.getEventType() == AgentRuntimeEventType.TOOL_APPROVAL_REQUIRED) + .findFirst() + .orElseThrow(); + + String toolCallId = String.valueOf(approval.getPayload().get("toolCallId")); + Assert.assertFalse(toolCallId.isBlank()); + @SuppressWarnings("unchecked") + Map metadata = + (Map) approval.getPayload().get("approvalMetadata"); + String approvalBatchId = String.valueOf(metadata.get("approvalBatchId")); + Assert.assertFalse(approvalBatchId.isBlank()); + } + @Test public void shouldCancelRejectedToolResumeWithoutExecutingTool() { InMemoryAgentSessionStore sessionStore = new InMemoryAgentSessionStore(); @@ -1105,6 +1370,22 @@ public class AgentScopeStatefulRuntimeTest { new AgentScopeMessageAdapter()); } + /** + * 根据审批事件创建恢复请求。 + * + * @param approvalEvent 工具审批事件 + * @param approved 是否批准 + * @return 恢复请求 + */ + private AgentResumeRequest resumeFromApproval(AgentRuntimeEvent approvalEvent, boolean approved) { + AgentResumeRequest request = new AgentResumeRequest(); + AgentResumeToken token = new AgentResumeToken(); + token.setValue(String.valueOf(approvalEvent.getPayload().get("resumeToken"))); + request.setResumeToken(token); + request.setApproved(approved); + return request; + } + private static class ScriptedModel implements Model { private final String modelName; diff --git a/easy-agents-agent-runtime/src/test/java/com/easyagents/agent/runtime/hitl/AgentToolApprovalCoordinatorTest.java b/easy-agents-agent-runtime/src/test/java/com/easyagents/agent/runtime/hitl/AgentToolApprovalCoordinatorTest.java new file mode 100644 index 0000000..c3aad9d --- /dev/null +++ b/easy-agents-agent-runtime/src/test/java/com/easyagents/agent/runtime/hitl/AgentToolApprovalCoordinatorTest.java @@ -0,0 +1,258 @@ +package com.easyagents.agent.runtime.hitl; + +import com.easyagents.agent.runtime.AgentResumeRequest; +import com.easyagents.agent.runtime.AgentRuntimeException; +import org.junit.Assert; +import org.junit.Test; + +import java.time.Instant; +import java.util.Map; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.Future; + +/** + * 测试工具审批协调器。 + */ +public class AgentToolApprovalCoordinatorTest { + + /** + * 验证同批次全部调用批准后才签发逐调用执行授权。 + */ + @Test + public void shouldAuthorizeBatchOnlyAfterAllCallsApproved() { + AgentToolApprovalCoordinator coordinator = AgentToolApprovalCoordinator.enabled(); + AgentPendingState first = register(coordinator, "call-1", "search", Map.of("q", "first"), "batch-1"); + AgentPendingState second = register(coordinator, "call-2", "search", Map.of("q", "second"), "batch-1"); + + AgentToolApprovalResolution firstResolution = coordinator.resolve(resume(first, true)); + + Assert.assertEquals(AgentToolApprovalResolution.Status.WAITING, firstResolution.getStatus()); + Assert.assertEquals(1, firstResolution.getRemainingStates().size()); + assertAuthorizationRejected(coordinator, "call-1", "search", Map.of("q", "first")); + + AgentToolApprovalResolution secondResolution = coordinator.resolve(resume(second, true)); + + Assert.assertEquals(AgentToolApprovalResolution.Status.READY, secondResolution.getStatus()); + coordinator.consumeExecutionAuthorization("call-1", "search", Map.of("q", "first")); + coordinator.consumeExecutionAuthorization("call-2", "search", Map.of("q", "second")); + assertAuthorizationRejected(coordinator, "call-1", "search", Map.of("q", "first")); + } + + /** + * 验证拒绝一个调用会关闭整个审批批次。 + */ + @Test + public void shouldRejectWholeBatchWhenAnyCallRejected() { + AgentToolApprovalCoordinator coordinator = AgentToolApprovalCoordinator.enabled(); + AgentPendingState first = register(coordinator, "call-1", "search", Map.of("q", "first"), "batch-1"); + AgentPendingState second = register(coordinator, "call-2", "search", Map.of("q", "second"), "batch-1"); + coordinator.resolve(resume(first, true)); + + AgentResumeRequest rejection = resume(second, false); + rejection.setRejectReason("not allowed"); + AgentToolApprovalResolution resolution = coordinator.resolve(rejection); + + Assert.assertEquals(AgentToolApprovalResolution.Status.REJECTED, resolution.getStatus()); + Assert.assertEquals("not allowed", resolution.getReason()); + assertAuthorizationRejected(coordinator, "call-1", "search", Map.of("q", "first")); + try { + coordinator.resolve(resume(first, true)); + Assert.fail("已消费的审批令牌不能重复使用"); + } catch (AgentRuntimeException expected) { + Assert.assertTrue(expected.getMessage().contains("invalid")); + } + } + + /** + * 验证过期令牌不能签发工具执行授权。 + */ + @Test + public void shouldExpireApprovalBeforeResolution() { + AgentToolApprovalCoordinator coordinator = AgentToolApprovalCoordinator.enabled(); + AgentPendingState expired = coordinator.register( + "session-1", + "agent-1", + "call-expired", + "search", + "approve", + Map.of("q", "expired"), + Map.of(), + Instant.now().minusSeconds(1), + "batch-expired"); + + AgentToolApprovalResolution resolution = coordinator.resolve(resume(expired, true)); + + Assert.assertEquals(AgentToolApprovalResolution.Status.EXPIRED, resolution.getStatus()); + assertAuthorizationRejected(coordinator, "call-expired", "search", Map.of("q", "expired")); + } + + /** + * 验证并发重复点击同一令牌时最多一个请求可以成功消费。 + * + * @throws Exception 并发任务执行失败时抛出 + */ + @Test + public void shouldConsumeConcurrentDuplicateApprovalOnlyOnce() throws Exception { + AgentToolApprovalCoordinator coordinator = AgentToolApprovalCoordinator.enabled(); + AgentPendingState pending = register( + coordinator, "call-1", "search", Map.of("q", "easyflow"), "batch-1"); + CountDownLatch start = new CountDownLatch(1); + ExecutorService executor = Executors.newFixedThreadPool(2); + try { + Future first = executor.submit(() -> resolveAfter(start, coordinator, pending)); + Future second = executor.submit(() -> resolveAfter(start, coordinator, pending)); + start.countDown(); + + int successCount = (first.get() ? 1 : 0) + (second.get() ? 1 : 0); + + Assert.assertEquals(1, successCount); + coordinator.consumeExecutionAuthorization( + "call-1", "search", Map.of("q", "easyflow")); + assertAuthorizationRejected( + coordinator, "call-1", "search", Map.of("q", "easyflow")); + } finally { + executor.shutdownNow(); + } + } + + /** + * 验证工具名称或入参变化时批准凭证立即失效。 + */ + @Test + public void shouldRejectExecutionWhenApprovedCallIsModified() { + AgentToolApprovalCoordinator coordinator = AgentToolApprovalCoordinator.enabled(); + AgentPendingState pending = register( + coordinator, "call-1", "search", Map.of("q", "easyflow"), "batch-1"); + coordinator.resolve(resume(pending, true)); + + assertAuthorizationRejected( + coordinator, "call-1", "search", Map.of("q", "modified")); + assertAuthorizationRejected( + coordinator, "call-1", "search", Map.of("q", "easyflow")); + } + + /** + * 验证跨节点受信任恢复仍需绑定明确的工具调用信息。 + */ + @Test + public void shouldAuthorizeTrustedExecutionByToolCallIdentity() { + AgentToolApprovalCoordinator coordinator = AgentToolApprovalCoordinator.enabled(); + AgentResumeRequest request = new AgentResumeRequest(); + AgentResumeToken token = new AgentResumeToken(); + token.setValue("persisted-token"); + request.setResumeToken(token); + request.setApproved(true); + request.setTrusted(true); + request.setMetadata(Map.of( + "toolCallId", "call-1", + "toolName", "search", + "toolInput", Map.of("q", "easyflow"))); + + coordinator.authorizeTrustedExecution(request); + + coordinator.consumeExecutionAuthorization( + "call-1", "search", Map.of("q", "easyflow")); + assertAuthorizationRejected( + coordinator, "call-1", "search", Map.of("q", "easyflow")); + } + + /** + * 验证同一 toolCallId 不能被重新绑定到不同工具内容。 + */ + @Test + public void shouldRejectDuplicateToolCallIdWithDifferentInput() { + AgentToolApprovalCoordinator coordinator = AgentToolApprovalCoordinator.enabled(); + register(coordinator, "call-1", "search", Map.of("q", "easyflow"), "batch-1"); + + try { + register(coordinator, "call-1", "search", Map.of("q", "modified"), "batch-2"); + Assert.fail("重复 toolCallId 不能绑定不同入参"); + } catch (AgentRuntimeException expected) { + Assert.assertTrue(expected.getMessage().contains("Duplicate toolCallId")); + } + } + + /** + * 注册测试用审批状态。 + * + * @param coordinator 审批协调器 + * @param toolCallId 工具调用ID + * @param toolName 工具名称 + * @param toolInput 工具入参 + * @param batchId 审批批次ID + * @return 待审批状态 + */ + private AgentPendingState register(AgentToolApprovalCoordinator coordinator, + String toolCallId, + String toolName, + Map toolInput, + String batchId) { + return coordinator.register( + "session-1", + "agent-1", + toolCallId, + toolName, + "approve", + toolInput, + Map.of(), + Instant.now().plusSeconds(60), + batchId); + } + + /** + * 创建测试用恢复请求。 + * + * @param state 待审批状态 + * @param approved 是否批准 + * @return 恢复请求 + */ + private AgentResumeRequest resume(AgentPendingState state, boolean approved) { + AgentResumeRequest request = new AgentResumeRequest(); + request.setResumeToken(state.getResumeToken()); + request.setApproved(approved); + return request; + } + + /** + * 等待并发起跑信号后消费审批令牌。 + * + * @param start 起跑信号 + * @param coordinator 审批协调器 + * @param pending 待审批状态 + * @return 成功消费时为 true + * @throws InterruptedException 等待被中断时抛出 + */ + private boolean resolveAfter(CountDownLatch start, + AgentToolApprovalCoordinator coordinator, + AgentPendingState pending) throws InterruptedException { + start.await(); + try { + coordinator.resolve(resume(pending, true)); + return true; + } catch (AgentRuntimeException expected) { + return false; + } + } + + /** + * 断言工具执行授权不可用。 + * + * @param coordinator 审批协调器 + * @param toolCallId 工具调用ID + * @param toolName 工具名称 + * @param toolInput 工具入参 + */ + private void assertAuthorizationRejected(AgentToolApprovalCoordinator coordinator, + String toolCallId, + String toolName, + Map toolInput) { + try { + coordinator.consumeExecutionAuthorization(toolCallId, toolName, toolInput); + Assert.fail("未授权或已消费的工具调用必须被拒绝"); + } catch (AgentToolApprovalRejectedException expected) { + Assert.assertNotNull(expected.getMessage()); + } + } +}