From 5efb44c644cf211f64e241ebf961abc36d9a51e1 Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 4 Dec 2024 06:56:25 +0800 Subject: feat: `--file/.file` accept file formats in document_loaders (#1033) --- src/config/input.rs | 18 ++++------- src/rag/loader.rs | 77 ---------------------------------------------- src/rag/mod.rs | 25 +++++++-------- src/rag/splitter/mod.rs | 6 ++-- src/utils/loader.rs | 82 +++++++++++++++++++++++++++++++++++++++++++++++++ src/utils/mod.rs | 2 ++ 6 files changed, 104 insertions(+), 106 deletions(-) delete mode 100644 src/rag/loader.rs create mode 100644 src/utils/loader.rs diff --git a/src/config/input.rs b/src/config/input.rs index fa6f00c..42bb8ec 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -90,7 +90,9 @@ impl Input { texts.push(String::new()); } for (path, contents) in files { - texts.push(format!("====== PATH: {path} ======\n{contents}\n")); + texts.push(format!( + "============ PATH: {path} ============\n\n{contents}\n" + )); } let (role, with_session, with_agent) = resolve_role(&config.read(), role); Ok(Self { @@ -392,9 +394,10 @@ async fn load_documents( data_urls.insert(sha256(&data_url), file_path); medias.push(data_url) } else { - let text = read_file(&file_path) + let document = load_file(&loaders, &file_path) + .await .with_context(|| format!("Unable to read file '{file_path}'"))?; - files.push((file_path, text)); + files.push((file_path, document.contents)); } } for file_url in remote_urls { @@ -459,12 +462,3 @@ fn read_media_to_data_url(image_path: &str) -> Result { Ok(data_url) } - -fn read_file>(file_path: P) -> Result { - let file_path = file_path.as_ref(); - - let mut text = String::new(); - let mut file = File::open(file_path)?; - file.read_to_string(&mut text)?; - Ok(text) -} diff --git a/src/rag/loader.rs b/src/rag/loader.rs deleted file mode 100644 index 658aa2f..0000000 --- a/src/rag/loader.rs +++ /dev/null @@ -1,77 +0,0 @@ -use super::*; - -use anyhow::{Context, Result}; -use std::collections::HashMap; - -pub const EXTENSION_METADATA: &str = "__extension__"; -pub const PATH_METADATA: &str = "__path__"; - -pub async fn load_recursive_url( - loaders: &HashMap, - path: &str, -) -> Result> { - let extension = RECURSIVE_URL_LOADER; - let pages: Vec = match loaders.get(extension) { - Some(loader_command) => { - let contents = run_loader_command(path, extension, loader_command)?; - serde_json::from_str(&contents).context(r#"The crawler response is invalid. It should follow the JSON format: `[{"path":"...", "text":"..."}]`."#)? - } - None => { - let options = CrawlOptions::preset(path); - crawl_website(path, options).await? - } - }; - let output = pages - .into_iter() - .map(|v| { - let Page { path, text } = v; - let mut metadata: RagMetadata = Default::default(); - metadata.insert(PATH_METADATA.into(), path); - metadata.insert(EXTENSION_METADATA.into(), "md".into()); - (text, metadata) - }) - .collect(); - Ok(output) -} - -pub async fn load_file( - loaders: &HashMap, - path: &str, -) -> Result<(String, RagMetadata)> { - let extension = get_patch_extension(path).unwrap_or_else(|| DEFAULT_EXTENSION.into()); - match loaders.get(&extension) { - Some(loader_command) => load_with_command(path, &extension, loader_command), - None => load_plain(path, &extension).await, - } -} - -pub async fn load_url( - loaders: &HashMap, - path: &str, -) -> Result<(String, RagMetadata)> { - let (contents, extension) = fetch(loaders, path, false).await?; - let mut metadata: RagMetadata = Default::default(); - metadata.insert(PATH_METADATA.into(), path.into()); - metadata.insert(EXTENSION_METADATA.into(), extension); - Ok((contents, metadata)) -} - -async fn load_plain(path: &str, extension: &str) -> Result<(String, RagMetadata)> { - let contents = tokio::fs::read_to_string(path).await?; - let mut metadata: RagMetadata = Default::default(); - metadata.insert(PATH_METADATA.into(), path.to_string()); - metadata.insert(EXTENSION_METADATA.into(), extension.to_string()); - Ok((contents, metadata)) -} - -fn load_with_command( - path: &str, - extension: &str, - loader_command: &str, -) -> Result<(String, RagMetadata)> { - let contents = run_loader_command(path, extension, loader_command)?; - let mut metadata: RagMetadata = Default::default(); - metadata.insert(PATH_METADATA.into(), path.to_string()); - metadata.insert(EXTENSION_METADATA.into(), DEFAULT_EXTENSION.to_string()); - Ok((contents, metadata)) -} diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 5760871..03b8f7e 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -1,11 +1,9 @@ -use self::loader::*; use self::splitter::*; use crate::client::*; use crate::config::*; use crate::utils::*; -mod loader; mod serde_vectors; mod splitter; @@ -367,7 +365,7 @@ impl Rag { } } - let mut files = vec![]; + let mut loaded_documents = vec![]; let mut has_error = false; let mut index = 0; let total = recursive_urls.len() + urls.len() + local_paths.len(); @@ -379,7 +377,7 @@ impl Rag { index += 1; println!("Load {start_url}** [{index}/{total}]"); match load_recursive_url(&loaders, &start_url).await { - Ok(v) => files.extend(v), + Ok(v) => loaded_documents.extend(v), Err(err) => handle_error(err, &mut has_error), } } @@ -387,7 +385,7 @@ impl Rag { index += 1; println!("Load {url} [{index}/{total}]"); match load_url(&loaders, &url).await { - Ok(v) => files.push(v), + Ok(v) => loaded_documents.push(v), Err(err) => handle_error(err, &mut has_error), } } @@ -395,7 +393,7 @@ impl Rag { index += 1; println!("Load {local_path} [{index}/{total}]"); match load_file(&loaders, &local_path).await { - Ok(v) => files.push(v), + Ok(v) => loaded_documents.push(v), Err(err) => handle_error(err, &mut has_error), } } @@ -414,11 +412,12 @@ impl Rag { } let mut rag_files = vec![]; - for (contents, mut metadata) in files { - let path = match metadata.swap_remove(PATH_METADATA) { - Some(v) => v, - None => continue, - }; + for LoadedDocument { + path, + contents, + mut metadata, + } in loaded_documents + { let hash = sha256(&contents); if let Some(file_ids) = to_deleted.get_mut(&hash) { if let Some((i, _)) = file_ids @@ -793,7 +792,7 @@ pub struct RagFile { #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct RagDocument { pub page_content: String, - pub metadata: RagMetadata, + pub metadata: DocumentMetadata, } impl RagDocument { @@ -814,8 +813,6 @@ impl Default for RagDocument { } } -pub type RagMetadata = IndexMap; - pub type FileId = usize; #[derive(Clone, Copy, Hash, Eq, PartialEq, Ord, PartialOrd)] diff --git a/src/rag/splitter/mod.rs b/src/rag/splitter/mod.rs index 351c560..a4a9167 100644 --- a/src/rag/splitter/mod.rs +++ b/src/rag/splitter/mod.rs @@ -2,7 +2,7 @@ mod language; pub use self::language::*; -use super::{RagDocument, RagMetadata}; +use super::{DocumentMetadata, RagDocument}; pub const DEFAULT_SEPARATES: [&str; 4] = ["\n\n", "\n", " ", ""]; @@ -75,7 +75,7 @@ impl RecursiveCharacterTextSplitter { chunk_header_options: &SplitterChunkHeaderOptions, ) -> Vec { let mut texts: Vec = Vec::new(); - let mut metadatas: Vec = Vec::new(); + let mut metadatas: Vec = Vec::new(); documents.iter().for_each(|d| { if !d.page_content.is_empty() { texts.push(d.page_content.clone()); @@ -89,7 +89,7 @@ impl RecursiveCharacterTextSplitter { pub fn create_documents( &self, texts: &[String], - metadatas: &[RagMetadata], + metadatas: &[DocumentMetadata], chunk_header_options: &SplitterChunkHeaderOptions, ) -> Vec { let SplitterChunkHeaderOptions { diff --git a/src/utils/loader.rs b/src/utils/loader.rs new file mode 100644 index 0000000..519563c --- /dev/null +++ b/src/utils/loader.rs @@ -0,0 +1,82 @@ +use super::*; + +use anyhow::{Context, Result}; +use indexmap::IndexMap; +use std::collections::HashMap; + +pub const EXTENSION_METADATA: &str = "__extension__"; + +pub type DocumentMetadata = IndexMap; + +#[derive(Debug, Clone)] +pub struct LoadedDocument { + pub path: String, + pub contents: String, + pub metadata: DocumentMetadata, +} + +impl LoadedDocument { + pub fn new(path: String, contents: String, metadata: DocumentMetadata) -> Self { + Self { + path, + contents, + metadata, + } + } +} + +pub async fn load_recursive_url( + loaders: &HashMap, + path: &str, +) -> Result> { + let extension = RECURSIVE_URL_LOADER; + let pages: Vec = match loaders.get(extension) { + Some(loader_command) => { + let contents = run_loader_command(path, extension, loader_command)?; + serde_json::from_str(&contents).context(r#"The crawler response is invalid. It should follow the JSON format: `[{"path":"...", "text":"..."}]`."#)? + } + None => { + let options = CrawlOptions::preset(path); + crawl_website(path, options).await? + } + }; + let output = pages + .into_iter() + .map(|v| { + let Page { path, text } = v; + let mut metadata: DocumentMetadata = Default::default(); + metadata.insert(EXTENSION_METADATA.into(), "md".into()); + LoadedDocument::new(path, text, metadata) + }) + .collect(); + Ok(output) +} + +pub async fn load_file(loaders: &HashMap, path: &str) -> Result { + let extension = get_patch_extension(path).unwrap_or_else(|| DEFAULT_EXTENSION.into()); + match loaders.get(&extension) { + Some(loader_command) => load_with_command(path, &extension, loader_command), + None => load_plain(path, &extension).await, + } +} + +pub async fn load_url(loaders: &HashMap, path: &str) -> Result { + let (contents, extension) = fetch(loaders, path, false).await?; + let mut metadata: DocumentMetadata = Default::default(); + metadata.insert(EXTENSION_METADATA.into(), extension); + Ok(LoadedDocument::new(path.into(), contents, metadata)) +} + +async fn load_plain(path: &str, extension: &str) -> Result { + let contents = tokio::fs::read_to_string(path).await?; + let mut metadata: DocumentMetadata = Default::default(); + metadata.insert(EXTENSION_METADATA.into(), extension.to_string()); + Ok(LoadedDocument::new(path.into(), contents, metadata)) +} + +fn load_with_command(path: &str, extension: &str, loader_command: &str) -> Result { + let contents = run_loader_command(path, extension, loader_command)?; + let mut metadata: DocumentMetadata = Default::default(); + metadata.insert(EXTENSION_METADATA.into(), DEFAULT_EXTENSION.to_string()); + Ok(LoadedDocument::new(path.into(), contents, metadata)) +} diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 5271571..c880252 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -3,6 +3,7 @@ mod clipboard; mod command; mod crypto; mod html_to_md; +mod loader; mod path; mod prompt_input; mod render_prompt; @@ -15,6 +16,7 @@ pub use self::clipboard::set_text; pub use self::command::*; pub use self::crypto::*; pub use self::html_to_md::*; +pub use self::loader::*; pub use self::path::*; pub use self::prompt_input::*; pub use self::render_prompt::render_prompt; -- cgit v1.2.3