diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/mod.rs | 17 | ||||
| -rw-r--r-- | src/rag/loader.rs | 77 | ||||
| -rw-r--r-- | src/rag/mod.rs | 13 |
3 files changed, 62 insertions, 45 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs index 62ed28e..26c0bda 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -115,6 +115,8 @@ pub struct Config { pub rag_min_score_vector_search: f32, pub rag_min_score_keyword_search: f32, pub rag_min_score_rerank: f32, + #[serde(default)] + pub rag_document_loaders: HashMap<String, String>, pub rag_template: Option<String>, pub highlight: bool, @@ -174,6 +176,7 @@ impl Default for Config { rag_min_score_vector_search: 0.0, rag_min_score_keyword_search: 0.0, rag_min_score_rerank: 0.0, + rag_document_loaders: Default::default(), rag_template: None, save_session: None, @@ -229,6 +232,7 @@ impl Config { config.setup_model()?; config.setup_highlight(); config.setup_light_theme()?; + config.setup_rag_document_loaders(); Ok(config) } @@ -1440,6 +1444,19 @@ impl Config { }; Ok(()) } + + fn setup_rag_document_loaders(&mut self) { + [ + ("pdf", "pdftotext $1 -"), + ("docx", "pandoc --to plain $1"), + ("url", "curl -fsSL $1"), + ] + .into_iter() + .for_each(|(k, v)| { + let (k, v) = (k.to_string(), v.to_string()); + self.rag_document_loaders.entry(k).or_insert(v); + }); + } } #[derive(Debug, Clone, Deserialize, Default)] diff --git a/src/rag/loader.rs b/src/rag/loader.rs index ba44dac..21fc79d 100644 --- a/src/rag/loader.rs +++ b/src/rag/loader.rs @@ -1,21 +1,17 @@ use super::*; -use anyhow::{bail, Result}; +use anyhow::{bail, Context, Result}; use async_recursion::async_recursion; -use lazy_static::lazy_static; -use std::{fs::read_to_string, path::Path}; -use which::which; +use std::{collections::HashMap, fs::read_to_string, path::Path}; -lazy_static! { - static ref EXIST_PANDOC: bool = which("pandoc").is_ok(); - static ref EXIST_PDFTOTEXT: bool = which("pdftotext").is_ok(); -} - -pub fn load(path: &str, extension: &str) -> Result<Vec<RagDocument>> { - match extension { - "docx" | "epub" => load_with_pandoc(path), - "pdf" => load_with_pdftotext(path), - _ => load_plain(path), +pub fn load_file( + loaders: &HashMap<String, String>, + path: &str, + loader_name: &str, +) -> Result<Vec<RagDocument>> { + match loaders.get(loader_name) { + Some(loader_command) => load_with_command(path, loader_name, loader_command), + None => load_plain(path), } } @@ -25,21 +21,33 @@ fn load_plain(path: &str) -> Result<Vec<RagDocument>> { Ok(vec![document]) } -fn load_with_pdftotext(path: &str) -> Result<Vec<RagDocument>> { - if !*EXIST_PDFTOTEXT { - bail!("Need to install pdftotext (part of the poppler package) to load the file.") - } - let contents = run_external_tool("pdftotext", &[path, "-"])?; - let document = RagDocument::new(contents); - Ok(vec![document]) -} - -fn load_with_pandoc(path: &str) -> Result<Vec<RagDocument>> { - if !*EXIST_PANDOC { - bail!("Need to install pandoc to load the file.") +fn load_with_command( + path: &str, + loader_name: &str, + loader_command: &str, +) -> Result<Vec<RagDocument>> { + let cmd_args = shell_words::split(loader_command) + .with_context(|| anyhow!("Invalid rag loader '{loader_name}': `{loader_command}`"))?; + let cmd_args: Vec<_> = cmd_args + .into_iter() + .map(|v| if v == "$1" { path.to_string() } else { v }) + .collect(); + let cmd_eval = shell_words::join(&cmd_args); + let (cmd, args) = cmd_args.split_at(1); + let cmd = &cmd[0]; + let (success, stdout, stderr) = + run_command_with_output(cmd, args, None).with_context(|| { + format!("Unable to run `{cmd_eval}`, Perhaps '{cmd}' is not installed?") + })?; + if !success { + let err = if !stderr.is_empty() { + stderr + } else { + format!("The command `{cmd_eval}` exited with non-zero.") + }; + bail!("{err}") } - let contents = run_external_tool("pandoc", &["--to", "plain", path])?; - let document = RagDocument::new(contents); + let document = RagDocument::new(stdout); Ok(vec![document]) } @@ -114,19 +122,6 @@ fn is_valid_extension(suffixes: Option<&Vec<String>>, path: &Path) -> bool { true } -fn run_external_tool(cmd: &str, args: &[&str]) -> Result<String> { - let (success, stdout, stderr) = run_command_with_output(cmd, args, None)?; - if success { - return Ok(stdout); - } - let err = if !stderr.is_empty() { - stderr - } else { - format!("`{cmd}` exited with non-zero.") - }; - bail!("{err}") -} - #[cfg(test)] mod tests { use super::*; 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<T: AsRef<Path>>( &mut self, + loaders: HashMap<String, String>, paths: &[T], progress_tx: Option<mpsc::UnboundedSender<String>>, ) -> 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!( "<document_metadata>\npath: {path}\n</document_metadata>\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()), |
