summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2025-02-08 08:46:34 +0800
committerGitHub <noreply@github.com>2025-02-08 08:46:34 +0800
commiteb527ea758b49adbc1e311638cc6999190ff1124 (patch)
treee2396eb8c92adb7cd236bbe1b6c0fa4ec26f6d05
parenta491820ab5ff9bb58e62fd25d88f31ab53a81db9 (diff)
downloadaichat-eb527ea758b49adbc1e311638cc6999190ff1124.tar.gz
feat: enhance `.file` for loading resources from diverse sources (#1155)
-rw-r--r--src/config/input.rs138
-rw-r--r--src/rag/mod.rs55
-rw-r--r--src/repl/mod.rs7
-rw-r--r--src/utils/command.rs5
-rw-r--r--src/utils/loader.rs47
-rw-r--r--src/utils/path.rs15
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();