From ec83167de6ae77dfde34d0afc2e87f03997be4d6 Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 27 Jun 2024 14:55:25 +0800 Subject: refactor: smart document splitter (#662) --- src/rag/mod.rs | 25 ++++++++++++++----------- 1 file changed, 14 insertions(+), 11 deletions(-) (limited to 'src/rag/mod.rs') 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 -- cgit v1.2.3