From 0862e3ce384389e7bb0b036830f5c503ebd812a1 Mon Sep 17 00:00:00 2001 From: Chris Beck Date: Sun, 7 Dec 2025 20:57:35 -0700 Subject: [PATCH] add support for requesting claude to stop what it's doing if it's taking too long --- signal-gateway/src/claude/mod.rs | 26 +++++++++---- signal-gateway/src/claude/worker.rs | 59 ++++++++++++++++++++++++----- signal-gateway/src/gateway/mod.rs | 12 ++++++ 3 files changed, 80 insertions(+), 17 deletions(-) diff --git a/signal-gateway/src/claude/mod.rs b/signal-gateway/src/claude/mod.rs index af4dfe3..1bf5bf1 100644 --- a/signal-gateway/src/claude/mod.rs +++ b/signal-gateway/src/claude/mod.rs @@ -73,9 +73,9 @@ pub enum ClaudeError { /// Worker has shut down. #[error("worker has shut down")] WorkerGone, - /// Worker returned an error. - #[error("{0}")] - WorkerError(String), + /// Stop was requested. + #[error("stop requested")] + StopRequested, } /// Claude API client. @@ -84,6 +84,7 @@ pub enum ClaudeError { /// concurrent API calls. pub struct ClaudeApi { request_tx: mpsc::Sender, + stop_tx: mpsc::Sender<()>, #[allow(dead_code)] worker_handle: tokio::task::JoinHandle<()>, } @@ -187,8 +188,9 @@ impl ClaudeApi { tool_executor: Weak, ) -> Result { let (request_tx, request_rx) = mpsc::channel(REQUEST_QUEUE_SIZE); + let (stop_tx, stop_rx) = mpsc::channel(REQUEST_QUEUE_SIZE); - let worker = ClaudeWorker::new(config, tool_executor, request_rx)?; + let worker = ClaudeWorker::new(config, tool_executor, request_rx, stop_rx)?; let worker_handle = tokio::spawn(async move { worker.run().await; @@ -197,6 +199,7 @@ impl ClaudeApi { Ok(Self { request_tx, + stop_tx, worker_handle, }) } @@ -217,9 +220,16 @@ impl ClaudeApi { .try_send(request) .map_err(|_| ClaudeError::QueueFull)?; - result_rx - .await - .map_err(|_| ClaudeError::WorkerGone)? - .map_err(ClaudeError::WorkerError) + // result_rx.await has type Result, RecvError> + // The outer Result is for channel errors, the inner is the actual response + result_rx.await.map_err(|_| ClaudeError::WorkerGone)? + } + + /// Request the worker to stop processing. + /// + /// This will cause the current request (if any) to be interrupted at the + /// next opportunity, and all pending requests to receive `StopRequested` errors. + pub fn request_stop(&self) { + let _ = self.stop_tx.try_send(()); } } diff --git a/signal-gateway/src/claude/worker.rs b/signal-gateway/src/claude/worker.rs index a03aa4a..2e4c461 100644 --- a/signal-gateway/src/claude/worker.rs +++ b/signal-gateway/src/claude/worker.rs @@ -11,7 +11,7 @@ use tracing::info; /// A request to be processed by the Claude worker. pub struct ClaudeRequest { pub prompt: String, - pub result_sender: oneshot::Sender>, + pub result_sender: oneshot::Sender>, } /// Background worker that processes Claude API requests serially. @@ -22,6 +22,7 @@ pub struct ClaudeWorker { system_prompt: String, tool_executor: Weak, request_rx: mpsc::Receiver, + stop_rx: mpsc::Receiver<()>, } impl ClaudeWorker { @@ -32,6 +33,7 @@ impl ClaudeWorker { config: ClaudeConfig, tool_executor: Weak, request_rx: mpsc::Receiver, + stop_rx: mpsc::Receiver<()>, ) -> Result { let api_key = std::fs::read_to_string(&config.api_key_file) .map_err(ClaudeError::ApiKeyRead)? @@ -48,23 +50,56 @@ impl ClaudeWorker { system_prompt, tool_executor, request_rx, + stop_rx, }) } /// Run the worker loop, processing requests serially. pub async fn run(mut self) { - while let Some(request) = self.request_rx.recv().await { - let result = self - .handle_request(&request.prompt) - .await - .map_err(|e| e.to_string()); - // Ignore send errors - the caller may have dropped the receiver - let _ = request.result_sender.send(result); + loop { + tokio::select! { + request = self.request_rx.recv() => { + let Some(request) = request else { + // Channel closed, exit + break; + }; + let result = self.handle_request(&request.prompt).await; + // If handle_request was interrupted by stop request, go on to drain the queues + if matches!(result, Err(ClaudeError::StopRequested)) { + self.handle_stop(); + } + // Ignore send errors - the caller may have dropped the receiver + let _ = request.result_sender.send(result); + } + _ = self.stop_rx.recv() => { + self.handle_stop(); + } + } + } + } + + /// Handle a stop request by draining queues and sending errors to pending requests. + fn handle_stop(&mut self) { + // Drain the stop_rx queue + while self.stop_rx.try_recv().is_ok() {} + + // Drain the request_rx queue and send StopRequested to each + while let Ok(request) = self.request_rx.try_recv() { + let _ = request.result_sender.send(Err(ClaudeError::StopRequested)); + } + } + + /// Check if stop has been requested. + fn check_stop(&mut self) -> Result<(), ClaudeError> { + match self.stop_rx.try_recv() { + Ok(()) => Err(ClaudeError::StopRequested), + Err(mpsc::error::TryRecvError::Empty) => Ok(()), + Err(mpsc::error::TryRecvError::Disconnected) => Err(ClaudeError::StopRequested), } } /// Handle a single request to the Claude API. - async fn handle_request(&self, prompt: &str) -> Result { + async fn handle_request(&mut self, prompt: &str) -> Result { let max_iterations = self.config.claude_max_iterations; // Get tools from the executor if still alive @@ -75,6 +110,9 @@ impl ClaudeWorker { info!("Claude request: {}", prompt); for iteration in 0..max_iterations { + // Check for stop before making API call + self.check_stop()?; + let request_body = MessagesRequest { model: &self.config.claude_model, max_tokens: self.config.claude_max_tokens, @@ -117,6 +155,9 @@ impl ClaudeWorker { // Execute each tool use and collect results for block in &response.content { if let ContentBlock::ToolUse { id, name, input } = block { + // Check for stop before each tool use + self.check_stop()?; + info!("Claude tool use: {}({})", name, input); let (result, is_error) = match executor.execute(name, input).await { Ok(result) => { diff --git a/signal-gateway/src/gateway/mod.rs b/signal-gateway/src/gateway/mod.rs index fe9229e..e0243c2 100644 --- a/signal-gateway/src/gateway/mod.rs +++ b/signal-gateway/src/gateway/mod.rs @@ -144,6 +144,9 @@ enum GatewayCommand { #[conf(repeat, pos)] prompt: Vec, }, + /// Stop current Claude request + #[conf(name = "cs", alias = "CS")] + ClaudeStop, } /// Parse a gateway command from a string (with or without leading /) @@ -700,6 +703,15 @@ impl Gateway { Err(err) => Err((500, err.to_string().into())), } } + GatewayCommand::ClaudeStop => { + let claude = self + .claude + .get() + .ok_or_else(|| (501u16, "claude was not configured".into()))?; + + claude.request_stop(); + Ok(AdminMessageResponse::new("stop requested")) + } } }