diff options
| author | sigoden <sigoden@gmail.com> | 2024-07-01 20:19:41 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-07-01 20:19:41 +0800 |
| commit | 7c6dac061b2a33aae334c08b725e02459c5e9ade (patch) | |
| tree | 3c2067aaf80a24f988600b42aa12b647675bc383 /src/rag/loader.rs | |
| parent | d193950d204ed82bdb3d6a3a111511d189f084bb (diff) | |
| download | aichat-7c6dac061b2a33aae334c08b725e02459c5e9ade.tar.gz | |
feat: support `.rebuild rag` repl command(#672)
Diffstat (limited to 'src/rag/loader.rs')
| -rw-r--r-- | src/rag/loader.rs | 240 |
1 files changed, 93 insertions, 147 deletions
diff --git a/src/rag/loader.rs b/src/rag/loader.rs index b9fe298..23f75cb 100644 --- a/src/rag/loader.rs +++ b/src/rag/loader.rs @@ -2,140 +2,108 @@ use super::*; use anyhow::{bail, Context, Result}; use async_recursion::async_recursion; -use serde_json::Value; use std::{collections::HashMap, path::Path}; pub const EXTENSION_METADATA: &str = "__extension__"; +pub const PATH_METADATA: &str = "__path__"; -pub async fn load( +pub async fn load_recrusive_url( loaders: &HashMap<String, String>, path: &str, - extension: &str, -) -> Result<Vec<RagDocument>> { - if extension == RECURSIVE_URL_LOADER { - let loader_command = loaders - .get(extension) - .with_context(|| format!("RAG document loader '{extension}' not configured"))?; - let contents = run_loader_command(path, extension, loader_command)?; - let output = match parse_json_documents(&contents) { - Some(v) => v, - None => vec![RagDocument::new(contents)], - }; - Ok(output) - } else if extension == URL_LOADER { - let (contents, extension) = fetch(loaders, path).await?; - let mut metadata: RagMetadata = Default::default(); - metadata.insert("path".into(), path.into()); - metadata.insert(EXTENSION_METADATA.into(), extension); - Ok(vec![RagDocument::new(contents).with_metadata(metadata)]) +) -> Result<Vec<(String, RagMetadata)>> { + let extension = RECURSIVE_URL_LOADER; + let loader_command = loaders + .get(extension) + .with_context(|| format!("RAG document loader '{extension}' not configured"))?; + let contents = run_loader_command(path, extension, loader_command)?; + let pages: Vec<WebPage> = serde_json::from_str(&contents).context(r#"The crawler response is invalid. It should follow the JSON format: `[{"path":"...", "text":"..."}]`."#)?; + let output = pages + .into_iter() + .map(|v| { + let WebPage { 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) +} + +#[derive(Debug, Deserialize)] +struct WebPage { + path: String, + text: String, +} + +pub async fn load_path( + loaders: &HashMap<String, String>, + path: &str, +) -> Result<Vec<(String, RagMetadata)>> { + let (path_str, suffixes) = parse_glob(path)?; + let suffixes = if suffixes.is_empty() { + None } else { - match loaders.get(extension) { - Some(loader_command) => load_with_command(path, extension, loader_command), - None => load_plain(path, extension).await, + Some(&suffixes) + }; + let mut file_paths = vec![]; + list_files(&mut file_paths, Path::new(&path_str), suffixes).await?; + let mut output = vec![]; + let file_paths_len = file_paths.len(); + match file_paths_len { + 0 => {} + 1 => output.push(load_file(loaders, &file_paths[0]).await?), + _ => { + for path in file_paths { + println!("🚀 Loading file {path}"); + output.push(load_file(loaders, &path).await?) + } + println!("✨ Load directory completed"); } } + Ok(output) } -async fn load_plain(path: &str, extension: &str) -> Result<Vec<RagDocument>> { - let contents = tokio::fs::read_to_string(path).await?; - if extension == "json" { - if let Some(documents) = parse_json_documents(&contents) { - return Ok(documents); - } +pub async fn load_file( + loaders: &HashMap<String, String>, + path: &str, +) -> Result<(String, RagMetadata)> { + let extension = get_extension(path); + match loaders.get(&extension) { + Some(loader_command) => load_with_command(path, &extension, loader_command), + None => load_plain(path, &extension).await, } - let mut document = RagDocument::new(contents); - document.metadata.insert("path".into(), path.to_string()); - Ok(vec![document]) +} + +pub async fn load_url( + loaders: &HashMap<String, String>, + path: &str, +) -> Result<(String, RagMetadata)> { + let (contents, extension) = fetch(loaders, path).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<Vec<RagDocument>> { +) -> Result<(String, RagMetadata)> { let contents = run_loader_command(path, extension, loader_command)?; - let mut document = RagDocument::new(contents); - document.metadata.insert("path".into(), path.to_string()); - document - .metadata - .insert(EXTENSION_METADATA.into(), DEFAULT_EXTENSION.to_string()); - Ok(vec![document]) -} - -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", - ] - .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 mut metadata: IndexMap<_, _> = obj - .into_iter() - .map(|(k, v)| { - if let Value::String(v) = v { - (k, v) - } else { - (k, v.to_string()) - } - }) - .collect(); - if key == "markdown" { - metadata.insert(EXTENSION_METADATA.into(), "md".into()); - } else if key == "html" { - metadata.insert(EXTENSION_METADATA.into(), "html".into()); - } - return Some(RagDocument { - page_content, - metadata, - }); - } - } - None - }) - .collect(); - if documents.is_empty() { - None - } else { - Some(documents) - } - } - _ => None, - } + 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)) } pub fn parse_glob(path_str: &str) -> Result<(String, Vec<String>)> { @@ -158,6 +126,8 @@ pub fn parse_glob(path_str: &str) -> Result<(String, Vec<String>)> { let extensions = vec![extensions_str.to_string()]; Ok((base_path, extensions)) } + } else if path_str.ends_with("/**") || path_str.ends_with(r"\**") { + Ok((path_str[0..path_str.len() - 3].to_string(), vec![])) } else { Ok((path_str.to_string(), vec![])) } @@ -209,6 +179,13 @@ fn is_valid_extension(suffixes: Option<&Vec<String>>, path: &Path) -> bool { true } +fn get_extension(path: &str) -> String { + Path::new(&path) + .extension() + .map(|v| v.to_string_lossy().to_lowercase()) + .unwrap_or_else(|| DEFAULT_EXTENSION.to_string()) +} + #[cfg(test)] mod tests { use super::*; @@ -216,6 +193,7 @@ mod tests { #[test] fn test_parse_glob() { assert_eq!(parse_glob("dir").unwrap(), ("dir".into(), vec![])); + assert_eq!(parse_glob("dir/**").unwrap(), ("dir".into(), vec![])); assert_eq!( parse_glob("dir/file.md").unwrap(), ("dir/file.md".into(), vec![]) @@ -233,36 +211,4 @@ 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, "text": "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()); - } } |
