summaryrefslogtreecommitdiffstats
path: root/src/main.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-21 06:52:22 +0800
committerGitHub <noreply@github.com>2024-04-21 06:52:22 +0800
commitb3162a72b50f88d78ca3d57c1f8b3421cecfe2b4 (patch)
treecee11116acce66ffcc7599dc13dcc9588d174620 /src/main.rs
parent8b75080973e6caa3128e9673440a57deb38fc0c9 (diff)
downloadaichat-b3162a72b50f88d78ca3d57c1f8b3421cecfe2b4.tar.gz
refactor: shell execution with spinner and select prompt (#425)
Diffstat (limited to 'src/main.rs')
-rw-r--r--src/main.rs53
1 files changed, 34 insertions, 19 deletions
diff --git a/src/main.rs b/src/main.rs
index 66492da..600381b 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -20,13 +20,14 @@ use crate::utils::{
use anyhow::{bail, Result};
use clap::Parser;
-use inquire::validator::Validation;
-use inquire::Text;
+use inquire::{Select, Text};
use is_terminal::IsTerminal;
use parking_lot::RwLock;
use std::io::{stderr, stdin, stdout, Read};
-use std::process;
-use std::sync::Arc;
+use std::sync::{mpsc, Arc};
+use std::time::Duration;
+use std::{process, thread};
+use utils::Spinner;
fn main() -> Result<()> {
let cli = Cli::parse();
@@ -153,7 +154,11 @@ fn start_interactive(config: &GlobalConfig) -> Result<()> {
fn execute(config: &GlobalConfig, mut input: Input) -> Result<()> {
let client = init_client(config)?;
config.read().maybe_print_send_tokens(&input);
- let mut eval_str = client.send_message(input.clone())?;
+ let (tx, rx) = mpsc::sync_channel::<()>(0);
+ thread::spawn(move || run_spinner(rx));
+ let ret = client.send_message(input.clone());
+ tx.send(())?;
+ let mut eval_str = ret?;
if let Ok(true) = CODE_BLOCK_RE.is_match(&eval_str) {
eval_str = extract_block(&eval_str);
}
@@ -166,27 +171,22 @@ fn execute(config: &GlobalConfig, mut input: Input) -> Result<()> {
return Ok(());
}
if stdout().is_terminal() {
- println!("{}", markdown_render.render(&eval_str).trim());
let mut explain = false;
loop {
- let answer = Text::new("[1]:execute [2]:explain [3]:revise [4]:cancel")
- .with_default("1")
- .with_validator(|input: &str| match matches!(input, "1" | "2" | "3" | "4") {
- true => Ok(Validation::Valid),
- false => Ok(Validation::Invalid(
- "Select a number between 1 and 4.".into(),
- )),
- })
- .prompt()?;
+ let answer = Select::new(
+ markdown_render.render(&eval_str).trim(),
+ vec!["✅ Execute", "📙 Explain", "🤔 Revise", "❌ Cancel"],
+ )
+ .prompt()?;
- match answer.as_str() {
- "1" => {
+ match answer {
+ "✅ Execute" => {
let code = run_command(&eval_str)?;
if code != 0 {
process::exit(code);
}
}
- "2" => {
+ "📙 Explain" => {
if !explain {
config.write().set_role(EXPLAIN_ROLE)?;
}
@@ -196,7 +196,7 @@ fn execute(config: &GlobalConfig, mut input: Input) -> Result<()> {
explain = true;
continue;
}
- "3" => {
+ "🤔 Revise" => {
let revision = Text::new("Enter your revision:").prompt()?;
let text = format!(
"[INST] {} [/INST]\n{eval_str}\n[INST] {revision} [/INST]\n",
@@ -246,3 +246,18 @@ fn create_input(
};
Ok(Some(input))
}
+
+fn run_spinner(rx: mpsc::Receiver<()>) -> Result<()> {
+ let mut writer = stdout();
+ let mut spinner = Spinner::new(" Generating");
+ loop {
+ spinner.step(&mut writer)?;
+ if let Ok(()) = rx.try_recv() {
+ spinner.stop(&mut writer)?;
+ break;
+ }
+ thread::sleep(Duration::from_millis(50))
+ }
+ spinner.stop(&mut writer)?;
+ Ok(())
+}