diff options
| author | sigoden <sigoden@gmail.com> | 2024-11-27 09:16:06 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-11-27 09:16:06 +0800 |
| commit | 2003159d322a96165ab949051c004c456de5be11 (patch) | |
| tree | dc3c7f8e1850267ba24d1a3bf9a76fdac4366b6a /src | |
| parent | 6f60ad47535b75064a79331c22758d86e6e46d4f (diff) | |
| download | aichat-2003159d322a96165ab949051c004c456de5be11.tar.gz | |
feat: handle Ctrl+C during every spinner in REPL (#1014)
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/input.rs | 17 | ||||
| -rw-r--r-- | src/main.rs | 14 | ||||
| -rw-r--r-- | src/rag/mod.rs | 58 | ||||
| -rw-r--r-- | src/render/stream.rs | 4 | ||||
| -rw-r--r-- | src/repl/mod.rs | 23 | ||||
| -rw-r--r-- | src/utils/spinner.rs | 131 |
6 files changed, 139 insertions, 108 deletions
diff --git a/src/config/input.rs b/src/config/input.rs index 5db172f..c527cfb 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -63,7 +63,6 @@ impl Input { paths: Vec<String>, role: Option<Role>, ) -> Result<Self> { - let spinner = create_spinner("Loading files").await; let mut raw_paths = vec![]; let mut local_paths = vec![]; let mut remote_urls = vec![]; @@ -82,7 +81,6 @@ impl Input { } } let ret = load_documents(config, local_paths, remote_urls).await; - spinner.stop(); let (files, medias, data_urls) = ret.context("Failed to load files")?; let mut texts = vec![]; if !raw_text.is_empty() { @@ -112,6 +110,21 @@ impl Input { }) } + pub async fn from_files_with_spinner( + config: &GlobalConfig, + raw_text: &str, + paths: Vec<String>, + role: Option<Role>, + abort_signal: AbortSignal, + ) -> Result<Self> { + abortable_run_with_spinner( + Input::from_files(config, raw_text, paths, role), + "Loading files", + abort_signal, + ) + .await + } + pub fn is_empty(&self) -> bool { self.text.is_empty() && self.medias.is_empty() } diff --git a/src/main.rs b/src/main.rs index 8ca1e31..c41b1e7 100644 --- a/src/main.rs +++ b/src/main.rs @@ -146,14 +146,14 @@ async fn run(config: GlobalConfig, cli: Cli, text: Option<String>) -> Result<()> if cfg!(target_os = "macos") && !stdin().is_terminal() { bail!("Unable to read the pipe for shell execution on MacOS") } - let input = create_input(&config, text, &cli.file).await?; + let input = create_input(&config, text, &cli.file, abort_signal.clone()).await?; shell_execute(&config, &SHELL, input, abort_signal.clone()).await?; return Ok(()); } config.write().apply_prelude()?; match is_repl { false => { - let mut input = create_input(&config, text, &cli.file).await?; + let mut input = create_input(&config, text, &cli.file, abort_signal.clone()).await?; input.use_embeddings(abort_signal.clone()).await?; start_directive(&config, input, cli.code, abort_signal).await } @@ -320,11 +320,19 @@ async fn create_input( config: &GlobalConfig, text: Option<String>, file: &[String], + abort_signal: AbortSignal, ) -> Result<Input> { let input = if file.is_empty() { Input::from_str(config, &text.unwrap_or_default(), None) } else { - Input::from_files(config, &text.unwrap_or_default(), file.to_vec(), None).await? + Input::from_files_with_spinner( + config, + &text.unwrap_or_default(), + file.to_vec(), + None, + abort_signal, + ) + .await? }; if input.is_empty() { bail!("No input"); diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 82da7c8..93f59a2 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -88,17 +88,13 @@ impl Rag { paths = add_documents()?; }; let loaders = config.read().document_loaders.clone(); - let spinner = create_spinner("Starting").await; - tokio::select! { - ret = rag.sync_documents(loaders, &paths, Some(spinner.clone())) => { - spinner.stop(); - ret?; - } - _ = wait_abort_signal(&abort_signal) => { - spinner.stop(); - bail!("Aborted!") - }, - }; + let (spinner, spinner_rx) = Spinner::create(""); + abortable_run_with_spinner_rx( + rag.sync_documents(loaders, &paths, Some(spinner)), + spinner_rx, + abort_signal, + ) + .await?; if rag.save()? { println!("✓ Saved RAG to '{}'.", save_path.display()); } @@ -143,17 +139,13 @@ impl Rag { T: AsRef<str>, { let loaders = config.read().document_loaders.clone(); - let spinner = create_spinner("Starting").await; - tokio::select! { - ret = self.sync_documents(loaders, document_paths, Some(spinner.clone())) => { - spinner.stop(); - ret?; - } - _ = wait_abort_signal(&abort_signal) => { - spinner.stop(); - bail!("Aborted!") - }, - }; + let (spinner, spinner_rx) = Spinner::create(""); + abortable_run_with_spinner_rx( + self.sync_documents(loaders, document_paths, Some(spinner)), + spinner_rx, + abort_signal, + ) + .await?; if self.save()? { println!("✓ Saved rag to '{}'.", self.path); } @@ -311,16 +303,18 @@ impl Rag { rerank_model: Option<&str>, abort_signal: AbortSignal, ) -> Result<(String, Vec<DocumentId>)> { - let spinner = create_spinner("Searching").await; - let ret = tokio::select! { - ret = self.hybird_search(text, top_k, min_score_vector_search, min_score_keyword_search, rerank_model) => { - ret - } - _ = wait_abort_signal(&abort_signal) => { - bail!("Aborted!") - }, - }; - spinner.stop(); + let ret = abortable_run_with_spinner( + self.hybird_search( + text, + top_k, + min_score_vector_search, + min_score_keyword_search, + rerank_model, + ), + "Searching", + abort_signal, + ) + .await; let (ids, documents): (Vec<_>, Vec<_>) = ret?.into_iter().unzip(); let embeddings = documents.join("\n\n"); Ok((embeddings, ids)) diff --git a/src/render/stream.rs b/src/render/stream.rs index 07a1a18..6ff6ea7 100644 --- a/src/render/stream.rs +++ b/src/render/stream.rs @@ -1,6 +1,6 @@ use super::{MarkdownRender, SseEvent}; -use crate::utils::{create_spinner, poll_abort_signal, AbortSignal}; +use crate::utils::{poll_abort_signal, spawn_spinner, AbortSignal}; use anyhow::Result; use crossterm::{ @@ -66,7 +66,7 @@ async fn markdown_stream_inner( let columns = terminal::size()?.0; - let mut spinner = Some(create_spinner("Generating").await); + let mut spinner = Some(spawn_spinner("Generating")); 'outer: loop { if abort_signal.aborted() { diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 98d166a..4519e23 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -10,7 +10,9 @@ use crate::client::{call_chat_completions, call_chat_completions_streaming}; use crate::config::{AssertState, Config, GlobalConfig, Input, StateFlags}; use crate::function::need_send_tool_results; use crate::render::render_error; -use crate::utils::{create_abort_signal, create_spinner, set_text, temp_file, AbortSignal}; +use crate::utils::{ + abortable_run_with_spinner, create_abort_signal, set_text, temp_file, AbortSignal, +}; use anyhow::{bail, Context, Result}; use fancy_regex::Regex; @@ -362,10 +364,12 @@ impl Repl { }, ".compress" => match args { Some("session") => { - let spinner = create_spinner("Compressing").await; - let ret = Config::compress_session(&self.config).await; - spinner.stop(); - ret?; + abortable_run_with_spinner( + Config::compress_session(&self.config), + "Compressing", + self.abort_signal.clone(), + ) + .await?; println!("✓ Successfully compressed the session."); } _ => { @@ -401,7 +405,14 @@ impl Repl { Some(args) => { let (files, text) = split_files_text(args); let files = shell_words::split(files).with_context(|| "Invalid args")?; - let input = Input::from_files(&self.config, text, files, None).await?; + let input = Input::from_files_with_spinner( + &self.config, + text, + files, + None, + self.abort_signal.clone(), + ) + .await?; ask(&self.config, self.abort_signal.clone(), input, true).await?; } None => println!("Usage: .file <files>... [-- <text>...]"), diff --git a/src/utils/spinner.rs b/src/utils/spinner.rs index 2fd21dd..dab9af1 100644 --- a/src/utils/spinner.rs +++ b/src/utils/spinner.rs @@ -1,20 +1,21 @@ use super::{poll_abort_signal, wait_abort_signal, AbortSignal, IS_STDOUT_TERMINAL}; use anyhow::{bail, Result}; -use crossterm::{ - cursor, queue, style, - terminal::{self, disable_raw_mode, enable_raw_mode}, -}; +use crossterm::{cursor, queue, style, terminal}; use std::{ future::Future, io::{stdout, Write}, time::Duration, }; use tokio::{ - sync::{mpsc, oneshot}, + sync::{ + mpsc::{self, UnboundedReceiver}, + oneshot, + }, time::interval, }; +#[derive(Debug, Default)] pub struct SpinnerInner { index: usize, message: String, @@ -23,13 +24,6 @@ pub struct SpinnerInner { impl SpinnerInner { const DATA: [&'static str; 10] = ["⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏"]; - fn new(message: &str) -> Self { - SpinnerInner { - index: 0, - message: message.to_string(), - } - } - fn step(&mut self) -> Result<()> { if !*IS_STDOUT_TERMINAL || self.message.is_empty() { return Ok(()); @@ -75,13 +69,14 @@ impl SpinnerInner { #[derive(Clone)] pub struct Spinner(mpsc::UnboundedSender<SpinnerEvent>); -impl Drop for Spinner { - fn drop(&mut self) { - self.stop(); +impl Spinner { + pub fn create(message: &str) -> (Self, UnboundedReceiver<SpinnerEvent>) { + let (tx, spinner_rx) = mpsc::unbounded_channel(); + let spinner = Spinner(tx); + let _ = spinner.set_message(message.to_string()); + (spinner, spinner_rx) } -} -impl Spinner { pub fn set_message(&self, message: String) -> Result<()> { self.0.send(SpinnerEvent::SetMessage(message))?; std::thread::sleep(Duration::from_millis(10)); @@ -94,43 +89,40 @@ impl Spinner { } } -enum SpinnerEvent { +pub enum SpinnerEvent { SetMessage(String), Stop, } -pub async fn create_spinner(message: &str) -> Spinner { - let message = format!(" {message}"); - let (tx, rx) = mpsc::unbounded_channel(); - tokio::spawn(run_spinner(message, rx)); - Spinner(tx) -} - -async fn run_spinner(message: String, mut rx: mpsc::UnboundedReceiver<SpinnerEvent>) -> Result<()> { - let mut spinner = SpinnerInner::new(&message); - let mut interval = interval(Duration::from_millis(50)); - loop { - tokio::select! { - _ = interval.tick() => { - let _ = spinner.step(); - } - evt = rx.recv() => { - if let Some(evt) = evt { - match evt { - SpinnerEvent::SetMessage(message) => { - spinner.set_message(message)?; - } - SpinnerEvent::Stop => { - spinner.clear_message()?; - break; +pub fn spawn_spinner(message: &str) -> Spinner { + let (spinner, mut spinner_rx) = Spinner::create(message); + tokio::spawn(async move { + let mut spinner = SpinnerInner::default(); + let mut interval = interval(Duration::from_millis(50)); + loop { + tokio::select! { + evt = spinner_rx.recv() => { + if let Some(evt) = evt { + match evt { + SpinnerEvent::SetMessage(message) => { + spinner.set_message(message)?; + } + SpinnerEvent::Stop => { + spinner.clear_message()?; + break; + } } - } + } + } + _ = interval.tick() => { + let _ = spinner.step(); } } } - } - Ok(()) + Ok::<(), anyhow::Error>(()) + }); + spinner } pub async fn abortable_run_with_spinner<F, T>( @@ -141,6 +133,18 @@ pub async fn abortable_run_with_spinner<F, T>( where F: Future<Output = Result<T>>, { + let (_, spinner_rx) = Spinner::create(message); + abortable_run_with_spinner_rx(task, spinner_rx, abort_signal).await +} + +pub async fn abortable_run_with_spinner_rx<F, T>( + task: F, + spinner_rx: UnboundedReceiver<SpinnerEvent>, + abort_signal: AbortSignal, +) -> Result<T> +where + F: Future<Output = Result<T>>, +{ if *IS_STDOUT_TERMINAL { let (done_tx, done_rx) = oneshot::channel(); let run_task = async { @@ -149,6 +153,11 @@ where let _ = done_tx.send(()); ret } + _ = tokio::signal::ctrl_c() => { + abort_signal.set_ctrlc(); + let _ = done_tx.send(()); + bail!("Aborted!") + }, _ = wait_abort_signal(&abort_signal) => { let _ = done_tx.send(()); bail!("Aborted."); @@ -157,7 +166,7 @@ where }; let (task_ret, spinner_ret) = tokio::join!( run_task, - run_abortable_spinner(message, abort_signal.clone(), done_rx) + run_abortable_spinner(spinner_rx, done_rx, abort_signal.clone()) ); spinner_ret?; task_ret @@ -167,25 +176,11 @@ where } 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 spinner_rx: UnboundedReceiver<SpinnerEvent>, mut done_rx: oneshot::Receiver<()>, + abort_signal: AbortSignal, ) -> Result<()> { - let message = format!(" {message}"); - let mut spinner = SpinnerInner::new(&message); + let mut spinner = SpinnerInner::default(); loop { if abort_signal.aborted() { break; @@ -200,6 +195,16 @@ async fn run_abortable_spinner_inner( _ => {} } + match spinner_rx.try_recv() { + Ok(SpinnerEvent::SetMessage(message)) => { + spinner.set_message(message)?; + } + Ok(SpinnerEvent::Stop) => { + spinner.clear_message()?; + } + Err(_) => {} + } + if poll_abort_signal(&abort_signal)? { break; } |
