diff options
Diffstat (limited to 'src/utils')
| -rw-r--r-- | src/utils/command.rs | 57 | ||||
| -rw-r--r-- | src/utils/mod.rs | 2 | ||||
| -rw-r--r-- | src/utils/request.rs | 75 |
3 files changed, 133 insertions, 1 deletions
diff --git a/src/utils/command.rs b/src/utils/command.rs index 557649a..1a7de5b 100644 --- a/src/utils/command.rs +++ b/src/utils/command.rs @@ -1,6 +1,8 @@ +use super::*; + use std::{collections::HashMap, env, ffi::OsStr, path::Path, process::Command}; -use anyhow::{Context, Result}; +use anyhow::{anyhow, bail, Context, Result}; pub fn detect_os() -> String { let os = env::consts::OS; @@ -91,6 +93,59 @@ pub fn run_command_with_output<T: AsRef<OsStr>>( Ok((status.success(), stdout.to_string(), stderr.to_string())) } +pub fn run_loader_command(path: &str, extension: &str, loader_command: &str) -> Result<String> { + let cmd_args = shell_words::split(loader_command).with_context(|| { + anyhow!("Invalid rag document loader '{extension}': `{loader_command}`") + })?; + let mut use_stdout = true; + let outpath = env::temp_dir() + .join(format!("aichat-output-{}", sha256(path))) + .display() + .to_string(); + let cmd_args: Vec<_> = cmd_args + .into_iter() + .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]; + 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 status != 0 { + bail!("The command `{cmd_eval}` exited with non-zero.") + } + let contents = std::fs::read_to_string(&outpath) + .context("Failed to read file generated by the loader")?; + Ok(contents) + } +} + pub fn edit_file(editor: &str, path: &Path) -> Result<()> { let mut child = Command::new(editor).arg(path).spawn()?; child.wait()?; diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 76ab208..e36a54b 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -4,6 +4,7 @@ mod command; mod crypto; mod prompt_input; mod render_prompt; +mod request; mod spinner; pub use self::abort_signal::*; @@ -12,6 +13,7 @@ pub use self::command::*; pub use self::crypto::*; pub use self::prompt_input::*; pub use self::render_prompt::render_prompt; +pub use self::request::*; pub use self::spinner::{create_spinner, Spinner}; use anyhow::{Context, Result}; diff --git a/src/utils/request.rs b/src/utils/request.rs new file mode 100644 index 0000000..bc5388c --- /dev/null +++ b/src/utils/request.rs @@ -0,0 +1,75 @@ +use super::*; + +use anyhow::{bail, Result}; +use http::header::CONTENT_TYPE; +use lazy_static::lazy_static; +use std::{collections::HashMap, env, time::Duration}; +use tokio::io::AsyncWriteExt; + +pub const URL_LOADER: &str = "url"; +pub const RECURSIVE_URL_LOADER: &str = "recursive_url"; +pub const DEFAULT_EXTENSION: &str = "txt"; + +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 fetch(loaders: &HashMap<String, String>, path: &str) -> Result<(String, String)> { + if let Some(loader_command) = loaders.get(URL_LOADER) { + let contents = run_loader_command(path, URL_LOADER, loader_command)?; + return Ok((contents, DEFAULT_EXTENSION.into())); + } + let client = match *CLIENT { + Ok(ref client) => client, + Err(ref err) => bail!("{err}"), + }; + let mut res = client.get(path).send().await?; + + let extension = path + .rsplit_once('/') + .and_then(|(_, pair)| pair.rsplit_once('.').map(|(_, ext)| ext)) + .unwrap_or(DEFAULT_EXTENSION); + let mut extension = extension.to_lowercase(); + let content_type = res + .headers() + .get(CONTENT_TYPE) + .and_then(|v| v.to_str().ok()) + .map(|v| match v.split_once(';') { + Some((mime, _)) => mime, + None => v, + }); + if let Some(true) = content_type.map(|v| v.contains("text/html")) { + extension = "html".into() + } + let result = match loaders.get(&extension) { + Some(loader_command) => { + let save_path = env::temp_dir() + .join(format!("aichat-download-{}.{extension}", sha256(path))) + .display() + .to_string(); + let mut save_file = tokio::fs::File::create(&save_path).await?; + let mut size = 0; + while let Some(chunk) = res.chunk().await? { + size += chunk.len(); + save_file.write_all(&chunk).await?; + } + let contents = if size == 0 { + println!("{}", warning_text(&format!("No content at '{path}'"))); + String::new() + } else { + run_loader_command(&save_path, &extension, loader_command)? + }; + (contents, DEFAULT_EXTENSION.into()) + } + None => { + let contents = res.text().await?; + (contents, extension) + } + }; + Ok(result) +} |
