From 3e5084a1fe12df12637d8eb2c6573a9e3ced043b Mon Sep 17 00:00:00 2001 From: Jie Shen <185767017+Sj295@users.noreply.github.com> Date: Fri, 17 Jul 2026 18:48:50 +0800 Subject: [PATCH] feat: permission-aware tool calling + progress display --- .../java/com/ccj/app/AppConfiguration.java | 20 ++++- .../ccj/app/ClaudeCodeJavaApplication.java | 17 ++++ .../com/ccj/app/invoker/AgentInvoker.java | 41 +++++++--- .../ccj/test/PermissionAwareAdapterTest.java | 78 +++++++++++++++++++ .../tools/adapter/ToolCallbackAdapter.java | 33 ++++++++ 5 files changed, 178 insertions(+), 11 deletions(-) create mode 100644 ccc-test/src/test/java/com/ccj/test/PermissionAwareAdapterTest.java diff --git a/ccc-app/src/main/java/com/ccj/app/AppConfiguration.java b/ccc-app/src/main/java/com/ccj/app/AppConfiguration.java index 9a2e7c4..3baec24 100644 --- a/ccc-app/src/main/java/com/ccj/app/AppConfiguration.java +++ b/ccc-app/src/main/java/com/ccj/app/AppConfiguration.java @@ -74,7 +74,8 @@ public Credentials credentials(AppSettings settings) { public ChatClient chatClient(AppSettings settings, Credentials credentials, SystemPromptProvider systemPromptProvider, CompactingChatMemory chatMemory, Session session, - ToolRegistry toolRegistry) { + ToolRegistry toolRegistry, + @org.springframework.context.annotation.Lazy com.ccj.tools.execution.PermissionCheckerImpl permissionChecker) { ChatClientFactory factory = new ChatClientFactory(settings, credentials); // 构建 advisor:注入 system prompt + 历史 + (Phase 5)cache_edits SystemPromptProvider.SessionContext sessionContext = new SystemPromptProvider.SessionContext( @@ -91,9 +92,24 @@ public ChatClient chatClient(AppSettings settings, Credentials credentials, com.ccj.core.tool.ToolUseContext ctxTemplate = new com.ccj.core.tool.ToolUseContext( session.workingDirectory(), java.util.List.of(), () -> false, new java.util.concurrent.ConcurrentHashMap<>(), session.id(), session); + // 权限检查回调(在工具执行前检查权限) + com.ccj.tools.adapter.ToolCallbackAdapter.PermissionChecker permChecker = (tool, input, ctx) -> { + if (permissionChecker == null) return null; // 无权限检查器时允许 + try { + com.ccj.core.tool.PermissionResult result = permissionChecker.check(tool, input, ctx); + if (result instanceof com.ccj.core.tool.PermissionResult.Deny d) { + return d.message(); + } + // Allow 和 Ask 都放行(Ask 在 REPL 层处理,此处简化为允许) + return null; + } catch (Exception e) { + log.warn("Permission check failed for {}: {}", tool.name(), e.getMessage()); + return null; // 出错时允许(fail-open) + } + }; java.util.List toolCallbacks = toolRegistry.all().stream() .filter(Tool::isEnabled) - .map(t -> (org.springframework.ai.tool.ToolCallback) new com.ccj.tools.adapter.ToolCallbackAdapter(t, ctxTemplate)) + .map(t -> (org.springframework.ai.tool.ToolCallback) new com.ccj.tools.adapter.ToolCallbackAdapter(t, ctxTemplate, permChecker)) .toList(); log.info("Registered {} tool callbacks on ChatClient", toolCallbacks.size()); return factory.createChatClientWithAdvisorAndTools(advisor, toolCallbacks); diff --git a/ccc-app/src/main/java/com/ccj/app/ClaudeCodeJavaApplication.java b/ccc-app/src/main/java/com/ccj/app/ClaudeCodeJavaApplication.java index 5bcad48..28a581b 100644 --- a/ccc-app/src/main/java/com/ccj/app/ClaudeCodeJavaApplication.java +++ b/ccc-app/src/main/java/com/ccj/app/ClaudeCodeJavaApplication.java @@ -38,6 +38,23 @@ public CommandLineRunner replRunner(AppSettings settings, Session session, // agent 调用器:AgentInvoker 构建结构化 Prompt,advisor 链负责 system prompt/历史/cache 注入 com.ccj.app.invoker.AgentInvoker agentInvoker = new com.ccj.app.invoker.AgentInvoker(chatClient); + + // 工具执行进度回调:在终端显示 "⏳ Working..." 状态 + agentInvoker.setToolProgressCallback(new com.ccj.app.invoker.AgentInvoker.ToolProgressCallback() { + @Override public void onStart() { + System.out.print("\r⏳ Working... "); + System.out.flush(); + } + @Override public void onComplete(String finalText) { + System.out.print("\r \r"); // 清除 spinner + System.out.flush(); + } + @Override public void onError(String errorMessage) { + System.out.print("\r \r"); + System.out.flush(); + } + }); + ReplMainLoop repl = new ReplMainLoop( settings, session, memory, commandRegistry, userInput -> agentInvoker.invoke(userInput).text()); diff --git a/ccc-app/src/main/java/com/ccj/app/invoker/AgentInvoker.java b/ccc-app/src/main/java/com/ccj/app/invoker/AgentInvoker.java index c284463..e1dacdb 100644 --- a/ccc-app/src/main/java/com/ccj/app/invoker/AgentInvoker.java +++ b/ccc-app/src/main/java/com/ccj/app/invoker/AgentInvoker.java @@ -1,7 +1,5 @@ package com.ccj.app.invoker; -import com.ccj.core.message.ContentBlock; -import com.ccj.core.message.Message; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.chat.client.ChatClient; @@ -10,36 +8,45 @@ import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.chat.prompt.Prompt; -import java.util.List; import java.util.Optional; +import java.util.function.Consumer; /** * Agent 调用器。替换原 ClaudeCodeJavaApplication 中的 Function<String,String> 裸 lambda。 * - * 对应 Claude Code services/api/claude.ts 的请求发起层: + * 职责: * - 构建 Prompt(用户文本 -> UserMessage) - * - 通过 ChatClient(已装配 advisor 链)发起调用 - * - 返回 InvocationResult(assistant 文本 + usage + 可选的多块内容) + * - 通过 ChatClient(已装配 advisor 链 + 工具回调)发起调用 + * - 返回 InvocationResult(assistant 文本 + usage) * - * advisor 链负责:system prompt 注入、历史注入、cache_control 标记、(Phase 5)cache_edits。 - * 本类只负责触发调用并解析响应。 + * 工具调用循环由 Spring AI 的 ToolCallAdvisor 自动驱动(.call() 路径)。 + * ToolProgressCallback 在工具执行期间提供进度反馈。 */ public class AgentInvoker { private static final Logger log = LoggerFactory.getLogger(AgentInvoker.class); private final ChatClient chatClient; + private volatile ToolProgressCallback toolProgressCallback; public AgentInvoker(ChatClient chatClient) { this.chatClient = chatClient; } + /** 设置工具执行进度回调(由 REPL 注册,显示 spinner)。 */ + public void setToolProgressCallback(ToolProgressCallback callback) { + this.toolProgressCallback = callback; + } + /** - * 调用 agent。 + * 调用 agent(阻塞式,工具调用循环由 ToolCallAdvisor 驱动)。 * @param userInput 用户输入文本 * @return 调用结果(assistant 文本 + usage) */ public InvocationResult invoke(String userInput) { + // 通知 REPL 开始处理 + if (toolProgressCallback != null) toolProgressCallback.onStart(); + try { UserMessage userMessage = new UserMessage(userInput); Prompt prompt = new Prompt(userMessage); @@ -50,9 +57,13 @@ public InvocationResult invoke(String userInput) { String text = extractText(response); Usage usage = response.getMetadata() != null ? response.getMetadata().getUsage() : null; + // 通知 REPL 处理完成 + if (toolProgressCallback != null) toolProgressCallback.onComplete(text); + return new InvocationResult(text, usage, null); } catch (Exception e) { log.error("Agent invocation failed", e); + if (toolProgressCallback != null) toolProgressCallback.onError(e.getMessage()); return new InvocationResult("[Error: " + e.getMessage() + "]", null, e); } } @@ -76,4 +87,16 @@ public record InvocationResult(String text, Usage usage, Throwable error) { public boolean isSuccess() { return error == null; } public Optional usageOpt() { return Optional.ofNullable(usage); } } + + /** + * 工具执行进度回调。由 REPL 实现,在工具调用期间显示进度。 + */ + public interface ToolProgressCallback { + /** 开始处理用户请求。 */ + void onStart(); + /** 处理完成,收到最终响应。 */ + void onComplete(String finalText); + /** 处理出错。 */ + void onError(String errorMessage); + } } diff --git a/ccc-test/src/test/java/com/ccj/test/PermissionAwareAdapterTest.java b/ccc-test/src/test/java/com/ccj/test/PermissionAwareAdapterTest.java new file mode 100644 index 0000000..ec5766d --- /dev/null +++ b/ccc-test/src/test/java/com/ccj/test/PermissionAwareAdapterTest.java @@ -0,0 +1,78 @@ +package com.ccj.test; + +import com.ccj.core.tool.Tool; +import com.ccj.core.tool.ToolResult; +import com.ccj.core.tool.ToolUseContext; +import com.ccj.tools.adapter.ToolCallbackAdapter; +import org.junit.jupiter.api.Test; + +import java.util.Map; +import java.util.concurrent.CompletableFuture; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * 验证 ToolCallbackAdapter 的权限检查功能。 + */ +class PermissionAwareAdapterTest { + + @Test + void noPermissionCheckerAllowsAll() { + TestTool tool = new TestTool("Echo"); + ToolCallbackAdapter adapter = new ToolCallbackAdapter(tool, null); + String result = adapter.call("{\"msg\":\"hello\"}"); + assertThat(result).contains("hello"); + assertThat(result).contains("\"isError\":false"); + } + + @Test + void permissionCheckerAllowsExecution() { + TestTool tool = new TestTool("Read"); + ToolCallbackAdapter.PermissionChecker checker = (t, i, c) -> null; // allow + ToolCallbackAdapter adapter = new ToolCallbackAdapter(tool, null, checker); + String result = adapter.call("{\"msg\":\"content\"}"); + assertThat(result).contains("content"); + } + + @Test + void permissionCheckerDeniesExecution() { + TestTool tool = new TestTool("Bash"); + ToolCallbackAdapter.PermissionChecker checker = (t, i, c) -> "Dangerous command blocked"; + ToolCallbackAdapter adapter = new ToolCallbackAdapter(tool, null, checker); + String result = adapter.call("{\"msg\":\"rm -rf /\"}"); + assertThat(result).contains("Permission denied"); + assertThat(result).contains("Dangerous command blocked"); + assertThat(result).contains("\"isError\":true"); + assertThat(tool.callCount).isZero(); // tool.call never invoked + } + + @Test + void permissionCheckerExceptionFailsOpen() { + TestTool tool = new TestTool("Read"); + ToolCallbackAdapter.PermissionChecker checker = (t, i, c) -> { + throw new RuntimeException("checker crashed"); + }; + ToolCallbackAdapter adapter = new ToolCallbackAdapter(tool, null, checker); + String result = adapter.call("{\"msg\":\"file.txt\"}"); + // fail-open: permission checker exception doesn't block execution + assertThat(result).contains("file.txt"); + } + + static class TestTool extends Tool { + final String name; + int callCount = 0; + + TestTool(String name) { this.name = name; } + + @Override public String name() { return name; } + @Override public String description() { return "test"; } + @Override public Map inputSchema() { + return Map.of("type", "object", "properties", Map.of("msg", Map.of("type", "string"))); + } + @Override + public CompletableFuture call(Map input, ToolUseContext context) { + callCount++; + return CompletableFuture.completedFuture(ToolResult.success((String) input.get("msg"))); + } + } +} diff --git a/ccc-tools/src/main/java/com/ccj/tools/adapter/ToolCallbackAdapter.java b/ccc-tools/src/main/java/com/ccj/tools/adapter/ToolCallbackAdapter.java index efafe9e..44b44bc 100644 --- a/ccc-tools/src/main/java/com/ccj/tools/adapter/ToolCallbackAdapter.java +++ b/ccc-tools/src/main/java/com/ccj/tools/adapter/ToolCallbackAdapter.java @@ -33,10 +33,26 @@ public class ToolCallbackAdapter implements ToolCallback { private final Tool tool; private final ToolUseContext contextTemplate; + private final PermissionChecker permissionChecker; + + /** 权限检查器接口(可选)。设置后在每次工具调用前检查权限。 */ + @FunctionalInterface + public interface PermissionChecker { + /** + * 检查工具调用权限。 + * @return null 表示允许;非 null 返回拒绝原因 + */ + String check(Tool tool, Map input, ToolUseContext context); + } public ToolCallbackAdapter(Tool tool, ToolUseContext contextTemplate) { + this(tool, contextTemplate, null); + } + + public ToolCallbackAdapter(Tool tool, ToolUseContext contextTemplate, PermissionChecker permissionChecker) { this.tool = tool; this.contextTemplate = contextTemplate; + this.permissionChecker = permissionChecker; } @Override @@ -60,6 +76,23 @@ public String call(String toolInput, ToolContext toolContext) { Map input = toolInput == null || toolInput.isBlank() ? Map.of() : M.readValue(toolInput, Map.class); ToolUseContext ctx = mergeContext(toolContext); + + // 权限检查(如果配置了 permissionChecker) + if (permissionChecker != null) { + try { + String denied = permissionChecker.check(tool, input, ctx); + if (denied != null) { + log.info("Tool {} denied by permission checker: {}", tool.name(), denied); + return M.writeValueAsString(Map.of( + "content", "Permission denied: " + denied, + "isError", true)); + } + } catch (Exception e) { + // 权限检查器异常:fail-open(记录但不阻断) + log.warn("Permission checker exception for {}: {}", tool.name(), e.getMessage()); + } + } + ToolResult result = tool.call(input, ctx).join(); return M.writeValueAsString(Map.of( "content", result.asText(),