summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/common.rs14
-rw-r--r--src/main.rs33
-rw-r--r--src/repl/mod.rs2
-rw-r--r--src/utils/spinner.rs88
4 files changed, 116 insertions, 21 deletions
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<ToolResult>)> {
- 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<String>) -> 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<SpinnerEve
Ok(())
}
-pub async fn run_with_spinner<F, T>(task: F, spinner_message: &str) -> Result<T>
+pub async fn abortable_run_with_spinner<F, T>(
+ task: F,
+ message: &str,
+ abort_signal: AbortSignal,
+) -> Result<T>
where
F: Future<Output = Result<T>>,
{
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(())
+}