summaryrefslogtreecommitdiffstats
path: root/src/repl
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-03-07 11:51:52 +0800
committerGitHub <noreply@github.com>2023-03-07 11:51:52 +0800
commit11dc4d104b21b75d63e10e42ea4ee767c7bd5fdf (patch)
tree338a2a1d9eed095762fd61ee41437aa2b4e3f991 /src/repl
parent1640456049cda9999bb27501af30e05e46b0360d (diff)
downloadaichat-11dc4d104b21b75d63e10e42ea4ee767c7bd5fdf.tar.gz
refactor: optimize ctrl+c/ctrl+d abort handling (#27)
Diffstat (limited to 'src/repl')
-rw-r--r--src/repl/abort.rs51
-rw-r--r--src/repl/handler.rs37
-rw-r--r--src/repl/mod.rs29
3 files changed, 87 insertions, 30 deletions
diff --git a/src/repl/abort.rs b/src/repl/abort.rs
new file mode 100644
index 0000000..f76abb6
--- /dev/null
+++ b/src/repl/abort.rs
@@ -0,0 +1,51 @@
+use std::sync::{
+ atomic::{AtomicBool, Ordering},
+ Arc,
+};
+
+pub type SharedAbortSignal = Arc<AbortSignal>;
+
+pub struct AbortSignal {
+ ctrlc: AtomicBool,
+ ctrld: AtomicBool,
+}
+
+impl AbortSignal {
+ pub fn new() -> SharedAbortSignal {
+ Arc::new(Self {
+ ctrlc: AtomicBool::new(false),
+ ctrld: AtomicBool::new(false),
+ })
+ }
+
+ pub fn aborted(&self) -> bool {
+ if self.aborted_ctrlc() {
+ return true;
+ }
+ if self.aborted_ctrld() {
+ return true;
+ }
+ false
+ }
+
+ pub fn aborted_ctrlc(&self) -> bool {
+ self.ctrlc.load(Ordering::SeqCst)
+ }
+
+ pub fn aborted_ctrld(&self) -> bool {
+ self.ctrld.load(Ordering::SeqCst)
+ }
+
+ pub fn reset(&self) {
+ self.ctrlc.store(false, Ordering::SeqCst);
+ self.ctrld.store(false, Ordering::SeqCst);
+ }
+
+ pub fn set_ctrlc(&self) {
+ self.ctrlc.store(true, Ordering::SeqCst);
+ }
+
+ pub fn set_ctrld(&self) {
+ self.ctrld.store(true, Ordering::SeqCst);
+ }
+}
diff --git a/src/repl/handler.rs b/src/repl/handler.rs
index 809c202..f2c34e2 100644
--- a/src/repl/handler.rs
+++ b/src/repl/handler.rs
@@ -8,10 +8,10 @@ use crossbeam::channel::{unbounded, Sender};
use crossbeam::sync::WaitGroup;
use std::cell::RefCell;
use std::fs::File;
-use std::sync::atomic::AtomicBool;
-use std::sync::Arc;
use std::thread::spawn;
+use super::abort::SharedAbortSignal;
+
pub enum ReplCmd {
Submit(String),
SetRole(String),
@@ -25,7 +25,7 @@ pub struct ReplCmdHandler {
client: ChatGptClient,
config: SharedConfig,
state: RefCell<ReplCmdHandlerState>,
- ctrlc: Arc<AtomicBool>,
+ abort: SharedAbortSignal,
}
pub struct ReplCmdHandlerState {
@@ -34,9 +34,12 @@ pub struct ReplCmdHandlerState {
}
impl ReplCmdHandler {
- pub fn init(client: ChatGptClient, config: SharedConfig) -> Result<Self> {
+ pub fn init(
+ client: ChatGptClient,
+ config: SharedConfig,
+ abort: SharedAbortSignal,
+ ) -> Result<Self> {
let save_file = config.as_ref().borrow().open_message_file()?;
- let ctrlc = Arc::new(AtomicBool::new(false));
let state = RefCell::new(ReplCmdHandlerState {
save_file,
reply: String::new(),
@@ -45,7 +48,7 @@ impl ReplCmdHandler {
client,
config,
state,
- ctrlc,
+ abort,
})
}
@@ -61,15 +64,15 @@ impl ReplCmdHandler {
let highlight = self.config.borrow().highlight;
let mut stream_handler = if highlight {
let (tx, rx) = unbounded();
- let ctrlc = self.ctrlc.clone();
+ let abort = self.abort.clone();
let wg = wg.clone();
spawn(move || {
- let _ = render_stream(rx, ctrlc);
+ let _ = render_stream(rx, abort);
drop(wg);
});
- ReplyStreamHandler::new(Some(tx), self.ctrlc.clone())
+ ReplyStreamHandler::new(Some(tx), self.abort.clone())
} else {
- ReplyStreamHandler::new(None, self.ctrlc.clone())
+ ReplyStreamHandler::new(None, self.abort.clone())
};
self.client
.send_message_streaming(&input, prompt, &mut stream_handler)?;
@@ -109,23 +112,19 @@ impl ReplCmdHandler {
pub fn get_reply(&self) -> String {
self.state.borrow().reply.to_string()
}
-
- pub fn get_ctrlc(&self) -> Arc<AtomicBool> {
- self.ctrlc.clone()
- }
}
pub struct ReplyStreamHandler {
sender: Option<Sender<ReplyStreamEvent>>,
buffer: String,
- ctrlc: Arc<AtomicBool>,
+ abort: SharedAbortSignal,
}
impl ReplyStreamHandler {
- pub fn new(sender: Option<Sender<ReplyStreamEvent>>, ctrlc: Arc<AtomicBool>) -> Self {
+ pub fn new(sender: Option<Sender<ReplyStreamEvent>>, abort: SharedAbortSignal) -> Self {
Self {
sender,
- ctrlc,
+ abort,
buffer: String::new(),
}
}
@@ -157,8 +156,8 @@ impl ReplyStreamHandler {
&self.buffer
}
- pub fn get_ctrlc(&self) -> Arc<AtomicBool> {
- self.ctrlc.clone()
+ pub fn get_abort(&self) -> SharedAbortSignal {
+ self.abort.clone()
}
}
diff --git a/src/repl/mod.rs b/src/repl/mod.rs
index 826ec6a..30e1c5e 100644
--- a/src/repl/mod.rs
+++ b/src/repl/mod.rs
@@ -1,3 +1,4 @@
+mod abort;
mod handler;
mod init;
@@ -8,9 +9,9 @@ use crate::utils::{copy, dump};
use anyhow::{Context, Result};
use reedline::{DefaultPrompt, Reedline, Signal};
-use std::sync::atomic::Ordering;
use std::sync::Arc;
+pub use self::abort::*;
pub use self::handler::*;
pub const REPL_COMMANDS: [(&str, &str, bool); 12] = [
@@ -35,23 +36,27 @@ pub struct Repl {
impl Repl {
pub fn run(&mut self, client: ChatGptClient, config: SharedConfig) -> Result<()> {
- let handler = ReplCmdHandler::init(client, config)?;
+ let abort = AbortSignal::new();
+ let handler = ReplCmdHandler::init(client, config, abort.clone())?;
dump(
format!("Welcome to aichat {}", env!("CARGO_PKG_VERSION")),
1,
);
dump("Type \".help\" for more information.", 1);
- let mut current_ctrlc = false;
+ let mut already_ctrlc = false;
let handler = Arc::new(handler);
loop {
- let handler_ctrlc = handler.get_ctrlc();
- if handler_ctrlc.load(Ordering::SeqCst) {
- handler_ctrlc.store(false, Ordering::SeqCst);
- current_ctrlc = true
+ if abort.aborted_ctrld() {
+ break;
}
- match self.editor.read_line(&self.prompt) {
+ if abort.aborted_ctrlc() && !already_ctrlc {
+ already_ctrlc = true;
+ }
+ let sig = self.editor.read_line(&self.prompt);
+ match sig {
Ok(Signal::Success(line)) => {
- current_ctrlc = false;
+ already_ctrlc = false;
+ abort.reset();
match self.handle_line(handler.clone(), line) {
Ok(quit) => {
if quit {
@@ -65,14 +70,16 @@ impl Repl {
}
}
Ok(Signal::CtrlC) => {
- if !current_ctrlc {
- current_ctrlc = true;
+ abort.set_ctrlc();
+ if !already_ctrlc {
+ already_ctrlc = true;
dump("(To exit, press Ctrl+C again or Ctrl+D or type .exit)", 2);
} else {
break;
}
}
Ok(Signal::CtrlD) => {
+ abort.set_ctrld();
break;
}
_ => {}