From 95bad975f4fb47fe86df1837660a1341009ea12a Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 26 Jun 2024 08:18:58 +0800 Subject: feat: custom rag document loaders (#650) --- src/rag/mod.rs | 13 +++++++++---- 1 file changed, 9 insertions(+), 4 deletions(-) (limited to 'src/rag/mod.rs') diff --git a/src/rag/mod.rs b/src/rag/mod.rs index ab7c3c7..e3b799d 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -18,6 +18,7 @@ use inquire::{required, validator::Validation, Select, Text}; use path_absolutize::Absolutize; use serde::{Deserialize, Serialize}; use serde_json::json; +use std::collections::HashMap; use std::{fmt::Debug, io::BufReader, path::Path}; use tokio::sync::mpsc; @@ -59,9 +60,10 @@ impl Rag { paths = add_doc_paths()?; }; debug!("doc paths: {paths:?}"); + let loaders = config.read().rag_document_loaders.clone(); let (stop_spinner_tx, set_spinner_message_tx) = run_spinner("Starting").await; tokio::select! { - ret = rag.add_paths(&paths, Some(set_spinner_message_tx)) => { + ret = rag.add_paths(loaders, &paths, Some(set_spinner_message_tx)) => { let _ = stop_spinner_tx.send(()); ret?; } @@ -221,6 +223,7 @@ impl Rag { pub async fn add_paths>( &mut self, + loaders: HashMap, paths: &[T], progress_tx: Option>, ) -> Result<()> { @@ -260,13 +263,15 @@ impl Rag { self.data.chunk_overlap, &separator, ); - let documents = load(&path, &extension) + let documents = load_file(&loaders, &path, &extension) .with_context(|| format!("Failed to load file at '{path}'"))?; let split_options = SplitterChunkHeaderOptions::default().with_chunk_header(&format!( "\npath: {path}\n\n\n" )); - let documents = splitter.split_documents(&documents, &split_options); - rag_files.push(RagFile { path, documents }); + if !documents.is_empty() { + let documents = splitter.split_documents(&documents, &split_options); + rag_files.push(RagFile { path, documents }); + } progress( &progress_tx, format!("Loading files [{}/{file_paths_len}]", rag_files.len()), -- cgit v1.2.3