summaryrefslogtreecommitdiffstats
path: root/src/rag/mod.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/rag/mod.rs')
-rw-r--r--src/rag/mod.rs265
1 files changed, 156 insertions, 109 deletions
diff --git a/src/rag/mod.rs b/src/rag/mod.rs
index ae208d2..f5f68ae 100644
--- a/src/rag/mod.rs
+++ b/src/rag/mod.rs
@@ -56,13 +56,13 @@ impl Rag {
let mut rag = Self::create(config, name, save_path, data)?;
let mut paths = doc_paths.to_vec();
if paths.is_empty() {
- paths = add_document_paths()?;
+ paths = add_documents()?;
};
debug!("doc paths: {paths:?}");
let loaders = config.read().document_loaders.clone();
let spinner = create_spinner("Starting").await;
tokio::select! {
- ret = rag.add_paths(loaders, &paths, Some(spinner.clone())) => {
+ ret = rag.load_paths(loaders, &paths, Some(spinner.clone())) => {
spinner.stop();
ret?;
}
@@ -103,6 +103,33 @@ impl Rag {
Ok(rag)
}
+ pub async fn rebuild(
+ &mut self,
+ config: &GlobalConfig,
+ save_path: &Path,
+ abort_signal: AbortSignal,
+ ) -> Result<()> {
+ debug!("rebuild rag: {}", self.name);
+ let loaders = config.read().document_loaders.clone();
+ let spinner = create_spinner("Starting").await;
+ let paths = self.data.document_paths.clone();
+ tokio::select! {
+ ret = self.load_paths(loaders, &paths, Some(spinner.clone())) => {
+ spinner.stop();
+ ret?;
+ }
+ _ = watch_abort_signal(abort_signal) => {
+ spinner.stop();
+ bail!("Aborted!")
+ },
+ };
+ if !self.is_temp() {
+ self.save(save_path)?;
+ println!("✨ Saved rag to '{}'", save_path.display());
+ }
+ Ok(())
+ }
+
pub fn config(config: &GlobalConfig) -> Result<(Model, usize, usize)> {
let (embedding_model_id, chunk_size, chunk_overlap) = {
let config = config.read();
@@ -176,12 +203,23 @@ impl Rag {
}
pub fn export(&self) -> Result<String> {
- let files: Vec<_> = self.data.files.iter().map(|v| &v.path).collect();
+ let files: Vec<_> = self
+ .data
+ .files
+ .iter()
+ .map(|(_, v)| {
+ json!({
+ "path": v.path,
+ "num_chunks": v.documents.len(),
+ })
+ })
+ .collect();
let data = json!({
"path": self.path,
"embedding_model": self.embedding_model.id(),
"chunk_size": self.data.chunk_size,
"chunk_overlap": self.data.chunk_overlap,
+ "document_paths": self.data.document_paths,
"files": files,
});
let output = serde_yaml::to_string(&data)
@@ -220,122 +258,103 @@ impl Rag {
Ok(output)
}
- pub async fn add_paths<T: AsRef<str>>(
+ pub async fn load_paths<T: AsRef<str>>(
&mut self,
loaders: HashMap<String, String>,
paths: &[T],
spinner: Option<Spinner>,
) -> Result<()> {
- let mut rag_files = vec![];
+ if let Some(spinner) = &spinner {
+ let _ = spinner.set_message(String::new());
+ }
- // List files
- let mut new_paths = vec![];
- progress(&spinner, "Gathering paths".into());
- for path in paths {
+ let mut document_paths = vec![];
+ let mut files = vec![];
+ let paths_len = paths.len();
+ for (index, path) in paths.iter().enumerate() {
let path = path.as_ref();
+ println!("Load {path} [{}/{paths_len}]", index + 1);
if Self::is_url_path(path) {
if let Some(path) = path.strip_suffix("**") {
- new_paths.push((path.to_string(), RECURSIVE_URL_LOADER.into()));
+ files.extend(load_recrusive_url(&loaders, path).await?);
} else {
- new_paths.push((path.to_string(), URL_LOADER.into()))
+ files.push(load_url(&loaders, path).await?);
}
+ document_paths.push(path.to_string());
} else {
let path = Path::new(path);
- let path = path
- .absolutize()
- .with_context(|| anyhow!("Invalid path '{}'", path.display()))?;
- let path_str = path.display().to_string();
- if self.data.files.iter().any(|v| v.path == path_str) {
- continue;
- }
- let (path_str, suffixes) = parse_glob(&path_str)?;
- let suffixes = if suffixes.is_empty() {
- None
- } else {
- Some(&suffixes)
- };
- let mut file_paths = vec![];
- list_files(&mut file_paths, Path::new(&path_str), suffixes).await?;
- for file_path in file_paths {
- let loader_name = Path::new(&file_path)
- .extension()
- .map(|v| v.to_string_lossy().to_lowercase())
- .unwrap_or_default();
- new_paths.push((file_path, loader_name))
- }
+ let path = path.absolutize()?.display().to_string();
+ files.extend(load_path(&loaders, &path).await?);
+ document_paths.push(path);
}
}
- // Load files
- let new_paths_len = new_paths.len();
- if new_paths_len > 0 {
- if let Some(spinner) = &spinner {
- let _ = spinner.set_message(String::new());
- }
- for (index, (path, extension)) in new_paths.into_iter().enumerate() {
- println!("Loading {path} [{}/{new_paths_len}]", index + 1);
- let documents = load(&loaders, &path, &extension)
- .await
- .with_context(|| format!("Failed to load '{path}'"))?;
- let splitted_documents: Vec<_> = documents
- .into_iter()
- .flat_map(|mut document| {
- let extension = document
- .metadata
- .swap_remove(EXTENSION_METADATA)
- .unwrap_or_else(|| extension.clone());
- let separator = get_separators(&extension);
- let splitter = RecursiveCharacterTextSplitter::new(
- self.data.chunk_size,
- self.data.chunk_overlap,
- &separator,
- );
- let metadata = document
- .metadata
- .iter()
- .map(|(k, v)| format!("{k}: {v}\n"))
- .collect::<Vec<String>>()
- .join("");
- let split_options = SplitterChunkHeaderOptions::default()
- .with_chunk_header(&format!(
- "<document_metadata>\n{metadata}</document_metadata>\n\n"
- ));
- splitter.split_documents(&[document], &split_options)
- })
- .collect();
- let display_path = if extension == RECURSIVE_URL_LOADER {
- format!("{path}**")
- } else {
- path
- };
- rag_files.push(RagFile {
- path: display_path,
- documents: splitted_documents,
- })
- }
+ let mut to_deleted: IndexMap<String, FileId> = Default::default();
+ for (file_id, file) in &self.data.files {
+ to_deleted.insert(file.hash.clone(), *file_id);
}
- if rag_files.is_empty() {
- return Ok(());
+ let mut rag_files = vec![];
+ for (contents, mut metadata) in files {
+ let path = match metadata.swap_remove(PATH_METADATA) {
+ Some(v) => v,
+ None => continue,
+ };
+ let hash = sha256(&contents);
+ if let Some(file_id) = to_deleted.get(&hash) {
+ if self.data.files[file_id].path == path {
+ to_deleted.swap_remove(&hash);
+ continue;
+ }
+ }
+ let extension = metadata
+ .swap_remove(EXTENSION_METADATA)
+ .unwrap_or_else(|| DEFAULT_EXTENSION.into());
+ let separator = get_separators(&extension);
+ let splitter = RecursiveCharacterTextSplitter::new(
+ self.data.chunk_size,
+ self.data.chunk_overlap,
+ &separator,
+ );
+ let split_options = SplitterChunkHeaderOptions::default().with_chunk_header(&format!(
+ "<document_metadata>\npath: {path}</document_metadata>\n\n"
+ ));
+ let document = RagDocument::new(contents);
+ let splitted_documents = splitter.split_documents(&[document], &split_options);
+ rag_files.push(RagFile {
+ hash: hash.clone(),
+ path,
+ documents: splitted_documents,
+ });
}
- // Convert vectors
- let mut vector_ids = vec![];
- let mut texts = vec![];
- for (file_index, file) in rag_files.iter().enumerate() {
- for (document_index, document) in file.documents.iter().enumerate() {
- vector_ids.push(combine_document_id(file_index, document_index));
- texts.push(document.page_content.clone())
+ let mut next_file_id = self.data.next_file_id;
+ let mut files = vec![];
+ let mut document_ids = vec![];
+ let mut embeddings = vec![];
+
+ if !rag_files.is_empty() {
+ let mut texts = vec![];
+ for file in rag_files.into_iter() {
+ for (document_index, document) in file.documents.iter().enumerate() {
+ document_ids.push(combine_document_id(next_file_id, document_index));
+ texts.push(document.page_content.clone())
+ }
+ files.push((next_file_id, file));
+ next_file_id += 1;
}
+
+ let embeddings_data = EmbeddingsData::new(texts, false);
+ embeddings = self
+ .create_embeddings(embeddings_data, spinner.clone())
+ .await?;
}
- let embeddings_data = EmbeddingsData::new(texts, false);
- let embeddings = self
- .create_embeddings(embeddings_data, spinner.clone())
- .await?;
+ self.data.del(to_deleted.values().cloned().collect());
+ self.data.add(next_file_id, files, document_ids, embeddings);
+ self.data.document_paths = document_paths;
- self.data.add(rag_files, vector_ids, embeddings);
- progress(&spinner, "Building vector store".into());
+ progress(&spinner, "Building database".into());
self.hnsw = self.data.build_hnsw();
self.bm25 = self.data.build_bm25();
@@ -485,21 +504,38 @@ impl Rag {
}
}
-#[derive(Debug, Clone, Serialize, Deserialize)]
+#[derive(Clone, Serialize, Deserialize)]
pub struct RagData {
pub embedding_model: String,
pub chunk_size: usize,
pub chunk_overlap: usize,
- pub files: Vec<RagFile>,
+ pub next_file_id: FileId,
+ pub document_paths: Vec<String>,
+ pub files: IndexMap<FileId, RagFile>,
pub vectors: IndexMap<DocumentId, Vec<f32>>,
}
+impl Debug for RagData {
+ fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
+ f.debug_struct("RagData")
+ .field("embedding_model", &self.embedding_model)
+ .field("chunk_size", &self.chunk_size)
+ .field("chunk_overlap", &self.chunk_overlap)
+ .field("next_file_id", &self.next_file_id)
+ .field("document_paths", &self.document_paths)
+ .field("files", &self.files)
+ .finish()
+ }
+}
+
impl RagData {
pub fn new(embedding_model: String, chunk_size: usize, chunk_overlap: usize) -> Self {
Self {
embedding_model,
chunk_size,
chunk_overlap,
+ next_file_id: 0,
+ document_paths: Default::default(),
files: Default::default(),
vectors: Default::default(),
}
@@ -507,19 +543,33 @@ impl RagData {
pub fn get(&self, id: DocumentId) -> Option<&RagDocument> {
let (file_index, document_index) = split_document_id(id);
- let file = self.files.get(file_index)?;
+ let file = self.files.get(&file_index)?;
let document = file.documents.get(document_index)?;
Some(document)
}
+ pub fn del(&mut self, file_ids: Vec<FileId>) {
+ for file_id in file_ids {
+ if let Some(file) = self.files.swap_remove(&file_id) {
+ for (document_index, _) in file.documents.iter().enumerate() {
+ let document_id = combine_document_id(file_id, document_index);
+ self.vectors.swap_remove(&document_id);
+ }
+ }
+ }
+ }
+
pub fn add(
&mut self,
- files: Vec<RagFile>,
- vector_ids: Vec<DocumentId>,
+ next_file_id: FileId,
+ files: Vec<(FileId, RagFile)>,
+ document_ids: Vec<DocumentId>,
embeddings: EmbeddingsOutput,
) {
+ self.next_file_id = next_file_id;
self.files.extend(files);
- self.vectors.extend(vector_ids.into_iter().zip(embeddings));
+ self.vectors
+ .extend(document_ids.into_iter().zip(embeddings));
}
pub fn build_hnsw(&self) -> Hnsw<'static, f32, DistCosine> {
@@ -531,9 +581,9 @@ impl RagData {
pub fn build_bm25(&self) -> BM25<DocumentId> {
let mut corpus = vec![];
- for (file_index, file) in self.files.iter().enumerate() {
+ for (file_index, file) in self.files.iter() {
for (document_index, document) in file.documents.iter().enumerate() {
- let id = combine_document_id(file_index, document_index);
+ let id = combine_document_id(*file_index, document_index);
corpus.push((id, document.page_content.clone()));
}
}
@@ -543,6 +593,7 @@ impl RagData {
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RagFile {
+ hash: String,
path: String,
documents: Vec<RagDocument>,
}
@@ -560,11 +611,6 @@ impl RagDocument {
metadata: IndexMap::new(),
}
}
-
- pub fn with_metadata(mut self, metadata: RagMetadata) -> Self {
- self.metadata = metadata;
- self
- }
}
impl Default for RagDocument {
@@ -578,6 +624,7 @@ impl Default for RagDocument {
pub type RagMetadata = IndexMap<String, String>;
+pub type FileId = usize;
pub type DocumentId = usize;
pub fn combine_document_id(file_index: usize, document_index: usize) -> DocumentId {
@@ -636,7 +683,7 @@ fn set_chunk_overlay(default_value: usize) -> Result<usize> {
value.parse().map_err(|_| anyhow!("Invalid chunk_overlay"))
}
-fn add_document_paths() -> Result<Vec<String>> {
+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;url;sites/**")