add support for requesting claude to stop what it's doing if it's taking too long
This commit is contained in:
@@ -73,9 +73,9 @@ pub enum ClaudeError {
|
|||||||
/// Worker has shut down.
|
/// Worker has shut down.
|
||||||
#[error("worker has shut down")]
|
#[error("worker has shut down")]
|
||||||
WorkerGone,
|
WorkerGone,
|
||||||
/// Worker returned an error.
|
/// Stop was requested.
|
||||||
#[error("{0}")]
|
#[error("stop requested")]
|
||||||
WorkerError(String),
|
StopRequested,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Claude API client.
|
/// Claude API client.
|
||||||
@@ -84,6 +84,7 @@ pub enum ClaudeError {
|
|||||||
/// concurrent API calls.
|
/// concurrent API calls.
|
||||||
pub struct ClaudeApi {
|
pub struct ClaudeApi {
|
||||||
request_tx: mpsc::Sender<ClaudeRequest>,
|
request_tx: mpsc::Sender<ClaudeRequest>,
|
||||||
|
stop_tx: mpsc::Sender<()>,
|
||||||
#[allow(dead_code)]
|
#[allow(dead_code)]
|
||||||
worker_handle: tokio::task::JoinHandle<()>,
|
worker_handle: tokio::task::JoinHandle<()>,
|
||||||
}
|
}
|
||||||
@@ -187,8 +188,9 @@ impl ClaudeApi {
|
|||||||
tool_executor: Weak<dyn ToolExecutor>,
|
tool_executor: Weak<dyn ToolExecutor>,
|
||||||
) -> Result<Self, ClaudeError> {
|
) -> Result<Self, ClaudeError> {
|
||||||
let (request_tx, request_rx) = mpsc::channel(REQUEST_QUEUE_SIZE);
|
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 {
|
let worker_handle = tokio::spawn(async move {
|
||||||
worker.run().await;
|
worker.run().await;
|
||||||
@@ -197,6 +199,7 @@ impl ClaudeApi {
|
|||||||
|
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
request_tx,
|
request_tx,
|
||||||
|
stop_tx,
|
||||||
worker_handle,
|
worker_handle,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -217,9 +220,16 @@ impl ClaudeApi {
|
|||||||
.try_send(request)
|
.try_send(request)
|
||||||
.map_err(|_| ClaudeError::QueueFull)?;
|
.map_err(|_| ClaudeError::QueueFull)?;
|
||||||
|
|
||||||
result_rx
|
// result_rx.await has type Result<Result<String, ClaudeError>, RecvError>
|
||||||
.await
|
// The outer Result is for channel errors, the inner is the actual response
|
||||||
.map_err(|_| ClaudeError::WorkerGone)?
|
result_rx.await.map_err(|_| ClaudeError::WorkerGone)?
|
||||||
.map_err(ClaudeError::WorkerError)
|
}
|
||||||
|
|
||||||
|
/// 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(());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ use tracing::info;
|
|||||||
/// A request to be processed by the Claude worker.
|
/// A request to be processed by the Claude worker.
|
||||||
pub struct ClaudeRequest {
|
pub struct ClaudeRequest {
|
||||||
pub prompt: String,
|
pub prompt: String,
|
||||||
pub result_sender: oneshot::Sender<Result<String, String>>,
|
pub result_sender: oneshot::Sender<Result<String, ClaudeError>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Background worker that processes Claude API requests serially.
|
/// Background worker that processes Claude API requests serially.
|
||||||
@@ -22,6 +22,7 @@ pub struct ClaudeWorker {
|
|||||||
system_prompt: String,
|
system_prompt: String,
|
||||||
tool_executor: Weak<dyn ToolExecutor>,
|
tool_executor: Weak<dyn ToolExecutor>,
|
||||||
request_rx: mpsc::Receiver<ClaudeRequest>,
|
request_rx: mpsc::Receiver<ClaudeRequest>,
|
||||||
|
stop_rx: mpsc::Receiver<()>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ClaudeWorker {
|
impl ClaudeWorker {
|
||||||
@@ -32,6 +33,7 @@ impl ClaudeWorker {
|
|||||||
config: ClaudeConfig,
|
config: ClaudeConfig,
|
||||||
tool_executor: Weak<dyn ToolExecutor>,
|
tool_executor: Weak<dyn ToolExecutor>,
|
||||||
request_rx: mpsc::Receiver<ClaudeRequest>,
|
request_rx: mpsc::Receiver<ClaudeRequest>,
|
||||||
|
stop_rx: mpsc::Receiver<()>,
|
||||||
) -> Result<Self, ClaudeError> {
|
) -> Result<Self, ClaudeError> {
|
||||||
let api_key = std::fs::read_to_string(&config.api_key_file)
|
let api_key = std::fs::read_to_string(&config.api_key_file)
|
||||||
.map_err(ClaudeError::ApiKeyRead)?
|
.map_err(ClaudeError::ApiKeyRead)?
|
||||||
@@ -48,23 +50,56 @@ impl ClaudeWorker {
|
|||||||
system_prompt,
|
system_prompt,
|
||||||
tool_executor,
|
tool_executor,
|
||||||
request_rx,
|
request_rx,
|
||||||
|
stop_rx,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Run the worker loop, processing requests serially.
|
/// Run the worker loop, processing requests serially.
|
||||||
pub async fn run(mut self) {
|
pub async fn run(mut self) {
|
||||||
while let Some(request) = self.request_rx.recv().await {
|
loop {
|
||||||
let result = self
|
tokio::select! {
|
||||||
.handle_request(&request.prompt)
|
request = self.request_rx.recv() => {
|
||||||
.await
|
let Some(request) = request else {
|
||||||
.map_err(|e| e.to_string());
|
// Channel closed, exit
|
||||||
// Ignore send errors - the caller may have dropped the receiver
|
break;
|
||||||
let _ = request.result_sender.send(result);
|
};
|
||||||
|
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.
|
/// Handle a single request to the Claude API.
|
||||||
async fn handle_request(&self, prompt: &str) -> Result<String, ClaudeError> {
|
async fn handle_request(&mut self, prompt: &str) -> Result<String, ClaudeError> {
|
||||||
let max_iterations = self.config.claude_max_iterations;
|
let max_iterations = self.config.claude_max_iterations;
|
||||||
|
|
||||||
// Get tools from the executor if still alive
|
// Get tools from the executor if still alive
|
||||||
@@ -75,6 +110,9 @@ impl ClaudeWorker {
|
|||||||
info!("Claude request: {}", prompt);
|
info!("Claude request: {}", prompt);
|
||||||
|
|
||||||
for iteration in 0..max_iterations {
|
for iteration in 0..max_iterations {
|
||||||
|
// Check for stop before making API call
|
||||||
|
self.check_stop()?;
|
||||||
|
|
||||||
let request_body = MessagesRequest {
|
let request_body = MessagesRequest {
|
||||||
model: &self.config.claude_model,
|
model: &self.config.claude_model,
|
||||||
max_tokens: self.config.claude_max_tokens,
|
max_tokens: self.config.claude_max_tokens,
|
||||||
@@ -117,6 +155,9 @@ impl ClaudeWorker {
|
|||||||
// Execute each tool use and collect results
|
// Execute each tool use and collect results
|
||||||
for block in &response.content {
|
for block in &response.content {
|
||||||
if let ContentBlock::ToolUse { id, name, input } = block {
|
if let ContentBlock::ToolUse { id, name, input } = block {
|
||||||
|
// Check for stop before each tool use
|
||||||
|
self.check_stop()?;
|
||||||
|
|
||||||
info!("Claude tool use: {}({})", name, input);
|
info!("Claude tool use: {}({})", name, input);
|
||||||
let (result, is_error) = match executor.execute(name, input).await {
|
let (result, is_error) = match executor.execute(name, input).await {
|
||||||
Ok(result) => {
|
Ok(result) => {
|
||||||
|
|||||||
@@ -144,6 +144,9 @@ enum GatewayCommand {
|
|||||||
#[conf(repeat, pos)]
|
#[conf(repeat, pos)]
|
||||||
prompt: Vec<String>,
|
prompt: Vec<String>,
|
||||||
},
|
},
|
||||||
|
/// Stop current Claude request
|
||||||
|
#[conf(name = "cs", alias = "CS")]
|
||||||
|
ClaudeStop,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Parse a gateway command from a string (with or without leading /)
|
/// 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())),
|
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"))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user