summaryrefslogtreecommitdiffstats
path: root/src/rag/mod.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-12-04 06:56:25 +0800
committerGitHub <noreply@github.com>2024-12-04 06:56:25 +0800
commit5efb44c644cf211f64e241ebf961abc36d9a51e1 (patch)
tree530e2e89716206107a75ec6b189dbdb19c57469a /src/rag/mod.rs
parente86a9f3ee93d6ea0d19e0af04c5a63b41ca72582 (diff)
downloadaichat-5efb44c644cf211f64e241ebf961abc36d9a51e1.tar.gz
feat: `--file/.file` accept file formats in document_loaders (#1033)
Diffstat (limited to 'src/rag/mod.rs')
-rw-r--r--src/rag/mod.rs25
1 files changed, 11 insertions, 14 deletions
diff --git a/src/rag/mod.rs b/src/rag/mod.rs
index 5760871..03b8f7e 100644
--- a/src/rag/mod.rs
+++ b/src/rag/mod.rs
@@ -1,11 +1,9 @@
-use self::loader::*;
use self::splitter::*;
use crate::client::*;
use crate::config::*;
use crate::utils::*;
-mod loader;
mod serde_vectors;
mod splitter;
@@ -367,7 +365,7 @@ impl Rag {
}
}
- let mut files = vec![];
+ let mut loaded_documents = vec![];
let mut has_error = false;
let mut index = 0;
let total = recursive_urls.len() + urls.len() + local_paths.len();
@@ -379,7 +377,7 @@ impl Rag {
index += 1;
println!("Load {start_url}** [{index}/{total}]");
match load_recursive_url(&loaders, &start_url).await {
- Ok(v) => files.extend(v),
+ Ok(v) => loaded_documents.extend(v),
Err(err) => handle_error(err, &mut has_error),
}
}
@@ -387,7 +385,7 @@ impl Rag {
index += 1;
println!("Load {url} [{index}/{total}]");
match load_url(&loaders, &url).await {
- Ok(v) => files.push(v),
+ Ok(v) => loaded_documents.push(v),
Err(err) => handle_error(err, &mut has_error),
}
}
@@ -395,7 +393,7 @@ impl Rag {
index += 1;
println!("Load {local_path} [{index}/{total}]");
match load_file(&loaders, &local_path).await {
- Ok(v) => files.push(v),
+ Ok(v) => loaded_documents.push(v),
Err(err) => handle_error(err, &mut has_error),
}
}
@@ -414,11 +412,12 @@ impl Rag {
}
let mut rag_files = vec![];
- for (contents, mut metadata) in files {
- let path = match metadata.swap_remove(PATH_METADATA) {
- Some(v) => v,
- None => continue,
- };
+ for LoadedDocument {
+ path,
+ contents,
+ mut metadata,
+ } in loaded_documents
+ {
let hash = sha256(&contents);
if let Some(file_ids) = to_deleted.get_mut(&hash) {
if let Some((i, _)) = file_ids
@@ -793,7 +792,7 @@ pub struct RagFile {
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct RagDocument {
pub page_content: String,
- pub metadata: RagMetadata,
+ pub metadata: DocumentMetadata,
}
impl RagDocument {
@@ -814,8 +813,6 @@ impl Default for RagDocument {
}
}
-pub type RagMetadata = IndexMap<String, String>;
-
pub type FileId = usize;
#[derive(Clone, Copy, Hash, Eq, PartialEq, Ord, PartialOrd)]