summaryrefslogtreecommitdiffstats
path: root/src/rag/mod.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-27 14:55:25 +0800
committerGitHub <noreply@github.com>2024-06-27 14:55:25 +0800
commitec83167de6ae77dfde34d0afc2e87f03997be4d6 (patch)
tree2516a680871d8665dafbd52425cb16e0f16be55d /src/rag/mod.rs
parentf82524fd154a1bf4a7277db868147afdb0cd507a (diff)
downloadaichat-ec83167de6ae77dfde34d0afc2e87f03997be4d6.tar.gz
refactor: smart document splitter (#662)
Diffstat (limited to 'src/rag/mod.rs')
-rw-r--r--src/rag/mod.rs25
1 files changed, 14 insertions, 11 deletions
diff --git a/src/rag/mod.rs b/src/rag/mod.rs
index 9e37e6f..536092f 100644
--- a/src/rag/mod.rs
+++ b/src/rag/mod.rs
@@ -272,20 +272,24 @@ impl Rag {
if let Some(spinner) = &spinner {
let _ = spinner.set_message(String::new());
}
- for (index, (path, loader_name)) in new_paths.into_iter().enumerate() {
+ for (index, (path, extension)) in new_paths.into_iter().enumerate() {
println!("Loading {path} [{}/{new_paths_len}]", index + 1);
- let documents = load(&loaders, &path, &loader_name)
+ let documents = load(&loaders, &path, &extension)
.await
.with_context(|| format!("Failed to load '{path}'"))?;
- let separator = get_separators(&loader_name);
- let splitter = RecursiveCharacterTextSplitter::new(
- self.data.chunk_size,
- self.data.chunk_overlap,
- &separator,
- );
let splitted_documents: Vec<_> = documents
.into_iter()
- .flat_map(|document| {
+ .flat_map(|mut document| {
+ let extension = document
+ .metadata
+ .swap_remove(EXTENSION_METADATA)
+ .unwrap_or_else(|| extension.clone());
+ let separator = get_separators(&extension);
+ let splitter = RecursiveCharacterTextSplitter::new(
+ self.data.chunk_size,
+ self.data.chunk_overlap,
+ &separator,
+ );
let metadata = document
.metadata
.iter()
@@ -299,7 +303,7 @@ impl Rag {
splitter.split_documents(&[document], &split_options)
})
.collect();
- let display_path = if loader_name == RECURSIVE_URL_LOADER {
+ let display_path = if extension == RECURSIVE_URL_LOADER {
format!("{path}**")
} else {
path
@@ -557,7 +561,6 @@ impl RagDocument {
}
}
- #[allow(unused)]
pub fn with_metadata(mut self, metadata: RagMetadata) -> Self {
self.metadata = metadata;
self