summaryrefslogtreecommitdiffstats
path: root/src/utils
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-12-04 06:56:25 +0800
committerGitHub <noreply@github.com>2024-12-04 06:56:25 +0800
commit5efb44c644cf211f64e241ebf961abc36d9a51e1 (patch)
tree530e2e89716206107a75ec6b189dbdb19c57469a /src/utils
parente86a9f3ee93d6ea0d19e0af04c5a63b41ca72582 (diff)
downloadaichat-5efb44c644cf211f64e241ebf961abc36d9a51e1.tar.gz
feat: `--file/.file` accept file formats in document_loaders (#1033)
Diffstat (limited to 'src/utils')
-rw-r--r--src/utils/loader.rs82
-rw-r--r--src/utils/mod.rs2
2 files changed, 84 insertions, 0 deletions
diff --git a/src/utils/loader.rs b/src/utils/loader.rs
new file mode 100644
index 0000000..519563c
--- /dev/null
+++ b/src/utils/loader.rs
@@ -0,0 +1,82 @@
+use super::*;
+
+use anyhow::{Context, Result};
+use indexmap::IndexMap;
+use std::collections::HashMap;
+
+pub const EXTENSION_METADATA: &str = "__extension__";
+
+pub type DocumentMetadata = IndexMap<String, String>;
+
+#[derive(Debug, Clone)]
+pub struct LoadedDocument {
+ pub path: String,
+ pub contents: String,
+ pub metadata: DocumentMetadata,
+}
+
+impl LoadedDocument {
+ pub fn new(path: String, contents: String, metadata: DocumentMetadata) -> Self {
+ Self {
+ path,
+ contents,
+ metadata,
+ }
+ }
+}
+
+pub async fn load_recursive_url(
+ loaders: &HashMap<String, String>,
+ path: &str,
+) -> Result<Vec<LoadedDocument>> {
+ let extension = RECURSIVE_URL_LOADER;
+ let pages: Vec<Page> = match loaders.get(extension) {
+ Some(loader_command) => {
+ let contents = run_loader_command(path, extension, loader_command)?;
+ serde_json::from_str(&contents).context(r#"The crawler response is invalid. It should follow the JSON format: `[{"path":"...", "text":"..."}]`."#)?
+ }
+ None => {
+ let options = CrawlOptions::preset(path);
+ crawl_website(path, options).await?
+ }
+ };
+ let output = pages
+ .into_iter()
+ .map(|v| {
+ let Page { path, text } = v;
+ let mut metadata: DocumentMetadata = Default::default();
+ metadata.insert(EXTENSION_METADATA.into(), "md".into());
+ LoadedDocument::new(path, text, metadata)
+ })
+ .collect();
+ Ok(output)
+}
+
+pub async fn load_file(loaders: &HashMap<String, String>, path: &str) -> Result<LoadedDocument> {
+ let extension = get_patch_extension(path).unwrap_or_else(|| DEFAULT_EXTENSION.into());
+ match loaders.get(&extension) {
+ Some(loader_command) => load_with_command(path, &extension, loader_command),
+ None => load_plain(path, &extension).await,
+ }
+}
+
+pub async fn load_url(loaders: &HashMap<String, String>, path: &str) -> Result<LoadedDocument> {
+ let (contents, extension) = fetch(loaders, path, false).await?;
+ let mut metadata: DocumentMetadata = Default::default();
+ metadata.insert(EXTENSION_METADATA.into(), extension);
+ Ok(LoadedDocument::new(path.into(), contents, metadata))
+}
+
+async fn load_plain(path: &str, extension: &str) -> Result<LoadedDocument> {
+ let contents = tokio::fs::read_to_string(path).await?;
+ let mut metadata: DocumentMetadata = Default::default();
+ metadata.insert(EXTENSION_METADATA.into(), extension.to_string());
+ Ok(LoadedDocument::new(path.into(), contents, metadata))
+}
+
+fn load_with_command(path: &str, extension: &str, loader_command: &str) -> Result<LoadedDocument> {
+ let contents = run_loader_command(path, extension, loader_command)?;
+ let mut metadata: DocumentMetadata = Default::default();
+ metadata.insert(EXTENSION_METADATA.into(), DEFAULT_EXTENSION.to_string());
+ Ok(LoadedDocument::new(path.into(), contents, metadata))
+}
diff --git a/src/utils/mod.rs b/src/utils/mod.rs
index 5271571..c880252 100644
--- a/src/utils/mod.rs
+++ b/src/utils/mod.rs
@@ -3,6 +3,7 @@ mod clipboard;
mod command;
mod crypto;
mod html_to_md;
+mod loader;
mod path;
mod prompt_input;
mod render_prompt;
@@ -15,6 +16,7 @@ pub use self::clipboard::set_text;
pub use self::command::*;
pub use self::crypto::*;
pub use self::html_to_md::*;
+pub use self::loader::*;
pub use self::path::*;
pub use self::prompt_input::*;
pub use self::render_prompt::render_prompt;