summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-11-27 09:16:06 +0800
committerGitHub <noreply@github.com>2024-11-27 09:16:06 +0800
commit2003159d322a96165ab949051c004c456de5be11 (patch)
treedc3c7f8e1850267ba24d1a3bf9a76fdac4366b6a /src
parent6f60ad47535b75064a79331c22758d86e6e46d4f (diff)
downloadaichat-2003159d322a96165ab949051c004c456de5be11.tar.gz
feat: handle Ctrl+C during every spinner in REPL (#1014)
Diffstat (limited to 'src')
-rw-r--r--src/config/input.rs17
-rw-r--r--src/main.rs14
-rw-r--r--src/rag/mod.rs58
-rw-r--r--src/render/stream.rs4
-rw-r--r--src/repl/mod.rs23
-rw-r--r--src/utils/spinner.rs131
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;
}