summaryrefslogtreecommitdiffstats
path: root/src/utils
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-05 09:02:23 +0800
committerGitHub <noreply@github.com>2024-06-05 09:02:23 +0800
commit1ec6abfaee2fdc189b348b7e3a8145bd9a84da74 (patch)
tree6f68a860a39b0fbe87784de925b5a31e00e74e33 /src/utils
parent71f2e94579511d7524f5534377001ab3f02a9597 (diff)
downloadaichat-1ec6abfaee2fdc189b348b7e3a8145bd9a84da74.tar.gz
feat: support RAG (#560)
* feat: support RAG * support more embeddings models and implement concurrent embedding api * show the progress of addings paths * ignore embedding context when saving message * embedding model max_chunk_size => default_chunk_size * support pdf and pandoc formats (docx, epub, ipynb)
Diffstat (limited to 'src/utils')
-rw-r--r--src/utils/abort_signal.rs9
-rw-r--r--src/utils/mod.rs2
-rw-r--r--src/utils/spinner.rs39
3 files changed, 43 insertions, 7 deletions
diff --git a/src/utils/abort_signal.rs b/src/utils/abort_signal.rs
index af58b35..ac93653 100644
--- a/src/utils/abort_signal.rs
+++ b/src/utils/abort_signal.rs
@@ -53,3 +53,12 @@ impl AbortSignalInner {
self.ctrld.store(true, Ordering::SeqCst);
}
}
+
+pub async fn watch_abort_signal(abort: AbortSignal) {
+ loop {
+ if abort.aborted() {
+ break;
+ }
+ tokio::time::sleep(std::time::Duration::from_millis(100)).await;
+ }
+}
diff --git a/src/utils/mod.rs b/src/utils/mod.rs
index fa67c63..95c6725 100644
--- a/src/utils/mod.rs
+++ b/src/utils/mod.rs
@@ -6,7 +6,7 @@ mod prompt_input;
mod render_prompt;
mod spinner;
-pub use self::abort_signal::{create_abort_signal, AbortSignal};
+pub use self::abort_signal::*;
pub use self::clipboard::set_text;
pub use self::command::*;
pub use self::crypto::*;
diff --git a/src/utils/spinner.rs b/src/utils/spinner.rs
index 0dd4d01..c746469 100644
--- a/src/utils/spinner.rs
+++ b/src/utils/spinner.rs
@@ -4,7 +4,10 @@ use std::{
io::{stdout, Stdout, Write},
time::Duration,
};
-use tokio::{sync::oneshot, time::interval};
+use tokio::{
+ sync::{mpsc, oneshot},
+ time::interval,
+};
pub struct Spinner {
index: usize,
@@ -23,6 +26,10 @@ impl Spinner {
}
}
+ pub fn set_message(&mut self, message: &str) {
+ self.message = format!(" {message}");
+ }
+
pub fn step(&mut self, writer: &mut Stdout) -> Result<()> {
if self.stopped {
return Ok(());
@@ -55,18 +62,38 @@ impl Spinner {
}
}
-pub async fn run_spinner(message: &str, rx: oneshot::Receiver<()>) -> Result<()> {
+pub async fn run_spinner(message: &str) -> (oneshot::Sender<()>, mpsc::UnboundedSender<String>) {
+ let message = format!(" {message}");
+ let (stop_tx, stop_rx) = oneshot::channel();
+ let (message_tx, message_rx) = mpsc::unbounded_channel();
+ tokio::spawn(run_spinner_inner(message, stop_rx, message_rx));
+ (stop_tx, message_tx)
+}
+
+async fn run_spinner_inner(
+ message: String,
+ stop_rx: oneshot::Receiver<()>,
+ mut message_rx: mpsc::UnboundedReceiver<String>,
+) -> Result<()> {
let mut writer = stdout();
- let mut spinner = Spinner::new(message);
+ let mut spinner = Spinner::new(&message);
let mut interval = interval(Duration::from_millis(50));
tokio::select! {
_ = async {
loop {
- interval.tick().await;
- let _ = spinner.step(&mut writer);
+ tokio::select! {
+ _ = interval.tick() => {
+ let _ = spinner.step(&mut writer);
+ }
+ message = message_rx.recv() => {
+ if let Some(message) = message {
+ spinner.set_message(&message);
+ }
+ }
+ }
}
} => {}
- _ = rx => {
+ _ = stop_rx => {
spinner.stop(&mut writer)?;
}
}