From 5985551abaf9418c1c37ef9fd6db39a6c9c94d69 Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 26 Jun 2024 21:51:06 +0800 Subject: feat: rag load websites (#655) --- src/rag/loader.rs | 201 ++++++++++++++++++++++++++++++++++++++++++++++++------ 1 file changed, 180 insertions(+), 21 deletions(-) (limited to 'src/rag/loader.rs') diff --git a/src/rag/loader.rs b/src/rag/loader.rs index 21fc79d..4ab0372 100644 --- a/src/rag/loader.rs +++ b/src/rag/loader.rs @@ -2,22 +2,43 @@ use super::*; use anyhow::{bail, Context, Result}; use async_recursion::async_recursion; -use std::{collections::HashMap, fs::read_to_string, path::Path}; +use serde_json::Value; +use std::{collections::HashMap, env, fs::read_to_string, path::Path}; -pub fn load_file( +pub const RECURSIVE_URL_LOADER: &str = "recursive_url"; + +pub fn load( loaders: &HashMap, path: &str, loader_name: &str, ) -> Result> { - match loaders.get(loader_name) { - Some(loader_command) => load_with_command(path, loader_name, loader_command), - None => load_plain(path), + if loader_name == RECURSIVE_URL_LOADER { + let loader_command = loaders + .get(loader_name) + .with_context(|| format!("RAG document loader '{loader_name}' not configured"))?; + let contents = run_loader_command(path, loader_name, loader_command)?; + let output = match parse_json_documents(&contents) { + Some(v) => v, + None => vec![RagDocument::new(contents)], + }; + Ok(output) + } else { + match loaders.get(loader_name) { + Some(loader_command) => load_with_command(path, loader_name, loader_command), + None => load_plain(path, loader_name), + } } } -fn load_plain(path: &str) -> Result> { +fn load_plain(path: &str, loader_name: &str) -> Result> { let contents = read_to_string(path)?; - let document = RagDocument::new(contents); + if loader_name == "json" { + if let Some(documents) = parse_json_documents(&contents) { + return Ok(documents); + } + } + let mut document = RagDocument::new(contents); + document.metadata.insert("path".into(), path.to_string()); Ok(vec![document]) } @@ -26,29 +47,135 @@ fn load_with_command( loader_name: &str, loader_command: &str, ) -> Result> { - let cmd_args = shell_words::split(loader_command) - .with_context(|| anyhow!("Invalid rag loader '{loader_name}': `{loader_command}`"))?; + let contents = run_loader_command(path, loader_name, loader_command)?; + let mut document = RagDocument::new(contents); + document.metadata.insert("path".into(), path.to_string()); + Ok(vec![document]) +} + +fn run_loader_command(path: &str, loader_name: &str, loader_command: &str) -> Result { + let cmd_args = shell_words::split(loader_command).with_context(|| { + anyhow!("Invalid rag document loader '{loader_name}': `{loader_command}`") + })?; + let mut use_stdout = true; + let outpath = env::temp_dir() + .join(format!("aichat-{}", sha256(path))) + .display() + .to_string(); let cmd_args: Vec<_> = cmd_args .into_iter() - .map(|v| if v == "$1" { path.to_string() } else { v }) + .map(|mut v| { + if v.contains("$1") { + v = v.replace("$1", path); + } + if v.contains("$2") { + use_stdout = false; + v = v.replace("$2", &outpath); + } + v + }) .collect(); let cmd_eval = shell_words::join(&cmd_args); + debug!("run `{cmd_eval}`"); let (cmd, args) = cmd_args.split_at(1); let cmd = &cmd[0]; - let (success, stdout, stderr) = - run_command_with_output(cmd, args, None).with_context(|| { + if use_stdout { + let (success, stdout, stderr) = + run_command_with_output(cmd, args, None).with_context(|| { + format!("Unable to run `{cmd_eval}`, Perhaps '{cmd}' is not installed?") + })?; + if !success { + let err = if !stderr.is_empty() { + stderr + } else { + format!("The command `{cmd_eval}` exited with non-zero.") + }; + bail!("{err}") + } + Ok(stdout) + } else { + let status = run_command(cmd, args, None).with_context(|| { format!("Unable to run `{cmd_eval}`, Perhaps '{cmd}' is not installed?") })?; - if !success { - let err = if !stderr.is_empty() { - stderr - } else { - format!("The command `{cmd_eval}` exited with non-zero.") - }; - bail!("{err}") + if status != 0 { + bail!("The command `{cmd_eval}` exited with non-zero.") + } + let contents = + read_to_string(&outpath).context("Failed to read file generated by the loader")?; + Ok(contents) + } +} + +fn parse_json_documents(data: &str) -> Option> { + 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", + "data", + ] + .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 metadata: IndexMap<_, _> = obj + .into_iter() + .map(|(k, v)| { + if let Value::String(v) = v { + (k, v) + } else { + (k, v.to_string()) + } + }) + .collect(); + return Some(RagDocument { + page_content, + metadata, + }); + } + } + None + }) + .collect(); + if documents.is_empty() { + None + } else { + Some(documents) + } + } + _ => None, } - let document = RagDocument::new(stdout); - Ok(vec![document]) } pub fn parse_glob(path_str: &str) -> Result<(String, Vec)> { @@ -146,4 +273,36 @@ 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, "data": "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()); + } } -- cgit v1.2.3