diff options
| -rw-r--r-- | Cargo.lock | 1 | ||||
| -rw-r--r-- | Cargo.toml | 1 | ||||
| -rw-r--r-- | config.example.yaml | 7 | ||||
| -rw-r--r-- | src/client/common.rs | 30 | ||||
| -rw-r--r-- | src/config/input.rs | 13 | ||||
| -rw-r--r-- | src/config/mod.rs | 16 | ||||
| -rw-r--r-- | src/rag/loader.rs | 63 | ||||
| -rw-r--r-- | src/rag/mod.rs | 3 | ||||
| -rw-r--r-- | src/utils/mod.rs | 23 |
9 files changed, 104 insertions, 53 deletions
@@ -71,7 +71,6 @@ dependencies = [ "json-patch", "lazy_static", "log", - "mime_guess", "nu-ansi-term 0.50.0", "parking_lot", "path-absolutize", @@ -42,7 +42,6 @@ reqwest-eventsource = "0.6.0" simplelog = "0.12.1" log = "0.4.20" shell-words = "1.1.0" -mime_guess = "2.0.4" sha2 = "0.10.8" unicode-width = "0.1.11" async-recursion = "1.1.1" diff --git a/config.example.yaml b/config.example.yaml index 262f41b..c7c5375 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -51,11 +51,12 @@ rag_min_score_keyword_search: 0 # Specifies the minimum relevance sc rag_min_score_rerank: 0 # Specifies the minimum relevance score for reranking # Defines document loaders rag_document_loaders: - # You can add more loaders, here is the syntax: - # <file-extension>: <command-to-load-the-file> + # You can add custom loaders using the following syntax: + # <file-extension>: <command-to-load-the-file> + # Note: Use `$1` for input filepath and `$2` for output filepath. If `$2` is not provided, output to stdout. pdf: 'pdftotext $1 -' # Load .pdf file, see https://poppler.freedesktop.org docx: 'pandoc --to plain $1' # Load .docx file - url: 'curl -fsSL $1' # Load url + # xlsx: 'ssconvert $1 $2' # Load .xlsx file # recursive_url: 'crawler $1 $2' # Load websites # Defines the query structure using variables like __CONTEXT__ and __INPUT__ to tailor searches to specific needs diff --git a/src/client/common.rs b/src/client/common.rs index 265ab0a..f25560d 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -4,10 +4,7 @@ use crate::{ config::{GlobalConfig, Input}, function::{eval_tool_calls, FunctionDeclaration, ToolCall, ToolResult}, render::{render_error, render_stream}, - utils::{ - prompt_input_integer, prompt_input_string, tokenize, watch_abort_signal, AbortSignal, - PromptKind, - }, + utils::*, }; use anyhow::{bail, Context, Result}; @@ -15,10 +12,10 @@ use async_trait::async_trait; use fancy_regex::Regex; use indexmap::IndexMap; use lazy_static::lazy_static; -use reqwest::{Client as ReqwestClient, ClientBuilder, Proxy, RequestBuilder}; +use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; -use std::{env, future::Future, time::Duration}; +use std::{future::Future, time::Duration}; use tokio::sync::mpsc::unbounded_channel; const MODELS_YAML: &str = include_str!("../../models.yaml"); @@ -340,7 +337,7 @@ pub trait Client: Sync + Send { let extra = self.extra_config(); let timeout = extra.and_then(|v| v.connect_timeout).unwrap_or(10); let proxy = extra.and_then(|v| v.proxy.clone()); - builder = set_proxy(builder, &proxy)?; + builder = set_proxy(builder, proxy.as_ref())?; let client = builder .connect_timeout(Duration::from_secs(timeout)) .build() @@ -771,22 +768,3 @@ fn to_json(kind: &PromptKind, value: &str) -> Value { }, } } - -fn set_proxy(builder: ClientBuilder, proxy: &Option<String>) -> Result<ClientBuilder> { - let proxy = if let Some(proxy) = proxy { - if proxy.is_empty() || proxy == "-" { - return Ok(builder); - } - proxy.clone() - } else if let Some(proxy) = ["HTTPS_PROXY", "https_proxy", "ALL_PROXY", "all_proxy"] - .into_iter() - .find_map(|v| env::var(v).ok()) - { - proxy - } else { - return Ok(builder); - }; - let builder = - builder.proxy(Proxy::all(&proxy).with_context(|| format!("Invalid proxy `{proxy}`"))?); - Ok(builder) -} diff --git a/src/config/input.rs b/src/config/input.rs index a41b1c7..3b15d1f 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -10,7 +10,6 @@ use crate::utils::{base64_encode, sha256, AbortSignal}; use anyhow::{bail, Context, Result}; use fancy_regex::Regex; use lazy_static::lazy_static; -use mime_guess::from_path; use std::{ collections::HashMap, fs::File, @@ -407,8 +406,16 @@ fn is_image_ext(path: &Path) -> bool { fn read_media_to_data_url<P: AsRef<Path>>(image_path: P) -> Result<String> { let image_path = image_path.as_ref(); - - let mime_type = from_path(image_path).first_or_octet_stream().to_string(); + let mime_type = match image_path.extension().and_then(|v| v.to_str()) { + Some(extension) => match extension { + "png" => "image/png", + "jpg" | "jpeg" => "image/jpeg", + "webp" => "image/webp", + "gif" => "image/gif", + _ => bail!("Unsupported media type"), + }, + None => bail!("Unknown media type"), + }; let mut file = File::open(image_path)?; let mut buffer = Vec::new(); file.read_to_end(&mut buffer)?; diff --git a/src/config/mod.rs b/src/config/mod.rs index 69609f4..ff337c0 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1434,16 +1434,12 @@ impl Config { } fn setup_rag_document_loaders(&mut self) { - [ - ("pdf", "pdftotext $1 -"), - ("docx", "pandoc --to plain $1"), - ("url", "curl -fsSL $1"), - ] - .into_iter() - .for_each(|(k, v)| { - let (k, v) = (k.to_string(), v.to_string()); - self.rag_document_loaders.entry(k).or_insert(v); - }); + [("pdf", "pdftotext $1 -"), ("docx", "pandoc --to plain $1")] + .into_iter() + .for_each(|(k, v)| { + let (k, v) = (k.to_string(), v.to_string()); + self.rag_document_loaders.entry(k).or_insert(v); + }); } } 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) } } diff --git a/src/rag/mod.rs b/src/rag/mod.rs index f01fc0c..9e37e6f 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -237,7 +237,7 @@ impl Rag { if let Some(path) = path.strip_suffix("**") { new_paths.push((path.to_string(), RECURSIVE_URL_LOADER.into())); } else { - new_paths.push((path.to_string(), "url".into())) + new_paths.push((path.to_string(), URL_LOADER.into())) } } else { let path = Path::new(path); @@ -275,6 +275,7 @@ impl Rag { for (index, (path, loader_name)) in new_paths.into_iter().enumerate() { println!("Loading {path} [{}/{new_paths_len}]", index + 1); let documents = load(&loaders, &path, &loader_name) + .await .with_context(|| format!("Failed to load '{path}'"))?; let separator = get_separators(&loader_name); let splitter = RecursiveCharacterTextSplitter::new( diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 7d1d07a..76ab208 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -14,6 +14,7 @@ pub use self::prompt_input::*; pub use self::render_prompt::render_prompt; pub use self::spinner::{create_spinner, Spinner}; +use anyhow::{Context, Result}; use fancy_regex::Regex; use is_terminal::IsTerminal; use lazy_static::lazy_static; @@ -180,6 +181,28 @@ pub fn safe_join_path<T1: AsRef<Path>, T2: AsRef<Path>>( } } +pub fn set_proxy( + builder: reqwest::ClientBuilder, + proxy: Option<&String>, +) -> Result<reqwest::ClientBuilder> { + let proxy = if let Some(proxy) = proxy { + if proxy.is_empty() || proxy == "-" { + return Ok(builder); + } + proxy.clone() + } else if let Some(proxy) = ["HTTPS_PROXY", "https_proxy", "ALL_PROXY", "all_proxy"] + .into_iter() + .find_map(|v| env::var(v).ok()) + { + proxy + } else { + return Ok(builder); + }; + let builder = builder + .proxy(reqwest::Proxy::all(&proxy).with_context(|| format!("Invalid proxy `{proxy}`"))?); + Ok(builder) +} + #[cfg(test)] mod tests { use super::*; |
