diff options
| author | sigoden <sigoden@gmail.com> | 2025-02-08 08:46:34 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2025-02-08 08:46:34 +0800 |
| commit | eb527ea758b49adbc1e311638cc6999190ff1124 (patch) | |
| tree | e2396eb8c92adb7cd236bbe1b6c0fa4ec26f6d05 | |
| parent | a491820ab5ff9bb58e62fd25d88f31ab53a81db9 (diff) | |
| download | aichat-eb527ea758b49adbc1e311638cc6999190ff1124.tar.gz | |
feat: enhance `.file` for loading resources from diverse sources (#1155)
| -rw-r--r-- | src/config/input.rs | 138 | ||||
| -rw-r--r-- | src/rag/mod.rs | 55 | ||||
| -rw-r--r-- | src/repl/mod.rs | 7 | ||||
| -rw-r--r-- | src/utils/command.rs | 5 | ||||
| -rw-r--r-- | src/utils/loader.rs | 47 | ||||
| -rw-r--r-- | src/utils/path.rs | 15 |
6 files changed, 195 insertions, 72 deletions
diff --git a/src/config/input.rs b/src/config/input.rs index 482b1d5..c16e1e2 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -5,11 +5,11 @@ use crate::client::{ MessageContent, MessageContentPart, MessageContentToolCalls, MessageRole, Model, }; use crate::function::ToolResult; -use crate::utils::{base64_encode, sha256, AbortSignal}; +use crate::utils::{base64_encode, is_loader_protocol, sha256, AbortSignal}; use anyhow::{bail, Context, Result}; -use path_absolutize::Absolutize; -use std::{collections::HashMap, fs::File, io::Read, path::Path}; +use indexmap::IndexSet; +use std::{collections::HashMap, fs::File, io::Read}; use unicode_width::{UnicodeWidthChar, UnicodeWidthStr}; const IMAGE_EXTS: [&str; 5] = ["png", "jpeg", "jpg", "webp", "gif"]; @@ -60,38 +60,19 @@ impl Input { paths: Vec<String>, role: Option<Role>, ) -> Result<Self> { - let mut raw_paths = vec![]; - let mut external_cmds = vec![]; - let mut local_paths = vec![]; - let mut remote_urls = vec![]; + let loaders = config.read().document_loaders.clone(); + let (raw_paths, local_paths, remote_urls, external_cmds, protocol_paths, with_last_reply) = + resolve_paths(&loaders, paths)?; let mut last_reply = None; - let mut with_last_reply = false; - for path in paths { - match resolve_local_path(&path) { - Some(v) => { - if v == "%%" { - with_last_reply = true; - raw_paths.push(v); - } else if v.len() > 2 && v.starts_with('`') && v.ends_with('`') { - external_cmds.push(v[1..v.len() - 1].to_string()); - raw_paths.push(v); - } else { - if let Ok(path) = Path::new(&v).absolutize() { - raw_paths.push(path.display().to_string()); - } - local_paths.push(v); - } - } - None => { - raw_paths.push(path.clone()); - remote_urls.push(path); - } - } - } - let (documents, medias, data_urls) = - load_documents(config, external_cmds, local_paths, remote_urls) - .await - .context("Failed to load files")?; + let (documents, medias, data_urls) = load_documents( + &loaders, + local_paths, + remote_urls, + external_cmds, + protocol_paths, + ) + .await + .context("Failed to load files")?; let mut texts = vec![]; if !raw_text.is_empty() { texts.push(raw_text.to_string()); @@ -409,11 +390,65 @@ fn resolve_role(config: &Config, role: Option<Role>) -> (Role, bool, bool) { } } +type ResolvePathsOutput = ( + Vec<String>, + Vec<String>, + Vec<String>, + Vec<String>, + Vec<String>, + bool, +); + +fn resolve_paths( + loaders: &HashMap<String, String>, + paths: Vec<String>, +) -> Result<ResolvePathsOutput> { + let mut raw_paths = IndexSet::new(); + let mut local_paths = IndexSet::new(); + let mut remote_urls = IndexSet::new(); + let mut external_cmds = IndexSet::new(); + let mut protocol_paths = IndexSet::new(); + let mut with_last_reply = false; + for path in paths { + if path == "%%" { + with_last_reply = true; + raw_paths.insert(path); + } else if path.starts_with('`') && path.len() > 2 && path.ends_with('`') { + external_cmds.insert(path[1..path.len() - 1].to_string()); + raw_paths.insert(path); + } else if is_url(&path) { + if path.strip_suffix("**").is_some() { + bail!("Invalid website '{path}'"); + } + remote_urls.insert(path.clone()); + raw_paths.insert(path); + } else if is_loader_protocol(loaders, &path) { + protocol_paths.insert(path.clone()); + raw_paths.insert(path); + } else { + let resolved_path = resolve_home_dir(&path); + let absolute_path = to_absolute_path(&resolved_path) + .with_context(|| format!("Invalid path '{path}'"))?; + local_paths.insert(resolved_path); + raw_paths.insert(absolute_path); + } + } + Ok(( + raw_paths.into_iter().collect(), + local_paths.into_iter().collect(), + remote_urls.into_iter().collect(), + external_cmds.into_iter().collect(), + protocol_paths.into_iter().collect(), + with_last_reply, + )) +} + async fn load_documents( - config: &GlobalConfig, - external_cmds: Vec<String>, + loaders: &HashMap<String, String>, local_paths: Vec<String>, remote_urls: Vec<String>, + external_cmds: Vec<String>, + protocol_paths: Vec<String>, ) -> Result<( Vec<(&'static str, String, String)>, Vec<String>, @@ -422,6 +457,7 @@ async fn load_documents( let mut files = vec![]; let mut medias = vec![]; let mut data_urls = HashMap::new(); + for cmd in external_cmds { let (success, stdout, stderr) = run_command_with_output(&SHELL.cmd, &[&SHELL.arg, &cmd], None)?; @@ -433,15 +469,14 @@ async fn load_documents( } let local_files = expand_glob_paths(&local_paths, true).await?; - let loaders = config.read().document_loaders.clone(); for file_path in local_files { if is_image(&file_path) { let contents = read_media_to_data_url(&file_path) - .with_context(|| format!("Unable to read media file '{file_path}'"))?; + .with_context(|| format!("Unable to read media '{file_path}'"))?; data_urls.insert(sha256(&contents), file_path); medias.push(contents) } else { - let document = load_file(&loaders, &file_path) + let document = load_file(loaders, &file_path) .await .with_context(|| format!("Unable to read file '{file_path}'"))?; files.push(("FILE", file_path, document.contents)); @@ -449,7 +484,7 @@ async fn load_documents( } for file_url in remote_urls { - let (contents, extension) = fetch_with_loaders(&loaders, &file_url, true) + let (contents, extension) = fetch_with_loaders(loaders, &file_url, true) .await .with_context(|| format!("Failed to load url '{file_url}'"))?; if extension == MEDIA_URL_EXTENSION { @@ -459,6 +494,17 @@ async fn load_documents( files.push(("URL", file_url, contents)); } } + + for protocol_path in protocol_paths { + let documents = load_protocol_path(loaders, &protocol_path) + .with_context(|| format!("Failed to load from '{protocol_path}'"))?; + files.extend( + documents + .into_iter() + .map(|document| ("FROM", document.path, document.contents)), + ); + } + Ok((files, medias, data_urls)) } @@ -474,18 +520,6 @@ pub fn resolve_data_url(data_urls: &HashMap<String, String>, data_url: String) - } } -fn resolve_local_path(path: &str) -> Option<String> { - if is_url(path) { - return None; - } - let new_path = if let (Some(file), Some(home)) = (path.strip_prefix("~/"), dirs::home_dir()) { - home.join(file).display().to_string() - } else { - path.to_string() - }; - Some(new_path) -} - fn is_image(path: &str) -> bool { get_patch_extension(path) .map(|v| IMAGE_EXTS.contains(&v.as_str())) diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 6165969..93a43c8 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -13,7 +13,6 @@ use hnsw_rs::prelude::*; use indexmap::{IndexMap, IndexSet}; use inquire::{required, validator::Validation, Confirm, Select, Text}; use parking_lot::RwLock; -use path_absolutize::Absolutize; use serde::{Deserialize, Serialize}; use serde_json::json; use std::{collections::HashMap, env, fmt::Debug, fs, hash::Hash, path::Path, time::Duration}; @@ -321,8 +320,8 @@ impl Rag { if let Some(spinner) = &spinner { let _ = spinner.set_message(String::new()); } - let (document_paths, mut recursive_urls, mut urls, mut local_paths) = - resolve_paths(paths).await?; + let (document_paths, mut recursive_urls, mut urls, mut protocol_paths, mut local_paths) = + resolve_paths(&loaders, paths).await?; let mut to_deleted: IndexMap<String, Vec<FileId>> = Default::default(); if refresh { for (file_id, file) in &self.data.files { @@ -342,6 +341,13 @@ impl Rag { .into_iter() .filter(|v| !self.data.document_paths.contains(&format!("{v}**"))) .collect(); + let protocol_paths_cloned = protocol_paths.clone(); + let match_protocol_path = + |v: &str| protocol_paths_cloned.iter().any(|root| v.starts_with(root)); + protocol_paths = protocol_paths + .into_iter() + .filter(|v| !self.data.document_paths.contains(v)) + .collect(); for (file_id, file) in &self.data.files { if is_url(&file.path) { if !urls.swap_remove(&file.path) && !match_recursive_url(&file.path) { @@ -350,6 +356,13 @@ impl Rag { .or_default() .push(*file_id); } + } else if is_loader_protocol(&loaders, &file.path) { + if !match_protocol_path(&file.path) { + to_deleted + .entry(file.hash.clone()) + .or_default() + .push(*file_id); + } } else if !local_paths.swap_remove(&file.path) { to_deleted .entry(file.hash.clone()) @@ -362,7 +375,7 @@ impl Rag { let mut loaded_documents = vec![]; let mut has_error = false; let mut index = 0; - let total = recursive_urls.len() + urls.len() + local_paths.len(); + let total = recursive_urls.len() + urls.len() + protocol_paths.len() + local_paths.len(); let handle_error = |error: anyhow::Error, has_error: &mut bool| { println!("{}", warning_text(&format!("⚠️ {error}"))); *has_error = true; @@ -383,6 +396,14 @@ impl Rag { Err(err) => handle_error(err, &mut has_error), } } + for protocol_path in protocol_paths { + index += 1; + println!("Load {protocol_path} [{index}/{total}]"); + match load_protocol_path(&loaders, &protocol_path) { + Ok(v) => loaded_documents.extend(v), + Err(err) => handle_error(err, &mut has_error), + } + } for local_path in local_paths { index += 1; println!("Load {local_path} [{index}/{total}]"); @@ -899,7 +920,7 @@ fn set_chunk_overlay(default_value: usize) -> Result<usize> { fn add_documents() -> Result<Vec<String>> { let text = Text::new("Add documents:") .with_validator(required!("This field is required")) - .with_help_message("e.g. file;dir/;dir/**/*.{md,mdx};solo-url;site-url/**") + .with_help_message("e.g. file;dir/;dir/**/*.{md,mdx};loader:resource;url;website/**") .prompt()?; let paths = text .split(';') @@ -916,16 +937,19 @@ fn add_documents() -> Result<Vec<String>> { } async fn resolve_paths<T: AsRef<str>>( + loaders: &HashMap<String, String>, paths: &[T], ) -> Result<( IndexSet<String>, IndexSet<String>, IndexSet<String>, IndexSet<String>, + IndexSet<String>, )> { let mut document_paths = IndexSet::new(); let mut recursive_urls = IndexSet::new(); let mut urls = IndexSet::new(); + let mut protocol_paths = IndexSet::new(); let mut absolute_paths = vec![]; for path in paths { let path = path.as_ref().trim(); @@ -936,18 +960,25 @@ async fn resolve_paths<T: AsRef<str>>( urls.insert(path.to_string()); } document_paths.insert(path.to_string()); + } else if is_loader_protocol(loaders, path) { + protocol_paths.insert(path.to_string()); + document_paths.insert(path.to_string()); } else { - let absolute_path = Path::new(path) - .absolutize() - .with_context(|| format!("Invalid path '{path}'"))? - .display() - .to_string(); - absolute_paths.push(absolute_path.clone()); + let resolved_path = resolve_home_dir(path); + let absolute_path = to_absolute_path(&resolved_path) + .with_context(|| format!("Invalid path '{path}'"))?; + absolute_paths.push(resolved_path); document_paths.insert(absolute_path); } } let local_paths = expand_glob_paths(&absolute_paths, false).await?; - Ok((document_paths, recursive_urls, urls, local_paths)) + Ok(( + document_paths, + recursive_urls, + urls, + protocol_paths, + local_paths, + )) } fn progress(spinner: &Option<Spinner>, message: String) { diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 780023a..466c93e 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -579,14 +579,15 @@ pub async fn run_repl_command( ask(config, abort_signal.clone(), input, true).await?; } None => println!( - r#"Usage: .file <file|dir|url|%%|cmd>... [-- <text>...] + r#"Usage: .file <file|dir|url|cmd|loader:resource|%%>... [-- <text>...] .file /tmp/file.txt .file src/ Cargo.toml -- analyze .file https://example.com/file.txt -- summarize .file https://example.com/image.png -- recognize text -.file %% -- translate last reply to english -.file `git diff` -- Generate git commit message"# +.file `git diff` -- Generate git commit message +.file jina:https://example.com +.file %% -- translate last reply to english"# ), }, ".continue" => { diff --git a/src/utils/command.rs b/src/utils/command.rs index 32e470d..3008028 100644 --- a/src/utils/command.rs +++ b/src/utils/command.rs @@ -107,9 +107,8 @@ pub fn run_command_with_output<T: AsRef<OsStr>>( } 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 cmd_args = shell_words::split(loader_command) + .with_context(|| anyhow!("Invalid document loader '{extension}': `{loader_command}`"))?; let mut use_stdout = true; let outpath = temp_file("-output-", "").display().to_string(); let cmd_args: Vec<_> = cmd_args diff --git a/src/utils/loader.rs b/src/utils/loader.rs index 2ac671d..b36bf4a 100644 --- a/src/utils/loader.rs +++ b/src/utils/loader.rs @@ -1,17 +1,19 @@ use super::*; -use anyhow::{Context, Result}; +use anyhow::{anyhow, Context, Result}; use indexmap::IndexMap; +use serde::{Deserialize, Serialize}; use std::collections::HashMap; pub const EXTENSION_METADATA: &str = "__extension__"; pub type DocumentMetadata = IndexMap<String, String>; -#[derive(Debug, Clone)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct LoadedDocument { pub path: String, pub contents: String, + #[serde(default)] pub metadata: DocumentMetadata, } @@ -80,3 +82,44 @@ fn load_with_command(path: &str, extension: &str, loader_command: &str) -> Resul metadata.insert(EXTENSION_METADATA.into(), DEFAULT_EXTENSION.to_string()); Ok(LoadedDocument::new(path.into(), contents, metadata)) } + +pub fn is_loader_protocol(loaders: &HashMap<String, String>, path: &str) -> bool { + match path.split_once(':') { + Some((protocol, _)) => loaders.contains_key(protocol), + None => false, + } +} + +pub fn load_protocol_path( + loaders: &HashMap<String, String>, + path: &str, +) -> Result<Vec<LoadedDocument>> { + let (protocol, loader_command, new_path) = path + .split_once(':') + .and_then(|(protocol, path)| { + let loader_command = loaders.get(protocol)?; + Some((protocol, loader_command, path)) + }) + .ok_or_else(|| anyhow!("No document loader for '{}'", path))?; + let contents = run_loader_command(new_path, protocol, loader_command)?; + let output = if let Ok(list) = serde_json::from_str::<Vec<LoadedDocument>>(&contents) { + list.into_iter() + .map(|mut v| { + if v.path.starts_with(path) { + } else if v.path.starts_with(new_path) { + v.path = format!("{}:{}", protocol, v.path); + } else { + v.path = format!("{}/{}", path, v.path); + } + v + }) + .collect() + } else { + vec![LoadedDocument::new( + path.into(), + contents, + Default::default(), + )] + }; + Ok(output) +} diff --git a/src/utils/path.rs b/src/utils/path.rs index bf3d936..20a365e 100644 --- a/src/utils/path.rs +++ b/src/utils/path.rs @@ -2,6 +2,7 @@ use std::path::{Component, Path, PathBuf}; use anyhow::{bail, Result}; use indexmap::IndexSet; +use path_absolutize::Absolutize; pub fn safe_join_path<T1: AsRef<Path>, T2: AsRef<Path>>( base_path: T1, @@ -75,6 +76,20 @@ pub fn get_patch_extension(path: &str) -> Option<String> { .map(|v| v.to_string_lossy().to_lowercase()) } +pub fn to_absolute_path(path: &str) -> Result<String> { + Ok(Path::new(&path).absolutize()?.display().to_string()) +} + +pub fn resolve_home_dir(path: &str) -> String { + let mut path = path.to_string(); + if path.starts_with("~/") || path.starts_with("~\\") { + if let Some(home_dir) = dirs::home_dir() { + path.replace_range(..1, &home_dir.display().to_string()); + } + } + path +} + fn parse_glob(path_str: &str) -> Result<(String, Vec<String>)> { if let Some(start) = path_str.find("/**/*.").or_else(|| path_str.find(r"\**\*.")) { let base_path = path_str[..start].to_string(); |
