diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-28 06:24:20 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-28 06:24:20 +0800 |
| commit | 4fbbbd2d991b37ac04b77151ef862de9649bbfec (patch) | |
| tree | 7e4343fb19b8d105b39ad6137d23944d8bffce60 | |
| parent | 10bd71297db11c163f95625080d956469a1d8689 (diff) | |
| download | aichat-4fbbbd2d991b37ac04b77151ef862de9649bbfec.tar.gz | |
feat: `.file`/`--file` support URLs (#665)
| -rw-r--r-- | config.example.yaml | 21 | ||||
| -rw-r--r-- | src/config/input.rs | 23 | ||||
| -rw-r--r-- | src/config/mod.rs | 10 | ||||
| -rw-r--r-- | src/main.rs | 17 | ||||
| -rw-r--r-- | src/rag/loader.rs | 122 | ||||
| -rw-r--r-- | src/rag/mod.rs | 4 | ||||
| -rw-r--r-- | src/repl/mod.rs | 2 | ||||
| -rw-r--r-- | src/utils/command.rs | 57 | ||||
| -rw-r--r-- | src/utils/mod.rs | 2 | ||||
| -rw-r--r-- | src/utils/request.rs | 75 |
10 files changed, 191 insertions, 142 deletions
diff --git a/config.example.yaml b/config.example.yaml index c7c5375..ed4855c 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -25,6 +25,17 @@ summarize_prompt: 'Summarize the discussion briefly in 200 words or less to use # Text prompt used for including the summary of the entire session summary_prompt: 'This is a summary of the chat history as a recap: ' +# Define document loaders to control how `.file`/`--file` and RAG load files of specific formats. +document_loaders: + # 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 + # xlsx: 'ssconvert $1 $2' # Load .xlsx file + # html: 'pandoc --to plain $1' # Load .html file + # recursive_url: 'crawler $1 $2' # Load websites + # ---- function-calling & agent ---- # Controls the function calling feature. For setup instructions, visit https://github.com/sigoden/llm-functions function_calling: true @@ -49,16 +60,6 @@ rag_chunk_overlap: null # Specifies the chunk overlap rag_min_score_vector_search: 0 # Specifies the minimum relevance score for vector-based searching rag_min_score_keyword_search: 0 # Specifies the minimum relevance score for keyword-based searching rag_min_score_rerank: 0 # Specifies the minimum relevance score for reranking -# Defines document loaders -rag_document_loaders: - # 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 - # 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 rag_template: | Use the following context as your learned knowledge, inside <context></context> XML tags. 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) +} |
