diff options
Diffstat (limited to 'src/rag')
| -rw-r--r-- | src/rag/loader.rs | 201 | ||||
| -rw-r--r-- | src/rag/mod.rs | 162 | ||||
| -rw-r--r-- | src/rag/splitter/mod.rs | 32 |
3 files changed, 285 insertions, 110 deletions
diff --git a/src/rag/loader.rs b/src/rag/loader.rs index 21fc79d..4ab0372 100644 --- a/src/rag/loader.rs +++ b/src/rag/loader.rs @@ -2,22 +2,43 @@ use super::*; use anyhow::{bail, Context, Result}; use async_recursion::async_recursion; -use std::{collections::HashMap, fs::read_to_string, path::Path}; +use serde_json::Value; +use std::{collections::HashMap, env, fs::read_to_string, path::Path}; -pub fn load_file( +pub const RECURSIVE_URL_LOADER: &str = "recursive_url"; + +pub fn load( 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), + if loader_name == RECURSIVE_URL_LOADER { + let loader_command = loaders + .get(loader_name) + .with_context(|| format!("RAG document loader '{loader_name}' not configured"))?; + let contents = run_loader_command(path, loader_name, loader_command)?; + let output = match parse_json_documents(&contents) { + Some(v) => v, + None => vec![RagDocument::new(contents)], + }; + Ok(output) + } else { + match loaders.get(loader_name) { + Some(loader_command) => load_with_command(path, loader_name, loader_command), + None => load_plain(path, loader_name), + } } } -fn load_plain(path: &str) -> Result<Vec<RagDocument>> { +fn load_plain(path: &str, loader_name: &str) -> Result<Vec<RagDocument>> { let contents = read_to_string(path)?; - let document = RagDocument::new(contents); + if loader_name == "json" { + if let Some(documents) = parse_json_documents(&contents) { + return Ok(documents); + } + } + let mut document = RagDocument::new(contents); + document.metadata.insert("path".into(), path.to_string()); Ok(vec![document]) } @@ -26,29 +47,135 @@ fn load_with_command( 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 contents = run_loader_command(path, loader_name, loader_command)?; + let mut document = RagDocument::new(contents); + document.metadata.insert("path".into(), path.to_string()); + Ok(vec![document]) +} + +fn run_loader_command(path: &str, loader_name: &str, loader_command: &str) -> Result<String> { + let cmd_args = shell_words::split(loader_command).with_context(|| { + anyhow!("Invalid rag document loader '{loader_name}': `{loader_command}`") + })?; + let mut use_stdout = true; + let outpath = env::temp_dir() + .join(format!("aichat-{}", sha256(path))) + .display() + .to_string(); let cmd_args: Vec<_> = cmd_args .into_iter() - .map(|v| if v == "$1" { path.to_string() } else { v }) + .map(|mut v| { + if v.contains("$1") { + v = v.replace("$1", path); + } + if v.contains("$2") { + use_stdout = false; + v = v.replace("$2", &outpath); + } + v + }) .collect(); let cmd_eval = shell_words::join(&cmd_args); + debug!("run `{cmd_eval}`"); 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(|| { + if use_stdout { + 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}") + } + Ok(stdout) + } else { + let status = run_command(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}") + if status != 0 { + bail!("The command `{cmd_eval}` exited with non-zero.") + } + let contents = + read_to_string(&outpath).context("Failed to read file generated by the loader")?; + Ok(contents) + } +} + +fn parse_json_documents(data: &str) -> Option<Vec<RagDocument>> { + let value: Value = serde_json::from_str(data).ok()?; + let items = match value { + Value::Array(v) => v, + _ => return None, + }; + if items.is_empty() { + return None; + } + match &items[0] { + Value::String(_) => { + let documents: Vec<_> = items + .into_iter() + .flat_map(|item| { + if let Value::String(content) = item { + Some(RagDocument::new(content)) + } else { + None + } + }) + .collect(); + Some(documents) + } + Value::Object(obj) => { + let key = [ + "page_content", + "pageContent", + "content", + "html", + "markdown", + "text", + "data", + ] + .into_iter() + .map(|v| v.to_string()) + .find(|key| obj.get(key).and_then(|v| v.as_str()).is_some())?; + let documents: Vec<_> = items + .into_iter() + .flat_map(|item| { + if let Value::Object(mut obj) = item { + if let Some(page_content) = obj.get(&key).and_then(|v| v.as_str()) { + let page_content = page_content.to_string(); + obj.remove(&key); + let metadata: IndexMap<_, _> = obj + .into_iter() + .map(|(k, v)| { + if let Value::String(v) = v { + (k, v) + } else { + (k, v.to_string()) + } + }) + .collect(); + return Some(RagDocument { + page_content, + metadata, + }); + } + } + None + }) + .collect(); + if documents.is_empty() { + None + } else { + Some(documents) + } + } + _ => None, } - let document = RagDocument::new(stdout); - Ok(vec![document]) } pub fn parse_glob(path_str: &str) -> Result<(String, Vec<String>)> { @@ -146,4 +273,36 @@ mod tests { ("C:\\dir".into(), vec!["md".into(), "txt".into()]) ); } + + #[test] + fn test_parse_json_documents() { + let data = r#"["foo", "bar"]"#; + assert_eq!( + parse_json_documents(data).unwrap(), + vec![RagDocument::new("foo"), RagDocument::new("bar")] + ); + + let data = r#"[{"content": "foo"}, {"content": "bar"}]"#; + assert_eq!( + parse_json_documents(data).unwrap(), + vec![RagDocument::new("foo"), RagDocument::new("bar")] + ); + + let mut metadata = IndexMap::new(); + metadata.insert("k1".into(), "1".into()); + let data = r#"[{"k1": 1, "data": "foo" }]"#; + assert_eq!( + parse_json_documents(data).unwrap(), + vec![RagDocument::new("foo").with_metadata(metadata.clone())] + ); + + let data = r#""hello""#; + assert!(parse_json_documents(data).is_none()); + + let data = r#"{"key":"value"}"#; + assert!(parse_json_documents(data).is_none()); + + let data = r#"[{"key":"value"}]"#; + assert!(parse_json_documents(data).is_none()); + } } diff --git a/src/rag/mod.rs b/src/rag/mod.rs index e3b799d..3139b7e 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -20,7 +20,6 @@ use serde::{Deserialize, Serialize}; use serde_json::json; use std::collections::HashMap; use std::{fmt::Debug, io::BufReader, path::Path}; -use tokio::sync::mpsc; pub struct Rag { name: String, @@ -61,14 +60,14 @@ impl Rag { }; debug!("doc paths: {paths:?}"); let loaders = config.read().rag_document_loaders.clone(); - let (stop_spinner_tx, set_spinner_message_tx) = run_spinner("Starting").await; + let spinner = create_spinner("Starting").await; tokio::select! { - ret = rag.add_paths(loaders, &paths, Some(set_spinner_message_tx)) => { - let _ = stop_spinner_tx.send(()); + ret = rag.add_paths(loaders, &paths, Some(spinner.clone())) => { + spinner.stop(); ret?; } _ = watch_abort_signal(abort_signal) => { - let _ = stop_spinner_tx.send(()); + spinner.stop(); bail!("Aborted!") }, }; @@ -207,7 +206,7 @@ impl Rag { rerank: Option<(Box<dyn Client>, f32)>, abort_signal: AbortSignal, ) -> Result<String> { - let (stop_spinner_tx, _) = run_spinner("Searching").await; + let spinner = create_spinner("Searching").await; let ret = tokio::select! { ret = self.hybird_search(text, top_k, min_score_vector_search, min_score_keyword_search, rerank) => { ret @@ -216,66 +215,99 @@ impl Rag { bail!("Aborted!") }, }; - let _ = stop_spinner_tx.send(()); + spinner.stop(); let output = ret?.join("\n\n"); Ok(output) } - pub async fn add_paths<T: AsRef<Path>>( + pub async fn add_paths<T: AsRef<str>>( &mut self, loaders: HashMap<String, String>, paths: &[T], - progress_tx: Option<mpsc::UnboundedSender<String>>, + spinner: Option<Spinner>, ) -> Result<()> { + let mut rag_files = vec![]; + // List files - let mut file_paths = vec![]; - progress(&progress_tx, "Listing paths".into()); + let mut new_paths = vec![]; + progress(&spinner, "Gathering paths".into()); for path in paths { - let path = path - .as_ref() - .absolutize() - .with_context(|| anyhow!("Invalid path '{}'", path.as_ref().display()))?; - let path_str = path.display().to_string(); - if self.data.files.iter().any(|v| v.path == path_str) { - continue; - } - let (path_str, suffixes) = parse_glob(&path_str)?; - let suffixes = if suffixes.is_empty() { - None + let path = path.as_ref(); + if path.starts_with("http://") || path.starts_with("https://") { + if let Some(path) = path.strip_suffix("**") { + new_paths.push((path.to_string(), RECURSIVE_URL_LOADER.into())); + } else { + new_paths.push((path.to_string(), "url".into())) + } } else { - Some(&suffixes) - }; - list_files(&mut file_paths, Path::new(&path_str), suffixes).await?; + let path = Path::new(path); + let path = path + .absolutize() + .with_context(|| anyhow!("Invalid path '{}'", path.display()))?; + let path_str = path.display().to_string(); + if self.data.files.iter().any(|v| v.path == path_str) { + continue; + } + let (path_str, suffixes) = parse_glob(&path_str)?; + let suffixes = if suffixes.is_empty() { + None + } else { + Some(&suffixes) + }; + let mut file_paths = vec![]; + list_files(&mut file_paths, Path::new(&path_str), suffixes).await?; + for file_path in file_paths { + let loader_name = Path::new(&file_path) + .extension() + .map(|v| v.to_string_lossy().to_lowercase()) + .unwrap_or_default(); + new_paths.push((file_path, loader_name)) + } + } } // Load files - let mut rag_files = vec![]; - let file_paths_len = file_paths.len(); - progress(&progress_tx, format!("Loading files [1/{file_paths_len}]")); - for path in file_paths { - let extension = Path::new(&path) - .extension() - .map(|v| v.to_string_lossy().to_lowercase()) - .unwrap_or_default(); - let separator = detect_separators(&extension); - let splitter = RecursiveCharacterTextSplitter::new( - self.data.chunk_size, - self.data.chunk_overlap, - &separator, - ); - 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" - )); - if !documents.is_empty() { - let documents = splitter.split_documents(&documents, &split_options); - rag_files.push(RagFile { path, documents }); + let new_paths_len = new_paths.len(); + if new_paths_len > 0 { + if let Some(spinner) = &spinner { + let _ = spinner.set_message(String::new()); + } + for (index, (path, loader_name)) in new_paths.into_iter().enumerate() { + println!("Loading {path} [{}/{new_paths_len}]", index + 1); + let documents = load(&loaders, &path, &loader_name) + .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| { + let metadata = document + .metadata + .iter() + .map(|(k, v)| format!("{k}: {v}\n")) + .collect::<Vec<String>>() + .join(""); + let split_options = SplitterChunkHeaderOptions::default() + .with_chunk_header(&format!( + "<document_metadata>\n{metadata}</document_metadata>\n\n" + )); + splitter.split_documents(&[document], &split_options) + }) + .collect(); + let display_path = if loader_name == RECURSIVE_URL_LOADER { + format!("{path}**") + } else { + path + }; + rag_files.push(RagFile { + path: display_path, + documents: splitted_documents, + }) } - progress( - &progress_tx, - format!("Loading files [{}/{file_paths_len}]", rag_files.len()), - ); } if rag_files.is_empty() { @@ -294,11 +326,11 @@ impl Rag { let embeddings_data = EmbeddingsData::new(texts, false); let embeddings = self - .create_embeddings(embeddings_data, progress_tx.clone()) + .create_embeddings(embeddings_data, spinner.clone()) .await?; self.data.add(rag_files, vector_ids, embeddings); - progress(&progress_tx, "Building vector store".into()); + progress(&spinner, "Building vector store".into()); self.hnsw = self.data.build_hnsw(); self.bm25 = self.data.build_bm25(); @@ -418,17 +450,17 @@ impl Rag { async fn create_embeddings( &self, data: EmbeddingsData, - progress_tx: Option<mpsc::UnboundedSender<String>>, + spinner: Option<Spinner>, ) -> Result<EmbeddingsOutput> { let EmbeddingsData { texts, query } = data; let mut output = vec![]; let batch_chunks = texts.chunks(self.embedding_model.max_batch_size()); let batch_chunks_len = batch_chunks.len(); - progress( - &progress_tx, - format!("Creating embeddings [1/{batch_chunks_len}]"), - ); for (index, texts) in batch_chunks.enumerate() { + progress( + &spinner, + format!("Creating embeddings [{}/{batch_chunks_len}]", index + 1), + ); let chunk_data = EmbeddingsData { texts: texts.to_vec(), query, @@ -439,10 +471,6 @@ impl Rag { .await .context("Failed to create embedding")?; output.extend(chunk_output); - progress( - &progress_tx, - format!("Creating embeddings [{}/{batch_chunks_len}]", index + 1), - ); } Ok(output) } @@ -510,7 +538,7 @@ pub struct RagFile { documents: Vec<RagDocument>, } -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct RagDocument { pub page_content: String, pub metadata: RagMetadata, @@ -603,15 +631,15 @@ fn set_chunk_overlay(default_value: usize) -> Result<usize> { fn add_doc_paths() -> Result<Vec<String>> { let text = Text::new("Add document paths:") .with_validator(required!("This field is required")) - .with_help_message("e.g. file1;dir2/;dir3/**/*.md") + .with_help_message("e.g. file;dir/;dir/**/*.md;url;sites/**") .prompt()?; let paths = text.split(';').map(|v| v.to_string()).collect(); Ok(paths) } -fn progress(spinner_message_tx: &Option<mpsc::UnboundedSender<String>>, message: String) { - if let Some(tx) = spinner_message_tx { - let _ = tx.send(message); +fn progress(spinner: &Option<Spinner>, message: String) { + if let Some(spinner) = spinner { + let _ = spinner.set_message(message); } } diff --git a/src/rag/splitter/mod.rs b/src/rag/splitter/mod.rs index 88054f7..c3d697e 100644 --- a/src/rag/splitter/mod.rs +++ b/src/rag/splitter/mod.rs @@ -8,7 +8,7 @@ use std::cmp::Ordering; pub const DEFAULT_SEPARATES: [&str; 4] = ["\n\n", "\n", " ", ""]; -pub fn detect_separators(extension: &str) -> Vec<&'static str> { +pub fn get_separators(extension: &str) -> Vec<&'static str> { match extension { "c" | "cc" | "cpp" => Language::Cpp.separators(), "go" => Language::Go.separators(), @@ -149,16 +149,7 @@ impl RecursiveCharacterTextSplitter { } let newlines_count = self.number_of_newlines(&chunk, 0, chunk.len()); - - let mut metadata = metadatas[i].clone(); - metadata.insert( - "loc".into(), - format!( - "{}:{}", - line_counter_index, - line_counter_index + newlines_count - ), - ); + let metadata = metadatas[i].clone(); page_content += &chunk; documents.push(RagDocument { page_content, @@ -348,11 +339,8 @@ mod tests { use pretty_assertions::assert_eq; use serde_json::{json, Value}; - fn build_metadata(source: &str, loc_from_line: usize, loc_to_line: usize) -> Value { - json!({ - "source": source, - "loc": format!("{loc_from_line}:{loc_to_line}"), - }) + fn build_metadata(source: &str) -> Value { + json!({ "source": source }) } #[test] fn test_split_text() { @@ -385,15 +373,15 @@ mod tests { json!([ { "page_content": "foo", - "metadata": build_metadata("1", 1, 1), + "metadata": build_metadata("1"), }, { "page_content": "bar", - "metadata": build_metadata("1", 1, 1), + "metadata": build_metadata("1"), }, { "page_content": "baz", - "metadata": build_metadata("2", 1, 1), + "metadata": build_metadata("2"), }, ]) ); @@ -420,15 +408,15 @@ mod tests { json!([ { "page_content": "SOURCE NAME: testing\n-----\nfoo", - "metadata": build_metadata("1", 1, 1), + "metadata": build_metadata("1"), }, { "page_content": "SOURCE NAME: testing\n-----\n(cont'd) bar", - "metadata": build_metadata("1", 1, 1), + "metadata": build_metadata("1"), }, { "page_content": "SOURCE NAME: testing\n-----\nbaz", - "metadata": build_metadata("2", 1, 1), + "metadata": build_metadata("2"), }, ]) ); |
