diff options
| author | sigoden <sigoden@gmail.com> | 2024-07-01 20:19:41 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-07-01 20:19:41 +0800 |
| commit | 7c6dac061b2a33aae334c08b725e02459c5e9ade (patch) | |
| tree | 3c2067aaf80a24f988600b42aa12b647675bc383 /src/rag/mod.rs | |
| parent | d193950d204ed82bdb3d6a3a111511d189f084bb (diff) | |
| download | aichat-7c6dac061b2a33aae334c08b725e02459c5e9ade.tar.gz | |
feat: support `.rebuild rag` repl command(#672)
Diffstat (limited to 'src/rag/mod.rs')
| -rw-r--r-- | src/rag/mod.rs | 265 |
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/**") |
