summaryrefslogtreecommitdiffstats
path: root/src/rag
diff options
context:
space:
mode:
Diffstat (limited to 'src/rag')
-rw-r--r--src/rag/loader.rs77
-rw-r--r--src/rag/mod.rs25
-rw-r--r--src/rag/splitter/mod.rs6
3 files changed, 14 insertions, 94 deletions
diff --git a/src/rag/loader.rs b/src/rag/loader.rs
deleted file mode 100644
index 658aa2f..0000000
--- a/src/rag/loader.rs
+++ /dev/null
@@ -1,77 +0,0 @@
-use super::*;
-
-use anyhow::{Context, Result};
-use std::collections::HashMap;
-
-pub const EXTENSION_METADATA: &str = "__extension__";
-pub const PATH_METADATA: &str = "__path__";
-
-pub async fn load_recursive_url(
- loaders: &HashMap<String, String>,
- path: &str,
-) -> Result<Vec<(String, RagMetadata)>> {
- 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: RagMetadata = Default::default();
- metadata.insert(PATH_METADATA.into(), path);
- metadata.insert(EXTENSION_METADATA.into(), "md".into());
- (text, metadata)
- })
- .collect();
- Ok(output)
-}
-
-pub async fn load_file(
- loaders: &HashMap<String, String>,
- path: &str,
-) -> Result<(String, RagMetadata)> {
- 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<(String, RagMetadata)> {
- let (contents, extension) = fetch(loaders, path, false).await?;
- let mut metadata: RagMetadata = Default::default();
- metadata.insert(PATH_METADATA.into(), path.into());
- metadata.insert(EXTENSION_METADATA.into(), extension);
- Ok((contents, metadata))
-}
-
-async fn load_plain(path: &str, extension: &str) -> Result<(String, RagMetadata)> {
- let contents = tokio::fs::read_to_string(path).await?;
- let mut metadata: RagMetadata = Default::default();
- metadata.insert(PATH_METADATA.into(), path.to_string());
- metadata.insert(EXTENSION_METADATA.into(), extension.to_string());
- Ok((contents, metadata))
-}
-
-fn load_with_command(
- path: &str,
- extension: &str,
- loader_command: &str,
-) -> Result<(String, RagMetadata)> {
- let contents = run_loader_command(path, extension, loader_command)?;
- let mut metadata: RagMetadata = Default::default();
- metadata.insert(PATH_METADATA.into(), path.to_string());
- metadata.insert(EXTENSION_METADATA.into(), DEFAULT_EXTENSION.to_string());
- Ok((contents, metadata))
-}
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 {