summaryrefslogtreecommitdiffstats
path: root/src/rag/loader.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-26 21:51:06 +0800
committerGitHub <noreply@github.com>2024-06-26 21:51:06 +0800
commit5985551abaf9418c1c37ef9fd6db39a6c9c94d69 (patch)
treed7003f3706572ee5740fba0dfdff703806379501 /src/rag/loader.rs
parent03b40036a177468b1c471eefd365d6f6bc283a7c (diff)
downloadaichat-5985551abaf9418c1c37ef9fd6db39a6c9c94d69.tar.gz
feat: rag load websites (#655)
Diffstat (limited to 'src/rag/loader.rs')
-rw-r--r--src/rag/loader.rs201
1 files changed, 180 insertions, 21 deletions
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<String, String>,
path: &str,
loader_name: &str,
) -> Result<Vec<RagDocument>> {
- 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<Vec<RagDocument>> {
+fn load_plain(path: &str, loader_name: &str) -> Result<Vec<RagDocument>> {
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<Vec<RagDocument>> {
- 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<String> {
+ 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<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",
+ "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<String>)> {
@@ -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());
+ }
}