summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-26 08:18:58 +0800
committerGitHub <noreply@github.com>2024-06-26 08:18:58 +0800
commit95bad975f4fb47fe86df1837660a1341009ea12a (patch)
tree1d88411576d07d6e46311ac726a8577d1e5e38f2 /src
parent34a6d13fb6e492860e1340ab7a97cbd1bdeb9d4d (diff)
downloadaichat-95bad975f4fb47fe86df1837660a1341009ea12a.tar.gz
feat: custom rag document loaders (#650)
Diffstat (limited to 'src')
-rw-r--r--src/config/mod.rs17
-rw-r--r--src/rag/loader.rs77
-rw-r--r--src/rag/mod.rs13
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()),