diff options
| author | sigoden <sigoden@gmail.com> | 2024-11-27 09:16:06 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-11-27 09:16:06 +0800 |
| commit | 2003159d322a96165ab949051c004c456de5be11 (patch) | |
| tree | dc3c7f8e1850267ba24d1a3bf9a76fdac4366b6a /src/rag | |
| parent | 6f60ad47535b75064a79331c22758d86e6e46d4f (diff) | |
| download | aichat-2003159d322a96165ab949051c004c456de5be11.tar.gz | |
feat: handle Ctrl+C during every spinner in REPL (#1014)
Diffstat (limited to 'src/rag')
| -rw-r--r-- | src/rag/mod.rs | 58 |
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)) |
