summaryrefslogtreecommitdiffstats
path: root/src/rag/loader.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/rag/loader.rs')
-rw-r--r--src/rag/loader.rs63
1 files changed, 55 insertions, 8 deletions
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<reqwest::Client> = {
+ 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<String, String>,
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<Vec<RagDocument>> {
- let contents = read_to_string(path)?;
+async fn load_plain(path: &str, loader_name: &str) -> Result<Vec<RagDocument>> {
+ 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<Vec<RagDocument>> {
Ok(vec![document])
}
+async fn load_url(loaders: &HashMap<String, String>, path: &str) -> Result<Vec<RagDocument>> {
+ 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)
}
}