Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 18 additions & 2 deletions ccc-app/src/main/java/com/ccj/app/AppConfiguration.java
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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)
}
Comment on lines +99 to +108
};
java.util.List<org.springframework.ai.tool.ToolCallback> 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);
Expand Down
17 changes: 17 additions & 0 deletions ccc-app/src/main/java/com/ccj/app/ClaudeCodeJavaApplication.java
Original file line number Diff line number Diff line change
Expand Up @@ -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());
Expand Down
41 changes: 32 additions & 9 deletions ccc-app/src/main/java/com/ccj/app/invoker/AgentInvoker.java
Original file line number Diff line number Diff line change
@@ -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;
Expand All @@ -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;
Comment on lines 11 to +12

/**
* Agent 调用器。替换原 ClaudeCodeJavaApplication 中的 Function&lt;String,String&gt; 裸 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);
Expand All @@ -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);
}
}
Expand All @@ -76,4 +87,16 @@ public record InvocationResult(String text, Usage usage, Throwable error) {
public boolean isSuccess() { return error == null; }
public Optional<Usage> usageOpt() { return Optional.ofNullable(usage); }
}

/**
* 工具执行进度回调。由 REPL 实现,在工具调用期间显示进度。
*/
public interface ToolProgressCallback {
/** 开始处理用户请求。 */
void onStart();
/** 处理完成,收到最终响应。 */
void onComplete(String finalText);
/** 处理出错。 */
void onError(String errorMessage);
}
}
Original file line number Diff line number Diff line change
@@ -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<String, Object> inputSchema() {
return Map.of("type", "object", "properties", Map.of("msg", Map.of("type", "string")));
}
@Override
public CompletableFuture<ToolResult> call(Map<String, Object> input, ToolUseContext context) {
callCount++;
return CompletableFuture.completedFuture(ToolResult.success((String) input.get("msg")));
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Comment on lines +38 to +40
/**
* 检查工具调用权限。
* @return null 表示允许;非 null 返回拒绝原因
*/
String check(Tool tool, Map<String, Object> 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;
}
Comment on lines +52 to 56

@Override
Expand All @@ -60,6 +76,23 @@ public String call(String toolInput, ToolContext toolContext) {
Map<String, Object> 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(),
Expand Down