summaryrefslogtreecommitdiffstats
path: root/src/rag/mod.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-07-28 10:15:43 +0800
committerGitHub <noreply@github.com>2024-07-28 10:15:43 +0800
commitc5b2be641acd43375f92fa04fb679911840c1c56 (patch)
tree1654b09766b87c41e707da7fcf75dbc6b1b2c7c9 /src/rag/mod.rs
parent49b61129c95a3528eaf25dabcb55825b5ed7be72 (diff)
downloadaichat-c5b2be641acd43375f92fa04fb679911840c1c56.tar.gz
feat: ask for confirmation when some rag documents fail to load (#760)
Diffstat (limited to 'src/rag/mod.rs')
-rw-r--r--src/rag/mod.rs65
1 files changed, 46 insertions, 19 deletions
diff --git a/src/rag/mod.rs b/src/rag/mod.rs
index d97d35c..184f51b 100644
--- a/src/rag/mod.rs
+++ b/src/rag/mod.rs
@@ -14,7 +14,7 @@ use anyhow::bail;
use anyhow::{anyhow, Context, Result};
use hnsw_rs::prelude::*;
use indexmap::{IndexMap, IndexSet};
-use inquire::{required, validator::Validation, Select, Text};
+use inquire::{required, validator::Validation, Confirm, Select, Text};
use path_absolutize::Absolutize;
use serde::{Deserialize, Serialize};
use serde_json::json;
@@ -62,7 +62,7 @@ impl Rag {
let loaders = config.read().document_loaders.clone();
let spinner = create_spinner("Starting").await;
tokio::select! {
- ret = rag.load_paths(loaders, &paths, Some(spinner.clone())) => {
+ ret = rag.sync_documents(loaders, &paths, Some(spinner.clone())) => {
spinner.stop();
ret?;
}
@@ -114,7 +114,7 @@ impl Rag {
let spinner = create_spinner("Starting").await;
let paths = self.data.document_paths.clone();
tokio::select! {
- ret = self.load_paths(loaders, &paths, Some(spinner.clone())) => {
+ ret = self.sync_documents(loaders, &paths, Some(spinner.clone())) => {
spinner.stop();
ret?;
}
@@ -258,7 +258,7 @@ impl Rag {
Ok(output)
}
- pub async fn load_paths<T: AsRef<str>>(
+ pub async fn sync_documents<T: AsRef<str>>(
&mut self,
loaders: HashMap<String, String>,
paths: &[T],
@@ -271,21 +271,32 @@ impl Rag {
let mut document_paths = vec![];
let mut files = vec![];
let paths_len = paths.len();
+ let mut has_error = false;
for (index, path) in paths.iter().enumerate() {
let path = path.as_ref();
println!("Load {path} [{}/{paths_len}]", index + 1);
- if Self::is_url_path(path) {
- if let Some(path) = path.strip_suffix("**") {
- files.extend(load_recursive_url(&loaders, path).await?);
- } else {
- files.push(load_url(&loaders, path).await?);
+ match load_document(&loaders, path).await {
+ Ok((path, document_files)) => {
+ files.extend(document_files);
+ document_paths.push(path);
}
- document_paths.push(path.to_string());
- } else {
- let path = Path::new(path);
- let path = path.absolutize()?.display().to_string();
- files.extend(load_path(&loaders, &path).await?);
- document_paths.push(path);
+ Err(err) => {
+ has_error = true;
+ println!("{}", warning_text(&format!("Error: {err:?}")));
+ }
+ }
+ }
+
+ if has_error {
+ let mut aborted = true;
+ if *IS_STDOUT_TERMINAL && !document_paths.is_empty() {
+ let ans = Confirm::new("Some documents failed to load. Continue?")
+ .with_default(false)
+ .prompt()?;
+ aborted = !ans;
+ }
+ if aborted {
+ bail!("Aborted");
}
}
@@ -367,10 +378,6 @@ impl Rag {
Ok(())
}
- pub fn is_url_path(path: &str) -> bool {
- path.starts_with("http://") || path.starts_with("https://")
- }
-
async fn hybird_search(
&self,
query: &str,
@@ -708,6 +715,26 @@ fn add_documents() -> Result<Vec<String>> {
Ok(paths)
}
+async fn load_document(
+ loaders: &HashMap<String, String>,
+ path: &str,
+) -> Result<(String, Vec<(String, RagMetadata)>)> {
+ let mut files = vec![];
+ if is_url(path) {
+ if let Some(path) = path.strip_suffix("**") {
+ files.extend(load_recursive_url(loaders, path).await?);
+ } else {
+ files.push(load_url(loaders, path).await?);
+ }
+ Ok((path.to_string(), files))
+ } else {
+ let path = Path::new(path);
+ let path = path.absolutize()?.display().to_string();
+ files.extend(load_path(loaders, &path).await?);
+ Ok((path.to_string(), files))
+ }
+}
+
fn progress(spinner: &Option<Spinner>, message: String) {
if let Some(spinner) = spinner {
let _ = spinner.set_message(message);