summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/config/input.rs18
-rw-r--r--src/rag/mod.rs25
-rw-r--r--src/rag/splitter/mod.rs6
-rw-r--r--src/utils/loader.rs (renamed from src/rag/loader.rs)61
-rw-r--r--src/utils/mod.rs2
5 files changed, 55 insertions, 57 deletions
diff --git a/src/config/input.rs b/src/config/input.rs
index fa6f00c..42bb8ec 100644
--- a/src/config/input.rs
+++ b/src/config/input.rs
@@ -90,7 +90,9 @@ impl Input {
texts.push(String::new());
}
for (path, contents) in files {
- texts.push(format!("====== PATH: {path} ======\n{contents}\n"));
+ texts.push(format!(
+ "============ PATH: {path} ============\n\n{contents}\n"
+ ));
}
let (role, with_session, with_agent) = resolve_role(&config.read(), role);
Ok(Self {
@@ -392,9 +394,10 @@ async fn load_documents(
data_urls.insert(sha256(&data_url), file_path);
medias.push(data_url)
} else {
- let text = read_file(&file_path)
+ let document = load_file(&loaders, &file_path)
+ .await
.with_context(|| format!("Unable to read file '{file_path}'"))?;
- files.push((file_path, text));
+ files.push((file_path, document.contents));
}
}
for file_url in remote_urls {
@@ -459,12 +462,3 @@ fn read_media_to_data_url(image_path: &str) -> Result<String> {
Ok(data_url)
}
-
-fn read_file<P: AsRef<Path>>(file_path: P) -> Result<String> {
- let file_path = file_path.as_ref();
-
- let mut text = String::new();
- let mut file = File::open(file_path)?;
- file.read_to_string(&mut text)?;
- Ok(text)
-}
diff --git a/src/rag/mod.rs b/src/rag/mod.rs
index 5760871..03b8f7e 100644
--- a/src/rag/mod.rs
+++ b/src/rag/mod.rs
@@ -1,11 +1,9 @@
-use self::loader::*;
use self::splitter::*;
use crate::client::*;
use crate::config::*;
use crate::utils::*;
-mod loader;
mod serde_vectors;
mod splitter;
@@ -367,7 +365,7 @@ impl Rag {
}
}
- let mut files = vec![];
+ let mut loaded_documents = vec![];
let mut has_error = false;
let mut index = 0;
let total = recursive_urls.len() + urls.len() + local_paths.len();
@@ -379,7 +377,7 @@ impl Rag {
index += 1;
println!("Load {start_url}** [{index}/{total}]");
match load_recursive_url(&loaders, &start_url).await {
- Ok(v) => files.extend(v),
+ Ok(v) => loaded_documents.extend(v),
Err(err) => handle_error(err, &mut has_error),
}
}
@@ -387,7 +385,7 @@ impl Rag {
index += 1;
println!("Load {url} [{index}/{total}]");
match load_url(&loaders, &url).await {
- Ok(v) => files.push(v),
+ Ok(v) => loaded_documents.push(v),
Err(err) => handle_error(err, &mut has_error),
}
}
@@ -395,7 +393,7 @@ impl Rag {
index += 1;
println!("Load {local_path} [{index}/{total}]");
match load_file(&loaders, &local_path).await {
- Ok(v) => files.push(v),
+ Ok(v) => loaded_documents.push(v),
Err(err) => handle_error(err, &mut has_error),
}
}
@@ -414,11 +412,12 @@ impl Rag {
}
let mut rag_files = vec![];
- for (contents, mut metadata) in files {
- let path = match metadata.swap_remove(PATH_METADATA) {
- Some(v) => v,
- None => continue,
- };
+ for LoadedDocument {
+ path,
+ contents,
+ mut metadata,
+ } in loaded_documents
+ {
let hash = sha256(&contents);
if let Some(file_ids) = to_deleted.get_mut(&hash) {
if let Some((i, _)) = file_ids
@@ -793,7 +792,7 @@ pub struct RagFile {
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct RagDocument {
pub page_content: String,
- pub metadata: RagMetadata,
+ pub metadata: DocumentMetadata,
}
impl RagDocument {
@@ -814,8 +813,6 @@ impl Default for RagDocument {
}
}
-pub type RagMetadata = IndexMap<String, String>;
-
pub type FileId = usize;
#[derive(Clone, Copy, Hash, Eq, PartialEq, Ord, PartialOrd)]
diff --git a/src/rag/splitter/mod.rs b/src/rag/splitter/mod.rs
index 351c560..a4a9167 100644
--- a/src/rag/splitter/mod.rs
+++ b/src/rag/splitter/mod.rs
@@ -2,7 +2,7 @@ mod language;
pub use self::language::*;
-use super::{RagDocument, RagMetadata};
+use super::{DocumentMetadata, RagDocument};
pub const DEFAULT_SEPARATES: [&str; 4] = ["\n\n", "\n", " ", ""];
@@ -75,7 +75,7 @@ impl RecursiveCharacterTextSplitter {
chunk_header_options: &SplitterChunkHeaderOptions,
) -> Vec<RagDocument> {
let mut texts: Vec<String> = Vec::new();
- let mut metadatas: Vec<RagMetadata> = Vec::new();
+ let mut metadatas: Vec<DocumentMetadata> = Vec::new();
documents.iter().for_each(|d| {
if !d.page_content.is_empty() {
texts.push(d.page_content.clone());
@@ -89,7 +89,7 @@ impl RecursiveCharacterTextSplitter {
pub fn create_documents(
&self,
texts: &[String],
- metadatas: &[RagMetadata],
+ metadatas: &[DocumentMetadata],
chunk_header_options: &SplitterChunkHeaderOptions,
) -> Vec<RagDocument> {
let SplitterChunkHeaderOptions {
diff --git a/src/rag/loader.rs b/src/utils/loader.rs
index 658aa2f..519563c 100644
--- a/src/rag/loader.rs
+++ b/src/utils/loader.rs
@@ -1,15 +1,34 @@
use super::*;
use anyhow::{Context, Result};
+use indexmap::IndexMap;
use std::collections::HashMap;
pub const EXTENSION_METADATA: &str = "__extension__";
-pub const PATH_METADATA: &str = "__path__";
+
+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<(String, RagMetadata)>> {
+) -> Result<Vec<LoadedDocument>> {
let extension = RECURSIVE_URL_LOADER;
let pages: Vec<Page> = match loaders.get(extension) {
Some(loader_command) => {
@@ -25,19 +44,15 @@ pub async fn load_recursive_url(
.into_iter()
.map(|v| {
let Page { path, text } = v;
- let mut metadata: RagMetadata = Default::default();
- metadata.insert(PATH_METADATA.into(), path);
+ let mut metadata: DocumentMetadata = Default::default();
metadata.insert(EXTENSION_METADATA.into(), "md".into());
- (text, metadata)
+ LoadedDocument::new(path, text, metadata)
})
.collect();
Ok(output)
}
-pub async fn load_file(
- loaders: &HashMap<String, String>,
- path: &str,
-) -> Result<(String, RagMetadata)> {
+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),
@@ -45,33 +60,23 @@ pub async fn load_file(
}
}
-pub async fn load_url(
- loaders: &HashMap<String, String>,
- path: &str,
-) -> Result<(String, RagMetadata)> {
+pub async fn load_url(loaders: &HashMap<String, String>, path: &str) -> Result<LoadedDocument> {
let (contents, extension) = fetch(loaders, path, false).await?;
- let mut metadata: RagMetadata = Default::default();
- metadata.insert(PATH_METADATA.into(), path.into());
+ let mut metadata: DocumentMetadata = Default::default();
metadata.insert(EXTENSION_METADATA.into(), extension);
- Ok((contents, metadata))
+ Ok(LoadedDocument::new(path.into(), contents, metadata))
}
-async fn load_plain(path: &str, extension: &str) -> Result<(String, RagMetadata)> {
+async fn load_plain(path: &str, extension: &str) -> Result<LoadedDocument> {
let contents = tokio::fs::read_to_string(path).await?;
- let mut metadata: RagMetadata = Default::default();
- metadata.insert(PATH_METADATA.into(), path.to_string());
+ let mut metadata: DocumentMetadata = Default::default();
metadata.insert(EXTENSION_METADATA.into(), extension.to_string());
- Ok((contents, metadata))
+ Ok(LoadedDocument::new(path.into(), contents, metadata))
}
-fn load_with_command(
- path: &str,
- extension: &str,
- loader_command: &str,
-) -> Result<(String, RagMetadata)> {
+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: RagMetadata = Default::default();
- metadata.insert(PATH_METADATA.into(), path.to_string());
+ let mut metadata: DocumentMetadata = Default::default();
metadata.insert(EXTENSION_METADATA.into(), DEFAULT_EXTENSION.to_string());
- Ok((contents, metadata))
+ 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;