summaryrefslogtreecommitdiffstats
path: root/src/rag
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/rag
parent6f60ad47535b75064a79331c22758d86e6e46d4f (diff)
downloadaichat-2003159d322a96165ab949051c004c456de5be11.tar.gz
feat: handle Ctrl+C during every spinner in REPL (#1014)
Diffstat (limited to 'src/rag')
-rw-r--r--src/rag/mod.rs58
1 files changed, 26 insertions, 32 deletions
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))