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 | |
| parent | e86a9f3ee93d6ea0d19e0af04c5a63b41ca72582 (diff) | |
| download | aichat-5efb44c644cf211f64e241ebf961abc36d9a51e1.tar.gz | |
feat: `--file/.file` accept file formats in document_loaders (#1033)
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/input.rs | 18 | ||||
| -rw-r--r-- | src/rag/mod.rs | 25 | ||||
| -rw-r--r-- | src/rag/splitter/mod.rs | 6 | ||||
| -rw-r--r-- | src/utils/loader.rs (renamed from src/rag/loader.rs) | 61 | ||||
| -rw-r--r-- | src/utils/mod.rs | 2 |
5 files changed, 55 insertions, 57 deletions
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<String> { Ok(data_url) } - -fn read_file<P: AsRef<Path>>(file_path: P) -> Result<String> { - 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/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 { diff --git a/src/rag/loader.rs b/src/utils/loader.rs index 658aa2f..519563c 100644 --- a/src/rag/loader.rs +++ b/src/utils/loader.rs @@ -1,15 +1,34 @@ use super::*; use anyhow::{Context, Result}; +use indexmap::IndexMap; use std::collections::HashMap; pub const EXTENSION_METADATA: &str = "__extension__"; -pub const PATH_METADATA: &str = "__path__"; + +pub type DocumentMetadata = IndexMap<String, String>; + +#[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<String, String>, path: &str, -) -> Result<Vec<(String, RagMetadata)>> { +) -> Result<Vec<LoadedDocument>> { let extension = RECURSIVE_URL_LOADER; let pages: Vec<Page> = match loaders.get(extension) { Some(loader_command) => { @@ -25,19 +44,15 @@ pub async fn load_recursive_url( .into_iter() .map(|v| { let Page { path, text } = v; - let mut metadata: RagMetadata = Default::default(); - metadata.insert(PATH_METADATA.into(), path); + let mut metadata: DocumentMetadata = Default::default(); metadata.insert(EXTENSION_METADATA.into(), "md".into()); - (text, metadata) + LoadedDocument::new(path, text, metadata) }) .collect(); Ok(output) } -pub async fn load_file( - loaders: &HashMap<String, String>, - path: &str, -) -> Result<(String, RagMetadata)> { +pub async fn load_file(loaders: &HashMap<String, String>, path: &str) -> Result<LoadedDocument> { 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), @@ -45,33 +60,23 @@ pub async fn load_file( } } -pub async fn load_url( - loaders: &HashMap<String, String>, - path: &str, -) -> Result<(String, RagMetadata)> { +pub async fn load_url(loaders: &HashMap<String, String>, path: &str) -> Result<LoadedDocument> { let (contents, extension) = fetch(loaders, path, false).await?; - let mut metadata: RagMetadata = Default::default(); - metadata.insert(PATH_METADATA.into(), path.into()); + let mut metadata: DocumentMetadata = Default::default(); metadata.insert(EXTENSION_METADATA.into(), extension); - Ok((contents, metadata)) + Ok(LoadedDocument::new(path.into(), contents, metadata)) } -async fn load_plain(path: &str, extension: &str) -> Result<(String, RagMetadata)> { +async fn load_plain(path: &str, extension: &str) -> Result<LoadedDocument> { let contents = tokio::fs::read_to_string(path).await?; - let mut metadata: RagMetadata = Default::default(); - metadata.insert(PATH_METADATA.into(), path.to_string()); + let mut metadata: DocumentMetadata = Default::default(); metadata.insert(EXTENSION_METADATA.into(), extension.to_string()); - Ok((contents, metadata)) + Ok(LoadedDocument::new(path.into(), contents, metadata)) } -fn load_with_command( - path: &str, - extension: &str, - loader_command: &str, -) -> Result<(String, RagMetadata)> { +fn load_with_command(path: &str, extension: &str, loader_command: &str) -> Result<LoadedDocument> { let contents = run_loader_command(path, extension, loader_command)?; - let mut metadata: RagMetadata = Default::default(); - metadata.insert(PATH_METADATA.into(), path.to_string()); + let mut metadata: DocumentMetadata = Default::default(); metadata.insert(EXTENSION_METADATA.into(), DEFAULT_EXTENSION.to_string()); - Ok((contents, metadata)) + 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; |
