summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--Cargo.lock1
-rw-r--r--Cargo.toml1
-rw-r--r--config.example.yaml7
-rw-r--r--src/client/common.rs30
-rw-r--r--src/config/input.rs13
-rw-r--r--src/config/mod.rs16
-rw-r--r--src/rag/loader.rs63
-rw-r--r--src/rag/mod.rs3
-rw-r--r--src/utils/mod.rs23
9 files changed, 104 insertions, 53 deletions
diff --git a/Cargo.lock b/Cargo.lock
index efb3c76..5dafdc2 100644
--- a/Cargo.lock
+++ b/Cargo.lock
@@ -71,7 +71,6 @@ dependencies = [
"json-patch",
"lazy_static",
"log",
- "mime_guess",
"nu-ansi-term 0.50.0",
"parking_lot",
"path-absolutize",
diff --git a/Cargo.toml b/Cargo.toml
index c99b2a5..f5f788e 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -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::*;