diff --git a/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/ClaudeAsyncClient.java b/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/ClaudeAsyncClient.java index cc8fdf8a..9b2de2e4 100644 --- a/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/ClaudeAsyncClient.java +++ b/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/ClaudeAsyncClient.java @@ -335,7 +335,9 @@ default Flux queryText(String prompt) { /** * Interrupts the current Claude operation. - * @return Mono that completes when the interrupt has been sent + * @return Mono that completes once the CLI accepts the interrupt, and fails with a + * {@link io.github.markpollack.claude.agent.sdk.exceptions.ClaudeSDKException} if it + * refuses or does not reply within the client timeout */ Mono interrupt(); @@ -343,14 +345,18 @@ default Flux queryText(String prompt) { * Sets the permission mode for tool execution. * @param mode the permission mode (e.g., "default", "acceptEdits", * "bypassPermissions") - * @return Mono that completes when the mode has been set + * @return Mono that completes once the CLI accepts the mode, and fails with a + * {@link io.github.markpollack.claude.agent.sdk.exceptions.ClaudeSDKException} if it + * refuses or does not reply within the client timeout */ Mono setPermissionMode(String mode); /** * Changes the Claude model during the session. * @param model the model ID to switch to - * @return Mono that completes when the model has been changed + * @return Mono that completes once the CLI accepts the model, and fails with a + * {@link io.github.markpollack.claude.agent.sdk.exceptions.ClaudeSDKException} if it + * refuses or does not reply within the client timeout */ Mono setModel(String model); diff --git a/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/DefaultClaudeAsyncClient.java b/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/DefaultClaudeAsyncClient.java index 14bed906..41a1b6fb 100644 --- a/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/DefaultClaudeAsyncClient.java +++ b/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/DefaultClaudeAsyncClient.java @@ -35,14 +35,12 @@ import io.github.markpollack.claude.agent.sdk.types.control.ControlResponse; import io.github.markpollack.claude.agent.sdk.types.control.HookEvent; import io.github.markpollack.claude.agent.sdk.types.control.HookInput; -import io.github.markpollack.claude.agent.sdk.types.control.HookOutput; import io.github.markpollack.claude.agent.sdk.permission.PermissionResult; import io.github.markpollack.claude.agent.sdk.permission.ToolPermissionCallback; import io.github.markpollack.claude.agent.sdk.permission.ToolPermissionContext; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; -import reactor.core.publisher.MonoSink; import reactor.core.publisher.Sinks; import reactor.core.scheduler.Schedulers; @@ -51,6 +49,7 @@ import java.util.*; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.TimeoutException; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.atomic.AtomicReference; @@ -151,12 +150,12 @@ public class DefaultClaudeAsyncClient implements ClaudeAsyncClient { */ private volatile Sinks.Many rawMessageSink; - // Control request handling (MCP SDK pattern using MonoSink for correlation) + // Control request handling (one Sinks.One per request, matched by request_id) private final AtomicInteger requestCounter = new AtomicInteger(0); private final String sessionPrefix = UUID.randomUUID().toString().substring(0, 8); - private final ConcurrentHashMap>> pendingResponses = new ConcurrentHashMap<>(); + private final ConcurrentHashMap>> pendingResponses = new ConcurrentHashMap<>(); // Cross-turn message handlers (thread-safe for concurrent registration) private final List> messageHandlers = new CopyOnWriteArrayList<>(); @@ -288,7 +287,10 @@ private void sendInitialize() throws ClaudeSDKException { request.put("hooks", hookConfig); logger.debug("Sending initialize with {} hook event types", hookConfig.size()); - sendControlRequest(request); + // Awaited, as in the sync client: the prompt follows right after, and the CLI + // must have registered the hooks before it starts the turn. A refusal fails the + // connect instead of leaving the hooks silently inactive. + sendControlRequest(request).block(); logger.info("Hook configuration sent to CLI: {} event types", hookConfig.size()); } @@ -409,58 +411,34 @@ public Flux receiveResponse() { @Override public Mono interrupt() { - return Mono.create(sink -> { - if (!connected.get() || closed.get()) { - sink.error(new IllegalStateException("Client is not connected")); - return; - } - try { - sendControlRequest(Map.of("subtype", "interrupt")); - sink.success(); - } - catch (Exception e) { - sink.error(new TransportException("Failed to send interrupt", e)); - } - }).subscribeOn(Schedulers.boundedElastic()); + return awaitControlRequest(Map.of("subtype", "interrupt")); } @Override public Mono setPermissionMode(String mode) { - return Mono.create(sink -> { - if (!connected.get() || closed.get()) { - sink.error(new IllegalStateException("Client is not connected")); - return; - } - try { - sendControlRequest(Map.of("subtype", "set_permission_mode", "mode", mode)); - currentPermissionMode.set(mode); - sink.success(); - } - catch (Exception e) { - sink.error(new TransportException("Failed to set permission mode", e)); - } - }).subscribeOn(Schedulers.boundedElastic()); + return awaitControlRequest(Map.of("subtype", "set_permission_mode", "mode", mode)) + .doOnSuccess(ignored -> currentPermissionMode.set(mode)); } @Override public Mono setModel(String model) { - return Mono.create(sink -> { + Map request = new LinkedHashMap<>(); + request.put("subtype", "set_model"); + request.put("model", model); + return awaitControlRequest(request).doOnSuccess(ignored -> currentModel.set(model)); + } + + /** + * Sends a control request on subscription and completes once the CLI accepts it. A + * refusal or a missing reply fails the returned Mono, as the sync client throws. + */ + private Mono awaitControlRequest(Map request) { + return Mono.defer(() -> { if (!connected.get() || closed.get()) { - sink.error(new IllegalStateException("Client is not connected")); - return; - } - try { - Map request = new LinkedHashMap<>(); - request.put("subtype", "set_model"); - request.put("model", model); - sendControlRequest(request); - currentModel.set(model); - sink.success(); + return Mono.>error(new IllegalStateException("Client is not connected")); } - catch (Exception e) { - sink.error(new TransportException("Failed to set model", e)); - } - }).subscribeOn(Schedulers.boundedElastic()); + return sendControlRequest(request); + }).subscribeOn(Schedulers.boundedElastic()).then(); } @Override @@ -663,31 +641,8 @@ else if (payload instanceof ControlRequest.McpMessageRequest mcpMessage) { private ControlResponse handleHookCallback(String requestId, ControlRequest.HookCallbackRequest hookCallback) { try { - String callbackId = hookCallback.callbackId(); - Map inputMap = hookCallback.input(); - - HookInput input = objectMapper.convertValue(inputMap, HookInput.class); - HookOutput output = hookRegistry.executeHook(callbackId, input); - - Map responsePayload = new LinkedHashMap<>(); - responsePayload.put("continue", output.continueExecution()); - if (output.decision() != null) { - responsePayload.put("decision", output.decision()); - } - if (output.reason() != null) { - responsePayload.put("reason", output.reason()); - } - if (output.hookSpecificOutput() != null) { - HookOutput.HookSpecificOutput specific = output.hookSpecificOutput(); - if (specific.permissionDecision() != null) { - responsePayload.put("permission_decision", specific.permissionDecision()); - } - if (specific.permissionDecisionReason() != null) { - responsePayload.put("permission_decision_reason", specific.permissionDecisionReason()); - } - } - - return ControlResponse.success(requestId, responsePayload); + HookInput input = objectMapper.convertValue(hookCallback.input(), HookInput.class); + return hookRegistry.handleCallback(requestId, hookCallback.callbackId(), input); } catch (Exception e) { logger.error("Hook callback failed", e); @@ -759,7 +714,7 @@ private void handleControlResponse(ControlResponse response) { logger.debug("Handling control response: requestId={}, subtype={}", requestId, response.response().subtype()); - MonoSink> sink = pendingResponses.remove(requestId); + Sinks.One> sink = pendingResponses.remove(requestId); if (sink == null) { logger.warn("Unexpected response for unknown request id {}", requestId); return; @@ -774,35 +729,58 @@ private void handleControlResponse(ControlResponse response) { Map typedMap = (Map) responseMap; payload.putAll(typedMap); } - sink.success(payload); + sink.tryEmitValue(payload); logger.debug("Control response delivered for requestId={}", requestId); } else if (response.response() instanceof ControlResponse.ErrorPayload error) { - sink.error(new ClaudeSDKException("Control request failed: " + error.error())); - logger.debug("Control response error delivered for requestId={}", requestId); + // Logged here, so a refusal shows even when nobody waits for the reply + logger.warn("Control request {} failed: {}", requestId, error.error()); + sink.tryEmitError(new ClaudeSDKException("Control request failed: " + error.error())); } else { - sink.success(payload); + sink.tryEmitValue(payload); } } - private void sendControlRequest(Map request) throws ClaudeSDKException { - try { - String requestId = sessionPrefix + "_" + requestCounter.incrementAndGet(); + /** + * Sends a control request and returns the CLI's reply. + * + *

+ * The reply is registered before the request is sent, so it is matched rather than + * reported as an unknown response, whether or not the caller waits for it. A caller + * that waits gets the reply, or the refusal as a {@link ClaudeSDKException}, within + * the client timeout. + *

+ * @throws TransportException if the request cannot be written to the CLI + */ + private Mono> sendControlRequest(Map request) throws ClaudeSDKException { + String requestId = sessionPrefix + "_" + requestCounter.incrementAndGet(); - Map fullRequest = new LinkedHashMap<>(); - fullRequest.put("type", "control"); - fullRequest.put("request_id", requestId); - fullRequest.putAll(request); + // The CLI only reads the control_request envelope with the payload nested + // under "request", as the sync client sends it; any other shape is dropped + // without a reply. + Map fullRequest = new LinkedHashMap<>(); + fullRequest.put("type", "control_request"); + fullRequest.put("request_id", requestId); + fullRequest.put("request", request); - String json = objectMapper.writeValueAsString(fullRequest); - transportRef.get().sendMessage(json); + Sinks.One> reply = Sinks.one(); + pendingResponses.put(requestId, reply); - logger.debug("Sent control request: id={}, subtype={}", requestId, request.get("subtype")); + try { + transportRef.get().sendMessage(objectMapper.writeValueAsString(fullRequest)); } catch (Exception e) { + pendingResponses.remove(requestId); throw new TransportException("Failed to send control request", e); } + + logger.debug("Sent control request: id={}, subtype={}", requestId, request.get("subtype")); + return reply.asMono() + .timeout(timeout) + .onErrorMap(TimeoutException.class, + e -> new ClaudeSDKException("Control request timed out: " + request.get("subtype"), e)) + .doOnError(e -> pendingResponses.remove(requestId)); } private void cleanup() { @@ -828,6 +806,9 @@ private void cleanup() { rawMessageSink = null; } + // Fail pending requests now, so a connect awaiting initialize is released + // instead of waiting out the timeout + pendingResponses.values().forEach(sink -> sink.tryEmitError(new TransportException("Client closed"))); pendingResponses.clear(); } diff --git a/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/DefaultClaudeSyncClient.java b/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/DefaultClaudeSyncClient.java index 00f7c5a2..7f93b300 100644 --- a/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/DefaultClaudeSyncClient.java +++ b/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/DefaultClaudeSyncClient.java @@ -39,7 +39,6 @@ import io.github.markpollack.claude.agent.sdk.types.control.ControlResponse; import io.github.markpollack.claude.agent.sdk.types.control.HookEvent; import io.github.markpollack.claude.agent.sdk.types.control.HookInput; -import io.github.markpollack.claude.agent.sdk.types.control.HookOutput; import io.github.markpollack.claude.agent.sdk.permission.PermissionResult; import io.github.markpollack.claude.agent.sdk.permission.ToolPermissionCallback; import io.github.markpollack.claude.agent.sdk.permission.ToolPermissionContext; @@ -485,34 +484,8 @@ else if (payload instanceof ControlRequest.McpMessageRequest mcpMessage) { private ControlResponse handleHookCallback(String requestId, ControlRequest.HookCallbackRequest hookCallback) { try { - String callbackId = hookCallback.callbackId(); - Map inputMap = hookCallback.input(); - - HookInput input = objectMapper.convertValue(inputMap, HookInput.class); - HookOutput output = hookRegistry.executeHook(callbackId, input); - - Map responsePayload = new LinkedHashMap<>(); - responsePayload.put("continue", output.continueExecution()); - if (output.decision() != null) { - responsePayload.put("decision", output.decision()); - } - if (output.reason() != null) { - responsePayload.put("reason", output.reason()); - } - if (output.hookSpecificOutput() != null) { - HookOutput.HookSpecificOutput specific = output.hookSpecificOutput(); - if (specific.permissionDecision() != null) { - responsePayload.put("permission_decision", specific.permissionDecision()); - } - if (specific.permissionDecisionReason() != null) { - responsePayload.put("permission_decision_reason", specific.permissionDecisionReason()); - } - if (specific.updatedInput() != null) { - responsePayload.put("updated_input", specific.updatedInput()); - } - } - - return ControlResponse.success(requestId, responsePayload); + HookInput input = objectMapper.convertValue(hookCallback.input(), HookInput.class); + return hookRegistry.handleCallback(requestId, hookCallback.callbackId(), input); } catch (Exception e) { logger.error("Error executing hook callback", e); diff --git a/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/hooks/HookRegistry.java b/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/hooks/HookRegistry.java index bbfd4d3b..1a890b51 100644 --- a/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/hooks/HookRegistry.java +++ b/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/hooks/HookRegistry.java @@ -19,6 +19,7 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import io.github.markpollack.claude.agent.sdk.types.control.ControlRequest; +import io.github.markpollack.claude.agent.sdk.types.control.ControlResponse; import io.github.markpollack.claude.agent.sdk.types.control.HookEvent; import io.github.markpollack.claude.agent.sdk.types.control.HookInput; import io.github.markpollack.claude.agent.sdk.types.control.HookOutput; @@ -247,20 +248,75 @@ public HookOutput executeHook(String hookId, HookInput input) { logger.warn("Hook not found: {}", hookId); return null; } + return execute(registration, input); + } + private HookOutput execute(HookRegistration registration, HookInput input) { try { - logger.debug("Executing hook: id={}, event={}", hookId, registration.event()); + logger.debug("Executing hook: id={}, event={}", registration.id(), registration.event()); HookOutput output = registration.callback().handle(input); - logger.debug("Hook result: id={}, continue={}", hookId, output.continueExecution()); + logger.debug("Hook result: id={}, continue={}", registration.id(), output.continueExecution()); return output; } catch (Exception e) { - logger.error("Hook execution failed: id={}", hookId, e); + logger.error("Hook execution failed: id={}", registration.id(), e); // Return a safe default on error return HookOutput.block("Hook execution failed: " + e.getMessage()); } } + /** + * Executes the hook a {@code hook_callback} control request names and wraps its + * output in the control response the CLI expects. + * + *

+ * The CLI validates the response against the hook JSON output format, the same one a + * command hook prints: {@code continue}, {@code decision}, {@code reason} and the + * other control fields at the top level, and {@code hookSpecificOutput} nested with + * camelCase keys ({@code hookEventName}, {@code permissionDecision}, + * {@code updatedInput}, {@code additionalContext}). {@link HookOutput} already + * serializes to exactly that and omits unset fields, so it is sent as is. An unset + * field must be absent rather than null: the CLI rejects {@code "continue": null} and + * then ignores the whole output, deny included. + *

+ * + *

+ * The CLI also rejects a {@code hookSpecificOutput} whose {@code hookEventName} is + * not the event it asked for, so a missing name is filled in from the registration. A + * name that differs is the hook's mistake: it is logged and sent as is. + *

+ * @param requestId the control request ID to answer + * @param hookId the hook ID the CLI asked for + * @param input the hook input + * @return a success response carrying the hook output, or an error response if no + * hook is registered under {@code hookId} + */ + public ControlResponse handleCallback(String requestId, String hookId, HookInput input) { + HookRegistration registration = hooksById.get(hookId); + if (registration == null) { + logger.warn("Hook not found: {}", hookId); + return ControlResponse.error(requestId, "No hook registered for callback ID: " + hookId); + } + return ControlResponse.success(requestId, withEventName(execute(registration, input), registration)); + } + + private static HookOutput withEventName(HookOutput output, HookRegistration registration) { + HookOutput.HookSpecificOutput specific = output.hookSpecificOutput(); + if (specific == null) { + return output; + } + String expected = registration.event().getProtocolName(); + if (specific.hookEventName() == null) { + return output.withHookSpecificOutput(specific.withHookEventName(expected)); + } + if (!expected.equals(specific.hookEventName())) { + logger.warn( + "Hook {} is registered for {} but returned hookSpecificOutput for {}; the CLI will ignore its whole output", + registration.id(), expected, specific.hookEventName()); + } + return output; + } + /** * Builds the hook configuration for CLI initialization. Returns a map from event * names to hook matcher configs. diff --git a/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/types/control/HookOutput.java b/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/types/control/HookOutput.java index ebd82dbd..4754f9ed 100644 --- a/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/types/control/HookOutput.java +++ b/claude-code-sdk/src/main/java/io/github/markpollack/claude/agent/sdk/types/control/HookOutput.java @@ -25,9 +25,10 @@ * Output from a hook execution. Sent back to CLI as part of control response. * *

- * Field naming note: Java uses camelCase but the protocol uses snake_case. - * Jackson @JsonProperty handles the conversion. Additionally, "continue" and "async" are - * Java keywords, so we use alternative names that get serialized correctly. + * Serializes to the CLI's hook JSON output format, the same one a command hook prints: + * camelCase keys, with {@link HookSpecificOutput} nested under + * {@code hookSpecificOutput}, and unset fields omitted. "continue" and "async" are Java + * keywords, so their components use alternative names that @JsonProperty maps back. */ @JsonInclude(JsonInclude.Include.NON_NULL) public record HookOutput( @@ -80,6 +81,14 @@ public static HookOutput async(int timeoutMs) { return builder().asyncExecution(true).asyncTimeout(timeoutMs).continueExecution(true).build(); } + /** + * Returns a copy of this output with {@code hookSpecificOutput} replaced. + */ + public HookOutput withHookSpecificOutput(HookSpecificOutput hookSpecificOutput) { + return new HookOutput(continueExecution, suppressOutput, stopReason, decision, systemMessage, reason, + asyncExecution, asyncTimeout, hookSpecificOutput); + } + /** * Create builder for fluent construction. */ @@ -225,6 +234,14 @@ public static HookSpecificOutput userPromptSubmit(String additionalContext) { return new HookSpecificOutput("UserPromptSubmit", null, null, null, additionalContext); } + /** + * Returns a copy of this output with {@code hookEventName} replaced. + */ + public HookSpecificOutput withHookEventName(String hookEventName) { + return new HookSpecificOutput(hookEventName, permissionDecision, permissionDecisionReason, updatedInput, + additionalContext); + } + /** * Builder for HookSpecificOutput. */ diff --git a/claude-code-sdk/src/test/java/io/github/markpollack/claude/agent/sdk/HookCallbackWireTest.java b/claude-code-sdk/src/test/java/io/github/markpollack/claude/agent/sdk/HookCallbackWireTest.java new file mode 100644 index 00000000..7945ee3d --- /dev/null +++ b/claude-code-sdk/src/test/java/io/github/markpollack/claude/agent/sdk/HookCallbackWireTest.java @@ -0,0 +1,319 @@ +/* + * Copyright 2026 Mark Pollack + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * 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 io.github.markpollack.claude.agent.sdk; + +import java.io.IOException; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.attribute.PosixFilePermissions; +import java.time.Duration; +import java.util.ArrayList; +import java.util.List; +import java.util.Optional; +import java.util.function.Predicate; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import io.github.markpollack.claude.agent.sdk.exceptions.ClaudeSDKException; +import io.github.markpollack.claude.agent.sdk.exceptions.TransportException; +import io.github.markpollack.claude.agent.sdk.hooks.HookRegistry; +import io.github.markpollack.claude.agent.sdk.types.control.HookOutput; +import io.github.markpollack.claude.agent.sdk.types.control.HookOutput.HookSpecificOutput; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.condition.DisabledOnOs; +import org.junit.jupiter.api.condition.OS; +import org.junit.jupiter.api.io.TempDir; +import reactor.core.Disposable; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +/** + * Drives each client through hook registration and one {@code hook_callback} round trip + * against a stub CLI, and checks what it writes to the CLI's stdin. + * + *

+ * The stub answers {@code initialize}, sends a PreToolUse {@code hook_callback} when the + * user message arrives, and records everything the SDK writes to its stdin. A PreToolUse + * deny must reach the CLI as the hook JSON output format, {@code hookSpecificOutput} + * nested with camelCase keys and no {@code "continue": null}; anything else is discarded + * by the CLI and the tool runs. + *

+ */ +@DisabledOnOs(OS.WINDOWS) +@DisplayName("Hook callback responses on the wire") +class HookCallbackWireTest { + + private static final Duration ARRIVAL_TIMEOUT = Duration.ofSeconds(10); + + private static final ObjectMapper MAPPER = new ObjectMapper(); + + private static final String HOOK_REQUEST_ID = "hook-req-1"; + + private static final String DENY_REASON = "branch changes are blocked"; + + private static final String INITIALIZE_SUCCESS = """ + {"type":"control_response","response":{"subtype":"success","request_id":"%s","response":{}}}"""; + + private static final String INITIALIZE_REFUSED = """ + {"type":"control_response","response":{"subtype":"error","request_id":"%s","error":"hooks rejected"}}"""; + + private static final String CONTROL_SUCCESS = """ + {"type":"control_response","response":{"subtype":"success","request_id":"%s","response":{}}}"""; + + private static final String CONTROL_REFUSED = """ + {"type":"control_response","response":{"subtype":"error","request_id":"%s","error":"request rejected"}}"""; + + @TempDir + Path tempDir; + + private Path recording; + + private String stubCli; + + private HookRegistry hooks; + + @BeforeEach + void setUp() throws IOException { + this.recording = tempDir.resolve("cli-stdin.jsonl"); + Files.createFile(recording); + this.stubCli = writeStubCli(INITIALIZE_SUCCESS); + this.hooks = new HookRegistry(); + hooks.registerPreToolUse("Bash", + input -> HookOutput.builder() + .hookSpecificOutput(HookSpecificOutput.preToolUseDeny(DENY_REASON)) + .build()); + } + + @Nested + @DisplayName("ClaudeSyncClient") + class Sync { + + @Test + @DisplayName("registers hooks with a control_request initialize and answers a PreToolUse deny in the hook JSON output format") + void preToolUseDeny() throws Exception { + try (ClaudeSyncClient client = ClaudeClient.sync() + .workingDirectory(tempDir) + .claudePath(stubCli) + .hookRegistry(hooks) + .timeout(ARRIVAL_TIMEOUT) + .build()) { + client.connect("run git switch -c probe"); + + assertInitializeRegistersHook(awaitInitialize()); + assertDenyResponse(awaitHookResponse()); + } + } + + } + + @Nested + @DisplayName("ClaudeAsyncClient") + class Async { + + @Test + @DisplayName("registers hooks with a control_request initialize and answers a PreToolUse deny in the hook JSON output format") + void preToolUseDeny() throws Exception { + ClaudeAsyncClient client = asyncClient(stubCli); + try { + Disposable turn = client.connect("run git switch -c probe").messages().subscribe(); + try { + assertInitializeRegistersHook(awaitInitialize()); + assertDenyResponse(awaitHookResponse()); + } + finally { + turn.dispose(); + } + } + finally { + client.close().block(ARRIVAL_TIMEOUT); + } + } + + @Test + @DisplayName("a refused initialize fails the connect") + void initializeRefused() throws Exception { + ClaudeAsyncClient client = asyncClient(writeStubCli(INITIALIZE_REFUSED)); + try { + assertThatThrownBy(() -> client.connect().block(ARRIVAL_TIMEOUT)).isInstanceOf(TransportException.class) + .rootCause() + .hasMessageContaining("hooks rejected"); + } + finally { + client.close().block(ARRIVAL_TIMEOUT); + } + } + + @Test + @DisplayName("setPermissionMode completes once the CLI accepts it") + void setPermissionModeAccepted() throws Exception { + DefaultClaudeAsyncClient client = (DefaultClaudeAsyncClient) asyncClient(stubCli); + try { + client.connect().block(ARRIVAL_TIMEOUT); + + client.setPermissionMode("plan").block(ARRIVAL_TIMEOUT); + + assertThat(client.getCurrentPermissionMode()).isEqualTo("plan"); + } + finally { + client.close().block(ARRIVAL_TIMEOUT); + } + } + + @Test + @DisplayName("a refused setModel fails and keeps the current model") + void setModelRefused() throws Exception { + DefaultClaudeAsyncClient client = (DefaultClaudeAsyncClient) asyncClient(stubCli); + try { + client.connect().block(ARRIVAL_TIMEOUT); + String before = client.getCurrentModel(); + + assertThatThrownBy(() -> client.setModel("no-such-model").block(ARRIVAL_TIMEOUT)) + .isInstanceOf(ClaudeSDKException.class) + .hasMessageContaining("request rejected"); + assertThat(client.getCurrentModel()).isEqualTo(before); + } + finally { + client.close().block(ARRIVAL_TIMEOUT); + } + } + + @Test + @DisplayName("a refused interrupt fails") + void interruptRefused() throws Exception { + ClaudeAsyncClient client = asyncClient(stubCli); + try { + client.connect().block(ARRIVAL_TIMEOUT); + + assertThatThrownBy(() -> client.interrupt().block(ARRIVAL_TIMEOUT)) + .isInstanceOf(ClaudeSDKException.class) + .hasMessageContaining("request rejected"); + } + finally { + client.close().block(ARRIVAL_TIMEOUT); + } + } + + } + + private ClaudeAsyncClient asyncClient(String cli) { + return ClaudeClient.async() + .workingDirectory(tempDir) + .claudePath(cli) + .hookRegistry(hooks) + .timeout(ARRIVAL_TIMEOUT) + .build(); + } + + private void assertDenyResponse(JsonNode response) throws IOException { + assertThat(response.at("/response/subtype").asText()).isEqualTo("success"); + assertThat(response.at("/response/response")).isEqualTo(MAPPER.readTree(""" + {"hookSpecificOutput": { + "hookEventName": "PreToolUse", + "permissionDecision": "deny", + "permissionDecisionReason": "%s"}} + """.formatted(DENY_REASON))); + } + + private void assertInitializeRegistersHook(JsonNode initialize) throws IOException { + assertThat(initialize.path("request_id").asText()).isNotBlank(); + assertThat(initialize.at("/request/hooks")).isEqualTo(MAPPER.readTree(""" + {"PreToolUse": [{"matcher": "Bash", "hookCallbackIds": ["hook_0"], "timeout": 60}]} + """)); + } + + /** + * The initialize request, in the only envelope the CLI reads: {@code control_request} + * with the payload nested under {@code request}. + */ + private JsonNode awaitInitialize() throws Exception { + return awaitLine(node -> "control_request".equals(node.path("type").asText()) + && "initialize".equals(node.at("/request/subtype").asText())); + } + + private JsonNode awaitHookResponse() throws Exception { + return awaitLine(node -> "control_response".equals(node.path("type").asText()) + && HOOK_REQUEST_ID.equals(node.at("/response/request_id").asText())); + } + + private JsonNode awaitLine(Predicate match) throws Exception { + long deadline = System.nanoTime() + ARRIVAL_TIMEOUT.toNanos(); + while (System.nanoTime() < deadline) { + Optional found = recordedLines().stream().filter(match).findFirst(); + if (found.isPresent()) { + return found.get(); + } + Thread.sleep(25); + } + throw new AssertionError("No matching line reached the CLI within " + ARRIVAL_TIMEOUT + "; recorded: " + + Files.readString(recording, StandardCharsets.UTF_8)); + } + + private List recordedLines() throws IOException { + List lines = new ArrayList<>(); + for (String line : Files.readAllLines(recording, StandardCharsets.UTF_8)) { + if (!line.isBlank()) { + lines.add(MAPPER.readTree(line)); + } + } + return lines; + } + + /** + * Writes a stand-in for the Claude CLI: it records each line the SDK sends, answers + * {@code initialize} with {@code initializeReply} (a printf format whose {@code %s} + * is the request ID), accepts {@code set_permission_mode}, refuses {@code set_model} + * and {@code interrupt}, and asks for the {@code hook_0} PreToolUse callback once a + * user message arrives. It contacts nothing. + */ + private String writeStubCli(String initializeReply) throws IOException { + Path stub = tempDir.resolve("claude-stub.sh"); + String hookCallback = """ + {"type":"control_request","request_id":"%s","request":{"subtype":"hook_callback",\ + "callback_id":"hook_0","tool_use_id":"tool_1","input":{"hook_event_name":"PreToolUse",\ + "session_id":"sess_1","transcript_path":"/tmp/t.jsonl","cwd":"/tmp","tool_name":"Bash",\ + "tool_input":{"command":"git switch -c probe"}}}}""".formatted(HOOK_REQUEST_ID); + String script = """ + #!/bin/sh + # Deterministic stand-in for the Claude CLI used by HookCallbackWireTest. + reply() { + id=$(printf '%%s\\n' "$line" | sed -n 's/.*"request_id":"\\([^"]*\\)".*/\\1/p') + printf "$1\\n" "$id" + } + while IFS= read -r line; do + printf '%%s\\n' "$line" >> '%s' + case "$line" in + *'"subtype":"initialize"'*) reply '%s' ;; + *'"subtype":"set_permission_mode"'*) reply '%s' ;; + *'"subtype":"set_model"'*|*'"subtype":"interrupt"'*) reply '%s' ;; + *'"type":"user"'*) + printf '%%s\\n' '%s' + ;; + esac + done + """.formatted(recording.toAbsolutePath(), initializeReply, CONTROL_SUCCESS, CONTROL_REFUSED, + hookCallback); + Files.writeString(stub, script, StandardCharsets.UTF_8); + Files.setPosixFilePermissions(stub, PosixFilePermissions.fromString("rwxr-xr-x")); + return stub.toAbsolutePath().toString(); + } + +} diff --git a/claude-code-sdk/src/test/java/io/github/markpollack/claude/agent/sdk/hooks/HookDecisionIT.java b/claude-code-sdk/src/test/java/io/github/markpollack/claude/agent/sdk/hooks/HookDecisionIT.java new file mode 100644 index 00000000..38777d6d --- /dev/null +++ b/claude-code-sdk/src/test/java/io/github/markpollack/claude/agent/sdk/hooks/HookDecisionIT.java @@ -0,0 +1,198 @@ +/* + * Copyright 2026 Mark Pollack + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * 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 io.github.markpollack.claude.agent.sdk.hooks; + +import io.github.markpollack.claude.agent.sdk.ClaudeAsyncClient; +import io.github.markpollack.claude.agent.sdk.ClaudeClient; +import io.github.markpollack.claude.agent.sdk.ClaudeSyncClient; +import io.github.markpollack.claude.agent.sdk.config.PermissionMode; +import io.github.markpollack.claude.agent.sdk.test.ClaudeCliTestBase; +import io.github.markpollack.claude.agent.sdk.transport.CLIOptions; +import io.github.markpollack.claude.agent.sdk.types.AssistantMessage; +import io.github.markpollack.claude.agent.sdk.types.Message; +import io.github.markpollack.claude.agent.sdk.types.ToolResultBlock; +import io.github.markpollack.claude.agent.sdk.types.UserMessage; +import io.github.markpollack.claude.agent.sdk.types.control.HookOutput; +import io.github.markpollack.claude.agent.sdk.types.control.HookOutput.HookSpecificOutput; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Tag; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Timeout; +import org.junit.jupiter.api.io.TempDir; + +import java.nio.file.Path; +import java.time.Duration; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.stream.Collectors; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * Proves, against the real CLI, that the decisions a hook returns through + * {@link HookOutput#hookSpecificOutput()} take effect: a PreToolUse deny stops the tool, + * an {@code updatedInput} replaces its input, and {@code additionalContext} reaches the + * model. + * + *

+ * The sessions run in {@link PermissionMode#BYPASS_PERMISSIONS}, so the hook is the only + * thing that can stop a tool. In a prompting mode the CLI could deny the tool on its own + * and a broken hook would still look like it worked. + *

+ */ +@Timeout(value = 180, unit = TimeUnit.SECONDS) +@Tag("live") +class HookDecisionIT extends ClaudeCliTestBase { + + private static final String HAIKU_MODEL = CLIOptions.MODEL_HAIKU; + + private static final String DENY_REASON = "touch is blocked by the HookDecisionIT policy"; + + private static final Duration SESSION_TIMEOUT = Duration.ofMinutes(2); + + @TempDir + Path tempDir; + + @Test + @DisplayName("A PreToolUse deny stops the Bash command") + void preToolUseDenyBlocksBash() throws Exception { + Path marker = tempDir.resolve("denied.txt"); + AtomicInteger hookCalls = new AtomicInteger(); + + List messages = runSync(denyBash(hookCalls), touchPrompt(marker)); + + assertDenied(hookCalls, marker, messages); + } + + @Test + @DisplayName("Async client: a PreToolUse deny stops the Bash command") + void asyncPreToolUseDenyBlocksBash() { + Path marker = tempDir.resolve("denied-async.txt"); + AtomicInteger hookCalls = new AtomicInteger(); + + List messages = runAsync(denyBash(hookCalls), touchPrompt(marker)); + + assertDenied(hookCalls, marker, messages); + } + + @Test + @DisplayName("A PreToolUse updatedInput replaces the Bash command") + void preToolUseUpdatedInputReplacesCommand() throws Exception { + Path requested = tempDir.resolve("requested.txt"); + Path rewritten = tempDir.resolve("rewritten.txt"); + HookRegistry hooks = new HookRegistry(); + hooks.registerPreToolUse("Bash", + input -> HookOutput.builder() + .hookSpecificOutput(HookSpecificOutput.preToolUseModify(Map.of("command", "touch " + rewritten))) + .build()); + + runSync(hooks, touchPrompt(requested)); + + assertThat(rewritten).as("the rewritten command should have run").exists(); + assertThat(requested).as("the command the model asked for should have been replaced").doesNotExist(); + } + + @Test + @DisplayName("UserPromptSubmit additionalContext reaches the model") + void userPromptSubmitAdditionalContextReachesModel() throws Exception { + HookRegistry hooks = new HookRegistry(); + hooks.registerUserPromptSubmit(input -> HookOutput.builder() + .hookSpecificOutput(HookSpecificOutput.userPromptSubmit("The codeword for this session is PELICAN-7731.")) + .build()); + + List messages = runSync(hooks, + "What is the codeword for this session? Reply with just the codeword, or NONE if you were not given one."); + + String text = messages.stream() + .filter(AssistantMessage.class::isInstance) + .map(m -> ((AssistantMessage) m).text()) + .collect(Collectors.joining()); + assertThat(text).contains("PELICAN-7731"); + } + + private static HookRegistry denyBash(AtomicInteger hookCalls) { + HookRegistry hooks = new HookRegistry(); + hooks.registerPreToolUse("Bash", input -> { + hookCalls.incrementAndGet(); + return HookOutput.builder().hookSpecificOutput(HookSpecificOutput.preToolUseDeny(DENY_REASON)).build(); + }); + return hooks; + } + + private static void assertDenied(AtomicInteger hookCalls, Path marker, List messages) { + assertThat(hookCalls).as("the PreToolUse hook should have been called").hasPositiveValue(); + assertThat(marker).as("the denied command must not have run").doesNotExist(); + assertThat(toolResults(messages)).as("the CLI reports the deny as an errored tool result") + .anySatisfy(result -> { + assertThat(result.isError()).isTrue(); + assertThat(String.valueOf(result.content())).contains(DENY_REASON); + }); + } + + private List runSync(HookRegistry hooks, String prompt) { + List messages = new ArrayList<>(); + try (ClaudeSyncClient client = ClaudeClient.sync() + .workingDirectory(tempDir) + .claudePath(getClaudeCliPath()) + .model(HAIKU_MODEL) + .permissionMode(PermissionMode.BYPASS_PERMISSIONS) + .hookRegistry(hooks) + .timeout(SESSION_TIMEOUT) + .build()) { + client.connectAndReceive(prompt).forEach(messages::add); + } + return messages; + } + + private List runAsync(HookRegistry hooks, String prompt) { + ClaudeAsyncClient client = ClaudeClient.async() + .workingDirectory(tempDir) + .claudePath(getClaudeCliPath()) + .model(HAIKU_MODEL) + .permissionMode(PermissionMode.BYPASS_PERMISSIONS) + .hookRegistry(hooks) + .timeout(SESSION_TIMEOUT) + .build(); + try { + return client.connect(prompt).messages().collectList().block(SESSION_TIMEOUT); + } + finally { + client.close().block(Duration.ofSeconds(30)); + } + } + + private static String touchPrompt(Path file) { + return "Use the Bash tool to run exactly this command, once, and do not try anything else afterwards: touch " + + file; + } + + private static List toolResults(List messages) { + return messages.stream() + .filter(UserMessage.class::isInstance) + .map(message -> ((UserMessage) message).getContentAsBlocks()) + .filter(Objects::nonNull) + .flatMap(List::stream) + .filter(ToolResultBlock.class::isInstance) + .map(ToolResultBlock.class::cast) + .toList(); + } + +} diff --git a/claude-code-sdk/src/test/java/io/github/markpollack/claude/agent/sdk/hooks/HookRegistryTest.java b/claude-code-sdk/src/test/java/io/github/markpollack/claude/agent/sdk/hooks/HookRegistryTest.java index a68a91fc..baa17422 100644 --- a/claude-code-sdk/src/test/java/io/github/markpollack/claude/agent/sdk/hooks/HookRegistryTest.java +++ b/claude-code-sdk/src/test/java/io/github/markpollack/claude/agent/sdk/hooks/HookRegistryTest.java @@ -16,11 +16,14 @@ package io.github.markpollack.claude.agent.sdk.hooks; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Nested; import org.junit.jupiter.api.Test; import io.github.markpollack.claude.agent.sdk.types.control.ControlRequest; +import io.github.markpollack.claude.agent.sdk.types.control.ControlResponse; import io.github.markpollack.claude.agent.sdk.types.control.HookEvent; import io.github.markpollack.claude.agent.sdk.types.control.HookInput; import io.github.markpollack.claude.agent.sdk.types.control.HookOutput; @@ -373,6 +376,166 @@ void buildMultiEventConfig() { } + /** + * The response to a {@code hook_callback} must be the CLI's hook JSON output format: + * {@code hookSpecificOutput} nested with camelCase keys, unset fields absent. The CLI + * rejects {@code "continue": null} and ignores unknown keys, and either one silently + * drops the hook's decision. + */ + @Nested + @DisplayName("Hook Callback Response") + class HookCallbackResponseTests { + + private static final String DENY_JSON = """ + {"hookSpecificOutput": { + "hookEventName": "PreToolUse", + "permissionDecision": "deny", + "permissionDecisionReason": "branch changes are blocked"}} + """; + + private final ObjectMapper mapper = new ObjectMapper(); + + private final HookInput preToolUseInput = new HookInput.PreToolUseInput("PreToolUse", "sess_1", "/tmp/t.md", + "/home", null, "Bash", "tool_123", Map.of("command", "git switch -c probe")); + + @Test + @DisplayName("A PreToolUse deny is sent nested, camelCase, with no continue") + void preToolUseDeny() throws Exception { + String id = registry.registerPreToolUse("Bash", + input -> HookOutput.builder() + .hookSpecificOutput(HookOutput.HookSpecificOutput.preToolUseDeny("branch changes are blocked")) + .build()); + + JsonNode wire = wire(registry.handleCallback("req_1", id, preToolUseInput)); + + assertThat(wire.at("/response/subtype").asText()).isEqualTo("success"); + assertThat(wire.at("/response/request_id").asText()).isEqualTo("req_1"); + assertThat(wire.at("/response/response")).isEqualTo(mapper.readTree(DENY_JSON)); + } + + @Test + @DisplayName("A PreToolUse allow is sent with its decision and reason") + void preToolUseAllow() throws Exception { + String id = registry.registerPreToolUse("Bash", + input -> HookOutput.builder() + .continueExecution(true) + .hookSpecificOutput(HookOutput.HookSpecificOutput.preToolUseAllow("read-only command")) + .build()); + + assertThat(sent(id)).isEqualTo(mapper.readTree(""" + {"continue": true, + "hookSpecificOutput": { + "hookEventName": "PreToolUse", + "permissionDecision": "allow", + "permissionDecisionReason": "read-only command"}} + """)); + } + + @Test + @DisplayName("A PreToolUse updatedInput is sent nested as updatedInput") + void preToolUseModify() throws Exception { + String id = registry.registerPreToolUse("Bash", + input -> HookOutput.builder() + .hookSpecificOutput( + HookOutput.HookSpecificOutput.preToolUseModify(Map.of("command", "git status"))) + .build()); + + assertThat(sent(id)).isEqualTo(mapper.readTree(""" + {"hookSpecificOutput": { + "hookEventName": "PreToolUse", + "updatedInput": {"command": "git status"}}} + """)); + } + + @Test + @DisplayName("PostToolUse and UserPromptSubmit additionalContext are sent nested") + void additionalContext() throws Exception { + String post = registry.registerPostToolUse(input -> HookOutput.builder() + .hookSpecificOutput(HookOutput.HookSpecificOutput.postToolUse("the build is red")) + .build()); + String prompt = registry.registerUserPromptSubmit(input -> HookOutput.builder() + .hookSpecificOutput(HookOutput.HookSpecificOutput.userPromptSubmit("today is release day")) + .build()); + + assertThat(sent(post)).isEqualTo(mapper.readTree(""" + {"hookSpecificOutput": { + "hookEventName": "PostToolUse", + "additionalContext": "the build is red"}} + """)); + assertThat(sent(prompt)).isEqualTo(mapper.readTree(""" + {"hookSpecificOutput": { + "hookEventName": "UserPromptSubmit", + "additionalContext": "today is release day"}} + """)); + } + + @Test + @DisplayName("The top-level control fields keep their names") + void topLevelFields() throws Exception { + String id = registry.registerPreToolUse("Bash", + input -> HookOutput.builder() + .continueExecution(false) + .stopReason("policy stop") + .suppressOutput(true) + .systemMessage("stopped by policy") + .decision("block") + .reason("not allowed") + .build()); + + assertThat(sent(id)).isEqualTo(mapper.readTree(""" + {"continue": false, "suppressOutput": true, "stopReason": "policy stop", + "decision": "block", "systemMessage": "stopped by policy", "reason": "not allowed"} + """)); + } + + @Test + @DisplayName("A missing hookEventName is filled in from the registration") + void missingEventName() throws Exception { + String id = registry.registerPreToolUse("Bash", + input -> HookOutput.builder() + .hookSpecificOutput(HookOutput.HookSpecificOutput.builder() + .permissionDecision("deny") + .permissionDecisionReason("branch changes are blocked") + .build()) + .build()); + + assertThat(sent(id)).isEqualTo(mapper.readTree(DENY_JSON)); + } + + @Test + @DisplayName("A hookEventName for another event is sent as the hook returned it") + void mismatchedEventName() throws Exception { + String id = registry.registerPostToolUse(input -> HookOutput.builder() + .hookSpecificOutput(HookOutput.HookSpecificOutput.preToolUseDeny("wrong event")) + .build()); + + assertThat(sent(id).at("/hookSpecificOutput/hookEventName").asText()).isEqualTo("PreToolUse"); + } + + @Test + @DisplayName("An unknown callback ID is answered with an error") + void unknownCallback() throws Exception { + JsonNode wire = wire(registry.handleCallback("req_7", "hook_missing", preToolUseInput)); + + assertThat(wire.at("/response/subtype").asText()).isEqualTo("error"); + assertThat(wire.at("/response/request_id").asText()).isEqualTo("req_7"); + assertThat(wire.at("/response/error").asText()).contains("hook_missing"); + } + + /** + * The hook output {@code handleCallback} sends for the hook registered under + * {@code id}. + */ + private JsonNode sent(String id) throws Exception { + return wire(registry.handleCallback("req", id, preToolUseInput)).at("/response/response"); + } + + private JsonNode wire(ControlResponse response) throws Exception { + return mapper.readTree(mapper.writeValueAsString(response)); + } + + } + @Nested @DisplayName("Initialize Request Creation") class InitializeRequestTests {