From f82524fd154a1bf4a7277db868147afdb0cd507a Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 27 Jun 2024 12:30:09 +0800 Subject: feat: implement native rag url loader (#660) --- src/rag/loader.rs | 63 ++++++++++++++++++++++++++++++++++++++++++++++++------- 1 file changed, 55 insertions(+), 8 deletions(-) (limited to 'src/rag/loader.rs') diff --git a/src/rag/loader.rs b/src/rag/loader.rs index 4ab0372..9f70c8b 100644 --- a/src/rag/loader.rs +++ b/src/rag/loader.rs @@ -2,12 +2,24 @@ use super::*; use anyhow::{bail, Context, Result}; use async_recursion::async_recursion; +use lazy_static::lazy_static; use serde_json::Value; -use std::{collections::HashMap, env, fs::read_to_string, path::Path}; +use std::{collections::HashMap, env, path::Path, time::Duration}; +use tokio::io::AsyncWriteExt; pub const RECURSIVE_URL_LOADER: &str = "recursive_url"; +pub const URL_LOADER: &str = "url"; -pub fn load( +lazy_static! { + static ref CLIENT: Result = { + let builder = reqwest::ClientBuilder::new().timeout(Duration::from_secs(30)); + let builder = set_proxy(builder, None)?; + let client = builder.build()?; + Ok(client) + }; +} + +pub async fn load( loaders: &HashMap, path: &str, loader_name: &str, @@ -25,13 +37,19 @@ pub fn load( } else { match loaders.get(loader_name) { Some(loader_command) => load_with_command(path, loader_name, loader_command), - None => load_plain(path, loader_name), + None => { + if loader_name == URL_LOADER { + load_url(loaders, path).await + } else { + load_plain(path, loader_name).await + } + } } } } -fn load_plain(path: &str, loader_name: &str) -> Result> { - let contents = read_to_string(path)?; +async fn load_plain(path: &str, loader_name: &str) -> Result> { + let contents = tokio::fs::read_to_string(path).await?; if loader_name == "json" { if let Some(documents) = parse_json_documents(&contents) { return Ok(documents); @@ -42,6 +60,35 @@ fn load_plain(path: &str, loader_name: &str) -> Result> { Ok(vec![document]) } +async fn load_url(loaders: &HashMap, path: &str) -> Result> { + let client = match *CLIENT { + Ok(ref client) => client, + Err(ref err) => bail!("{err}"), + }; + let mut res = client.get(path).send().await?; + let loader_name = path + .rsplit_once('/') + .and_then(|(_, pair)| pair.rsplit_once('.').map(|(_, ext)| ext)) + .unwrap_or("txt"); + let contents = match loaders.get(loader_name) { + Some(loader_command) => { + let save_path = env::temp_dir() + .join(format!("aichat-download-{}.{loader_name}", sha256(path))) + .display() + .to_string(); + let mut save_file = tokio::fs::File::create(&save_path).await?; + while let Some(chunk) = res.chunk().await? { + save_file.write_all(&chunk).await?; + } + run_loader_command(&save_path, loader_name, loader_command)? + } + None => res.text().await?, + }; + let mut document = RagDocument::new(contents); + document.metadata.insert("path".into(), path.to_string()); + Ok(vec![document]) +} + fn load_with_command( path: &str, loader_name: &str, @@ -59,7 +106,7 @@ fn run_loader_command(path: &str, loader_name: &str, loader_command: &str) -> Re })?; let mut use_stdout = true; let outpath = env::temp_dir() - .join(format!("aichat-{}", sha256(path))) + .join(format!("aichat-output-{}", sha256(path))) .display() .to_string(); let cmd_args: Vec<_> = cmd_args @@ -100,8 +147,8 @@ fn run_loader_command(path: &str, loader_name: &str, loader_command: &str) -> Re 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")?; + let contents = std::fs::read_to_string(&outpath) + .context("Failed to read file generated by the loader")?; Ok(contents) } } -- cgit v1.2.3