summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/config/input.rs23
-rw-r--r--src/config/mod.rs10
-rw-r--r--src/main.rs17
-rw-r--r--src/rag/loader.rs122
-rw-r--r--src/rag/mod.rs4
-rw-r--r--src/repl/mod.rs2
-rw-r--r--src/utils/command.rs57
-rw-r--r--src/utils/mod.rs2
-rw-r--r--src/utils/request.rs75
9 files changed, 180 insertions, 132 deletions
diff --git a/src/config/input.rs b/src/config/input.rs
index 3b15d1f..bfc3ea9 100644
--- a/src/config/input.rs
+++ b/src/config/input.rs
@@ -59,20 +59,25 @@ impl Input {
}
}
- pub fn new(
+ pub async fn from_files(
config: &GlobalConfig,
text: &str,
files: Vec<String>,
role: Option<Role>,
) -> Result<Self> {
- let mut texts = vec![text.to_string()];
+ let mut texts = vec![];
+ if !text.is_empty() {
+ texts.push(text.to_string());
+ };
let mut medias = vec![];
let mut data_urls = HashMap::new();
let files: Vec<_> = files
.iter()
.map(|f| (f, is_image_ext(Path::new(f))))
.collect();
- let include_filepath = files.iter().filter(|(_, is_image)| !*is_image).count() > 1;
+ let multi_files = files.iter().filter(|(_, is_image)| !*is_image).count() > 1;
+ let loaders = config.read().document_loaders.clone();
+ let spinner = create_spinner("Loading files").await;
for (file_item, is_image) in files {
match resolve_local_file(file_item) {
Some(file_path) => {
@@ -84,7 +89,7 @@ impl Input {
} else {
let text = read_file(&file_path)
.with_context(|| format!("Unable to read file '{file_item}'"))?;
- if include_filepath {
+ if multi_files {
texts.push(format!("`{file_item}`:\n~~~~~~\n{text}\n~~~~~~"));
} else {
texts.push(text);
@@ -95,11 +100,19 @@ impl Input {
if is_image {
medias.push(file_item.to_string())
} else {
- bail!("Unable to use remote file '{file_item}");
+ let (text, _) = fetch(&loaders, file_item)
+ .await
+ .with_context(|| format!("Failed to load '{file_item}'"))?;
+ if multi_files {
+ texts.push(format!("`{file_item}`:\n~~~~~~\n{text}\n~~~~~~"));
+ } else {
+ texts.push(text);
+ }
}
}
}
}
+ spinner.stop();
let (role, with_session, with_agent) = resolve_role(&config.read(), role);
Ok(Self {
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 338d7bd..879f511 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -114,7 +114,7 @@ pub struct Config {
pub rag_min_score_keyword_search: f32,
pub rag_min_score_rerank: f32,
#[serde(default)]
- pub rag_document_loaders: HashMap<String, String>,
+ pub document_loaders: HashMap<String, String>,
pub rag_template: Option<String>,
pub highlight: bool,
@@ -174,7 +174,7 @@ impl Default for Config {
rag_min_score_vector_search: 0.0,
rag_min_score_keyword_search: 0.0,
rag_min_score_rerank: 0.0,
- rag_document_loaders: Default::default(),
+ document_loaders: Default::default(),
rag_template: None,
save_session: None,
@@ -230,7 +230,7 @@ impl Config {
config.setup_model()?;
config.setup_highlight();
config.setup_light_theme()?;
- config.setup_rag_document_loaders();
+ config.setup_document_loaders();
Ok(config)
}
@@ -1433,12 +1433,12 @@ impl Config {
Ok(())
}
- fn setup_rag_document_loaders(&mut self) {
+ fn setup_document_loaders(&mut self) {
[("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);
+ self.document_loaders.entry(k).or_insert(v);
});
}
}
diff --git a/src/main.rs b/src/main.rs
index f89ec47..a2ed9ad 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -22,7 +22,10 @@ use crate::config::{
use crate::function::{eval_tool_calls, need_send_tool_results};
use crate::render::{render_error, MarkdownRender};
use crate::repl::Repl;
-use crate::utils::*;
+use crate::utils::{
+ create_abort_signal, create_spinner, detect_shell, extract_block, run_command, AbortSignal,
+ Shell, CODE_BLOCK_RE, IS_STDOUT_TERMINAL,
+};
use anyhow::{bail, Result};
use async_recursion::async_recursion;
@@ -138,7 +141,7 @@ async fn main() -> Result<()> {
if no_input {
bail!("No input");
}
- let input = create_input(&config, text, file)?;
+ let input = create_input(&config, text, file).await?;
let shell = detect_shell();
shell_execute(&config, &shell, input).await?;
return Ok(());
@@ -146,7 +149,7 @@ async fn main() -> Result<()> {
config.write().apply_prelude()?;
if let Err(err) = match no_input {
false => {
- let mut input = create_input(&config, text, file)?;
+ let mut input = create_input(&config, text, file).await?;
input.use_embeddings(abort_signal.clone()).await?;
start_directive(&config, input, cli.no_stream, cli.code, abort_signal).await
}
@@ -298,11 +301,15 @@ fn aggregate_text(text: Option<String>) -> Result<Option<String>> {
Ok(text)
}
-fn create_input(config: &GlobalConfig, text: Option<String>, file: &[String]) -> Result<Input> {
+async fn create_input(
+ config: &GlobalConfig,
+ text: Option<String>,
+ file: &[String],
+) -> Result<Input> {
let input = if file.is_empty() {
Input::from_str(config, &text.unwrap_or_default(), None)
} else {
- Input::new(config, &text.unwrap_or_default(), file.to_vec(), None)?
+ Input::from_files(config, &text.unwrap_or_default(), file.to_vec(), None).await?
};
if input.is_empty() {
bail!("No input");
diff --git a/src/rag/loader.rs b/src/rag/loader.rs
index 99bc489..b9fe298 100644
--- a/src/rag/loader.rs
+++ b/src/rag/loader.rs
@@ -2,24 +2,11 @@ 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, path::Path, time::Duration};
-use tokio::io::AsyncWriteExt;
+use std::{collections::HashMap, path::Path};
-pub const RECURSIVE_URL_LOADER: &str = "recursive_url";
-pub const URL_LOADER: &str = "url";
pub const EXTENSION_METADATA: &str = "__extension__";
-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,
@@ -35,16 +22,16 @@ pub async fn load(
None => vec![RagDocument::new(contents)],
};
Ok(output)
+ } else if extension == URL_LOADER {
+ let (contents, extension) = fetch(loaders, path).await?;
+ let mut metadata: RagMetadata = Default::default();
+ metadata.insert("path".into(), path.into());
+ metadata.insert(EXTENSION_METADATA.into(), extension);
+ Ok(vec![RagDocument::new(contents).with_metadata(metadata)])
} else {
match loaders.get(extension) {
Some(loader_command) => load_with_command(path, extension, loader_command),
- None => {
- if extension == URL_LOADER {
- load_url(loaders, path).await
- } else {
- load_plain(path, extension).await
- }
- }
+ None => load_plain(path, extension).await,
}
}
}
@@ -61,44 +48,6 @@ async fn load_plain(path: &str, extension: &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 mut metadata: RagMetadata = Default::default();
- metadata.insert("path".into(), path.to_string());
-
- let extension = path
- .rsplit_once('/')
- .and_then(|(_, pair)| pair.rsplit_once('.').map(|(_, ext)| ext))
- .unwrap_or("txt");
- let extension = extension.to_lowercase();
- let document = 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?;
- while let Some(chunk) = res.chunk().await? {
- save_file.write_all(&chunk).await?;
- }
- let contents = run_loader_command(&save_path, &extension, loader_command)?;
- metadata.insert(EXTENSION_METADATA.into(), "txt".to_string());
- RagDocument::new(contents).with_metadata(metadata)
- }
- None => {
- let contents = res.text().await?;
- metadata.insert(EXTENSION_METADATA.into(), extension);
- RagDocument::new(contents).with_metadata(metadata)
- }
- };
- Ok(vec![document])
-}
-
fn load_with_command(
path: &str,
extension: &str,
@@ -109,63 +58,10 @@ fn load_with_command(
document.metadata.insert("path".into(), path.to_string());
document
.metadata
- .insert(EXTENSION_METADATA.into(), "txt".to_string());
+ .insert(EXTENSION_METADATA.into(), DEFAULT_EXTENSION.to_string());
Ok(vec![document])
}
-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)
- }
-}
-
fn parse_json_documents(data: &str) -> Option<Vec<RagDocument>> {
let value: Value = serde_json::from_str(data).ok()?;
let items = match value {
diff --git a/src/rag/mod.rs b/src/rag/mod.rs
index 536092f..ae208d2 100644
--- a/src/rag/mod.rs
+++ b/src/rag/mod.rs
@@ -59,7 +59,7 @@ impl Rag {
paths = add_document_paths()?;
};
debug!("doc paths: {paths:?}");
- let loaders = config.read().rag_document_loaders.clone();
+ let loaders = config.read().document_loaders.clone();
let spinner = create_spinner("Starting").await;
tokio::select! {
ret = rag.add_paths(loaders, &paths, Some(spinner.clone())) => {
@@ -641,7 +641,7 @@ fn add_document_paths() -> Result<Vec<String>> {
.with_validator(required!("This field is required"))
.with_help_message("e.g. file;dir/;dir/**/*.md;url;sites/**")
.prompt()?;
- let paths = text.split(';').map(|v| v.to_string()).collect();
+ let paths = text.split(';').map(|v| v.trim().to_string()).collect();
Ok(paths)
}
diff --git a/src/repl/mod.rs b/src/repl/mod.rs
index c09395d..360af70 100644
--- a/src/repl/mod.rs
+++ b/src/repl/mod.rs
@@ -318,7 +318,7 @@ Tips: use <tab> to autocomplete conversation starter text.
Some(args) => {
let (files, text) = split_files_text(args);
let files = shell_words::split(files).with_context(|| "Invalid args")?;
- let input = Input::new(&self.config, text, files, None)?;
+ let input = Input::from_files(&self.config, text, files, None).await?;
ask(&self.config, self.abort_signal.clone(), input, true).await?;
}
None => println!("Usage: .file <files>... [-- <text>...]"),
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)
+}