From 740f4060f1272063ec89f2b93ed334d3ed1cbfe4 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 29 Oct 2024 14:31:04 +0800 Subject: fix: unexpected Ctrl-C/Ctrl-D handling in non-stream REPL Chat (#957) --- src/client/common.rs | 14 +++++++-- src/main.rs | 33 ++++++++++++++------ src/repl/mod.rs | 2 +- src/utils/spinner.rs | 88 ++++++++++++++++++++++++++++++++++++++++++++++------ 4 files changed, 116 insertions(+), 21 deletions(-) (limited to 'src') diff --git a/src/client/common.rs b/src/client/common.rs index fe12593..1275d0b 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -403,9 +403,15 @@ pub async fn call_chat_completions( input: &Input, extract_code: bool, client: &dyn Client, + abort_signal: AbortSignal, ) -> Result<(String, Vec)> { - let task = client.chat_completions(input.clone()); - let ret = run_with_spinner(task, "Generating").await; + let ret = abortable_run_with_spinner( + client.chat_completions(input.clone()), + "Generating", + abort_signal, + ) + .await; + match ret { Ok(ret) => { let ChatCompletionsOutput { @@ -438,6 +444,10 @@ pub async fn call_chat_completions_streaming( render_stream(rx, client.global_config(), abort_signal.clone()), ); + if handler.abort().aborted() { + bail!("Aborted."); + } + render_ret?; let (text, tool_calls) = handler.take(); diff --git a/src/main.rs b/src/main.rs index ee95045..4f24f25 100644 --- a/src/main.rs +++ b/src/main.rs @@ -143,7 +143,7 @@ async fn run(config: GlobalConfig, cli: Cli, text: Option) -> Result<()> bail!("Unable to read the pipe for shell execution on MacOS") } let input = create_input(&config, text, &cli.file).await?; - shell_execute(&config, &SHELL, input).await?; + shell_execute(&config, &SHELL, input, abort_signal.clone()).await?; return Ok(()); } config.write().apply_prelude()?; @@ -173,7 +173,7 @@ async fn start_directive( let extract_code = !*IS_STDOUT_TERMINAL && code_mode; config.write().before_chat_completion(&input)?; let (output, tool_results) = if !input.stream() || extract_code { - call_chat_completions(&input, extract_code, client.as_ref()).await? + call_chat_completions(&input, extract_code, client.as_ref(), abort_signal.clone()).await? } else { call_chat_completions_streaming(&input, client.as_ref(), abort_signal.clone()).await? }; @@ -201,10 +201,21 @@ async fn start_interactive(config: &GlobalConfig) -> Result<()> { } #[async_recursion::async_recursion] -async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) -> Result<()> { +async fn shell_execute( + config: &GlobalConfig, + shell: &Shell, + mut input: Input, + abort_signal: AbortSignal, +) -> Result<()> { let client = input.create_client()?; config.write().before_chat_completion(&input)?; - let (mut eval_str, _) = call_chat_completions(&input, false, client.as_ref()).await?; + let ret = abortable_run_with_spinner( + client.chat_completions(input.clone()), + "Generating", + abort_signal.clone(), + ) + .await; + let mut eval_str = ret?.text; if let Ok(true) = CODE_BLOCK_RE.is_match(&eval_str) { eval_str = extract_block(&eval_str); } @@ -253,17 +264,21 @@ async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) - let revision = Text::new("Enter your revision:").prompt()?; let text = format!("{}\n{revision}", input.text()); input.set_text(text); - return shell_execute(config, shell, input).await; + return shell_execute(config, shell, input, abort_signal.clone()).await; } "d" => { let role = config.read().retrieve_role(EXPLAIN_SHELL_ROLE)?; let input = Input::from_str(config, &eval_str, Some(role)); - let abort_signal = create_abort_signal(); if input.stream() { - call_chat_completions_streaming(&input, client.as_ref(), abort_signal) - .await?; + call_chat_completions_streaming( + &input, + client.as_ref(), + abort_signal.clone(), + ) + .await?; } else { - call_chat_completions(&input, false, client.as_ref()).await?; + call_chat_completions(&input, false, client.as_ref(), abort_signal.clone()) + .await?; } println!(); continue; diff --git a/src/repl/mod.rs b/src/repl/mod.rs index b392eb4..c08afc3 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -660,7 +660,7 @@ async fn ask( let (output, tool_results) = if input.stream() { call_chat_completions_streaming(&input, client.as_ref(), abort_signal.clone()).await? } else { - call_chat_completions(&input, false, client.as_ref()).await? + call_chat_completions(&input, false, client.as_ref(), abort_signal.clone()).await? }; config .write() diff --git a/src/utils/spinner.rs b/src/utils/spinner.rs index 53969f4..2fd21dd 100644 --- a/src/utils/spinner.rs +++ b/src/utils/spinner.rs @@ -1,13 +1,19 @@ -use super::IS_STDOUT_TERMINAL; +use super::{poll_abort_signal, wait_abort_signal, AbortSignal, IS_STDOUT_TERMINAL}; -use anyhow::Result; -use crossterm::{cursor, queue, style, terminal}; +use anyhow::{bail, Result}; +use crossterm::{ + cursor, queue, style, + terminal::{self, disable_raw_mode, enable_raw_mode}, +}; use std::{ future::Future, io::{stdout, Write}, time::Duration, }; -use tokio::{sync::mpsc, time::interval}; +use tokio::{ + sync::{mpsc, oneshot}, + time::interval, +}; pub struct SpinnerInner { index: usize, @@ -127,16 +133,80 @@ async fn run_spinner(message: String, mut rx: mpsc::UnboundedReceiver(task: F, spinner_message: &str) -> Result +pub async fn abortable_run_with_spinner( + task: F, + message: &str, + abort_signal: AbortSignal, +) -> Result where F: Future>, { if *IS_STDOUT_TERMINAL { - let spinner = create_spinner(spinner_message).await; - let ret = task.await; - spinner.stop(); - ret + let (done_tx, done_rx) = oneshot::channel(); + let run_task = async { + tokio::select! { + ret = task => { + let _ = done_tx.send(()); + ret + } + _ = wait_abort_signal(&abort_signal) => { + let _ = done_tx.send(()); + bail!("Aborted."); + }, + } + }; + let (task_ret, spinner_ret) = tokio::join!( + run_task, + run_abortable_spinner(message, abort_signal.clone(), done_rx) + ); + spinner_ret?; + task_ret } else { task.await } } + +async fn run_abortable_spinner( + message: &str, + abort_signal: AbortSignal, + done_rx: oneshot::Receiver<()>, +) -> Result<()> { + enable_raw_mode()?; + + let ret = run_abortable_spinner_inner(message, abort_signal, done_rx).await; + + disable_raw_mode()?; + ret +} + +async fn run_abortable_spinner_inner( + message: &str, + abort_signal: AbortSignal, + mut done_rx: oneshot::Receiver<()>, +) -> Result<()> { + let message = format!(" {message}"); + let mut spinner = SpinnerInner::new(&message); + loop { + if abort_signal.aborted() { + break; + } + + tokio::time::sleep(Duration::from_millis(25)).await; + + match done_rx.try_recv() { + Ok(_) | Err(oneshot::error::TryRecvError::Closed) => { + break; + } + _ => {} + } + + if poll_abort_signal(&abort_signal)? { + break; + } + + spinner.step()?; + } + + spinner.clear_message()?; + Ok(()) +} -- cgit v1.2.3