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/utils/spinner.rs | 88 ++++++++++++++++++++++++++++++++++++++++++++++------ 1 file changed, 79 insertions(+), 9 deletions(-) (limited to 'src/utils') 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