summaryrefslogtreecommitdiffstats
path: root/src/rag/loader.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-07-01 20:19:41 +0800
committerGitHub <noreply@github.com>2024-07-01 20:19:41 +0800
commit7c6dac061b2a33aae334c08b725e02459c5e9ade (patch)
tree3c2067aaf80a24f988600b42aa12b647675bc383 /src/rag/loader.rs
parentd193950d204ed82bdb3d6a3a111511d189f084bb (diff)
downloadaichat-7c6dac061b2a33aae334c08b725e02459c5e9ade.tar.gz
feat: support `.rebuild rag` repl command(#672)
Diffstat (limited to 'src/rag/loader.rs')
-rw-r--r--src/rag/loader.rs240
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());
- }
}