From 0afe5fa24b8990991635c7023f0fb89403f55d90 Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 10 Jul 2024 07:27:46 +0800 Subject: feat: `--file/.file` can load dirs (#693) --- src/config/input.rs | 126 +++++++++++++++++++++++++++------------------------- 1 file changed, 65 insertions(+), 61 deletions(-) (limited to 'src/config') diff --git a/src/config/input.rs b/src/config/input.rs index bfc3ea9..ee85c2b 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -10,12 +10,7 @@ use crate::utils::{base64_encode, sha256, AbortSignal}; use anyhow::{bail, Context, Result}; use fancy_regex::Regex; use lazy_static::lazy_static; -use std::{ - collections::HashMap, - fs::File, - io::Read, - path::{Path, PathBuf}, -}; +use std::{collections::HashMap, fs::File, io::Read, path::Path}; use unicode_width::{UnicodeWidthChar, UnicodeWidthStr}; const IMAGE_EXTS: [&str; 5] = ["png", "jpeg", "jpg", "webp", "gif"]; @@ -62,58 +57,29 @@ impl Input { pub async fn from_files( config: &GlobalConfig, text: &str, - files: Vec, + paths: Vec, role: Option, ) -> Result { 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 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) => { - if is_image { - let data_url = read_media_to_data_url(&file_path) - .with_context(|| format!("Unable to read media file '{file_item}'"))?; - data_urls.insert(sha256(&data_url), file_path.display().to_string()); - medias.push(data_url) - } else { - let text = read_file(&file_path) - .with_context(|| format!("Unable to read file '{file_item}'"))?; - if multi_files { - texts.push(format!("`{file_item}`:\n~~~~~~\n{text}\n~~~~~~")); - } else { - texts.push(text); - } - } - } - None => { - if is_image { - medias.push(file_item.to_string()) - } else { - 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); - } - } - } + let ret = load_paths(config, paths).await; + spinner.stop(); + let (files, medias, data_urls) = ret?; + let files_len = files.len(); + if files_len > 0 { + texts.push(String::new()); + } + let is_multi_files = files_len > 1; + for (path, contents) in files { + if is_multi_files { + texts.push(format!("`{path}`:\n\n{contents}\n\n")); + } else { + texts.push(contents); } } - spinner.stop(); - let (role, with_session, with_agent) = resolve_role(&config.read(), role); Ok(Self { config: config.clone(), @@ -383,6 +349,48 @@ fn resolve_role(config: &Config, role: Option) -> (Role, bool, bool) { } } +async fn load_paths( + config: &GlobalConfig, + paths: Vec, +) -> Result<(Vec<(String, String)>, Vec, HashMap)> { + let mut files = vec![]; + let mut medias = vec![]; + let mut data_urls = HashMap::new(); + let loaders = config.read().document_loaders.clone(); + let mut local_paths = vec![]; + let mut remote_urls = vec![]; + for path in paths { + match resolve_local_path(&path) { + Some(v) => local_paths.push(v), + None => remote_urls.push(path), + } + } + let local_files = expand_glob_paths(&local_paths).await?; + for file_path in local_files { + if is_image(&file_path) { + let data_url = read_media_to_data_url(&file_path) + .with_context(|| format!("Unable to read media file '{file_path}'"))?; + data_urls.insert(sha256(&data_url), file_path); + medias.push(data_url) + } else { + let text = read_file(&file_path) + .with_context(|| format!("Unable to read file '{file_path}'"))?; + files.push((file_path, text)); + } + } + for file_url in remote_urls { + if is_image(&file_url) { + medias.push(file_url) + } else { + let (text, _) = fetch(&loaders, &file_url) + .await + .with_context(|| format!("Failed to load url '{file_url}'"))?; + files.push((file_url, text)); + } + } + Ok((files, medias, data_urls)) +} + pub fn resolve_data_url(data_urls: &HashMap, data_url: String) -> String { if data_url.starts_with("data:") { let hash = sha256(&data_url); @@ -395,25 +403,21 @@ pub fn resolve_data_url(data_urls: &HashMap, data_url: String) - } } -fn resolve_local_file(file: &str) -> Option { - if let Ok(true) = URL_RE.is_match(file) { +fn resolve_local_path(path: &str) -> Option { + if let Ok(true) = URL_RE.is_match(path) { return None; } - let path = if let (Some(file), Some(home)) = (file.strip_prefix("~/"), dirs::home_dir()) { + let new_path = if let (Some(file), Some(home)) = (path.strip_prefix("~/"), dirs::home_dir()) { home.join(file) } else { - std::env::current_dir().ok()?.join(file) + std::env::current_dir().ok()?.join(path) }; - Some(path) + Some(new_path.display().to_string()) } -fn is_image_ext(path: &Path) -> bool { - path.extension() - .map(|v| { - IMAGE_EXTS - .iter() - .any(|ext| *ext == v.to_string_lossy().to_lowercase()) - }) +fn is_image(path: &str) -> bool { + path_extension(path) + .map(|v| IMAGE_EXTS.contains(&v.as_str())) .unwrap_or_default() } -- cgit v1.2.3