summaryrefslogtreecommitdiffstats
path: root/src/utils
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-28 06:24:20 +0800
committerGitHub <noreply@github.com>2024-06-28 06:24:20 +0800
commit4fbbbd2d991b37ac04b77151ef862de9649bbfec (patch)
tree7e4343fb19b8d105b39ad6137d23944d8bffce60 /src/utils
parent10bd71297db11c163f95625080d956469a1d8689 (diff)
downloadaichat-4fbbbd2d991b37ac04b77151ef862de9649bbfec.tar.gz
feat: `.file`/`--file` support URLs (#665)
Diffstat (limited to 'src/utils')
-rw-r--r--src/utils/command.rs57
-rw-r--r--src/utils/mod.rs2
-rw-r--r--src/utils/request.rs75
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)
+}