Skip to content
Open
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
Original file line number Diff line number Diff line change
Expand Up @@ -335,22 +335,28 @@ default Flux<String> 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<Void> interrupt();

/**
* 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<Void> 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<Void> setModel(String model);

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand All @@ -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;
Expand Down Expand Up @@ -151,12 +150,12 @@ public class DefaultClaudeAsyncClient implements ClaudeAsyncClient {
*/
private volatile Sinks.Many<ParsedMessage> 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<String, MonoSink<Map<String, Object>>> pendingResponses = new ConcurrentHashMap<>();
private final ConcurrentHashMap<String, Sinks.One<Map<String, Object>>> pendingResponses = new ConcurrentHashMap<>();

// Cross-turn message handlers (thread-safe for concurrent registration)
private final List<Consumer<Message>> messageHandlers = new CopyOnWriteArrayList<>();
Expand Down Expand Up @@ -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());
}

Expand Down Expand Up @@ -409,58 +411,34 @@ public Flux<Message> receiveResponse() {

@Override
public Mono<Void> interrupt() {
return Mono.<Void>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<Void> setPermissionMode(String mode) {
return Mono.<Void>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<Void> setModel(String model) {
return Mono.<Void>create(sink -> {
Map<String, Object> 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<Void> awaitControlRequest(Map<String, Object> request) {
return Mono.defer(() -> {
if (!connected.get() || closed.get()) {
sink.error(new IllegalStateException("Client is not connected"));
return;
}
try {
Map<String, Object> request = new LinkedHashMap<>();
request.put("subtype", "set_model");
request.put("model", model);
sendControlRequest(request);
currentModel.set(model);
sink.success();
return Mono.<Map<String, Object>>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
Expand Down Expand Up @@ -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<String, Object> inputMap = hookCallback.input();

HookInput input = objectMapper.convertValue(inputMap, HookInput.class);
HookOutput output = hookRegistry.executeHook(callbackId, input);

Map<String, Object> 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);
Expand Down Expand Up @@ -759,7 +714,7 @@ private void handleControlResponse(ControlResponse response) {

logger.debug("Handling control response: requestId={}, subtype={}", requestId, response.response().subtype());

MonoSink<Map<String, Object>> sink = pendingResponses.remove(requestId);
Sinks.One<Map<String, Object>> sink = pendingResponses.remove(requestId);
if (sink == null) {
logger.warn("Unexpected response for unknown request id {}", requestId);
return;
Expand All @@ -774,35 +729,58 @@ private void handleControlResponse(ControlResponse response) {
Map<String, Object> typedMap = (Map<String, Object>) 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<String, Object> request) throws ClaudeSDKException {
try {
String requestId = sessionPrefix + "_" + requestCounter.incrementAndGet();
/**
* Sends a control request and returns the CLI's reply.
*
* <p>
* 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.
* </p>
* @throws TransportException if the request cannot be written to the CLI
*/
private Mono<Map<String, Object>> sendControlRequest(Map<String, Object> request) throws ClaudeSDKException {
String requestId = sessionPrefix + "_" + requestCounter.incrementAndGet();

Map<String, Object> 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<String, Object> 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<Map<String, Object>> 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() {
Expand All @@ -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();
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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<String, Object> inputMap = hookCallback.input();

HookInput input = objectMapper.convertValue(inputMap, HookInput.class);
HookOutput output = hookRegistry.executeHook(callbackId, input);

Map<String, Object> 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);
Expand Down
Loading