diff options
| author | sigoden <sigoden@gmail.com> | 2024-12-04 06:56:25 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-12-04 06:56:25 +0800 |
| commit | 5efb44c644cf211f64e241ebf961abc36d9a51e1 (patch) | |
| tree | 530e2e89716206107a75ec6b189dbdb19c57469a /src/rag | |
| parent | e86a9f3ee93d6ea0d19e0af04c5a63b41ca72582 (diff) | |
| download | aichat-5efb44c644cf211f64e241ebf961abc36d9a51e1.tar.gz | |
feat: `--file/.file` accept file formats in document_loaders (#1033)
Diffstat (limited to 'src/rag')
| -rw-r--r-- | src/rag/loader.rs | 77 | ||||
| -rw-r--r-- | src/rag/mod.rs | 25 | ||||
| -rw-r--r-- | src/rag/splitter/mod.rs | 6 |
3 files changed, 14 insertions, 94 deletions
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<String, String>, - path: &str, -) -> Result<Vec<(String, RagMetadata)>> { - let extension = RECURSIVE_URL_LOADER; - let pages: Vec<Page> = 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<String, String>, - 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<String, String>, - 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<String, String>; - 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<RagDocument> { let mut texts: Vec<String> = Vec::new(); - let mut metadatas: Vec<RagMetadata> = Vec::new(); + let mut metadatas: Vec<DocumentMetadata> = 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<RagDocument> { let SplitterChunkHeaderOptions { |
