diff options
Diffstat (limited to 'src/rag')
| -rw-r--r-- | src/rag/loader.rs | 146 | ||||
| -rw-r--r-- | src/rag/mod.rs | 425 | ||||
| -rw-r--r-- | src/rag/splitter.rs | 564 |
3 files changed, 1135 insertions, 0 deletions
diff --git a/src/rag/loader.rs b/src/rag/loader.rs new file mode 100644 index 0000000..106802a --- /dev/null +++ b/src/rag/loader.rs @@ -0,0 +1,146 @@ +use super::RagDocument; + +use anyhow::{bail, Context, Result}; +use async_recursion::async_recursion; +use std::{path::Path, process::Command}; +use tokio::fs; + +pub async fn load(path: &str, extension: &str) -> Result<Vec<RagDocument>> { + match extension { + "docx" | "epub" | "ipynb" => load_pandoc(path) + .await + .context("Failed to load with pandoc"), + "pdf" => load_pdf(path).await, + _ => load_plain(path).await, + } +} + +async fn load_plain(path: &str) -> Result<Vec<RagDocument>> { + let contents = fs::read_to_string(path).await?; + let document = RagDocument::new(contents); + Ok(vec![document]) +} + +async fn load_pdf(path: &str) -> Result<Vec<RagDocument>> { + let contents = pdf_extract::extract_text(path)?; + let document = RagDocument::new(contents); + Ok(vec![document]) +} + +async fn load_pandoc(path: &str) -> Result<Vec<RagDocument>> { + let output = Command::new("pandoc") + .arg("--to") + .arg("plain") + .arg(path) + .output()?; + + if !output.status.success() { + let stderr = String::from_utf8_lossy(&output.stderr); + bail!( + "Pandoc conversion failed with exit code {:?}: {}", + output.status.code(), + stderr + ); + } + + let contents = std::str::from_utf8(&output.stdout)?; + let document = RagDocument::new(contents); + Ok(vec![document]) +} + +pub fn parse_glob(path_str: &str) -> Result<(String, Vec<String>)> { + if let Some(start) = path_str.find("/**/*.").or_else(|| path_str.find(r"\**\*.")) { + let base_path = path_str[..start].to_string(); + if let Some(curly_brace_end) = path_str[start..].find('}') { + let end = start + curly_brace_end; + let extensions_str = &path_str[start + 6..end + 1]; + let extensions = if extensions_str.starts_with('{') && extensions_str.ends_with('}') { + extensions_str[1..extensions_str.len() - 1] + .split(',') + .map(|s| s.to_string()) + .collect::<Vec<String>>() + } else { + bail!("Invalid path '{path_str}'"); + }; + Ok((base_path, extensions)) + } else { + let extensions_str = &path_str[start + 6..]; + let extensions = vec![extensions_str.to_string()]; + Ok((base_path, extensions)) + } + } else { + Ok((path_str.to_string(), vec![])) + } +} + +#[async_recursion] +pub async fn list_files( + files: &mut Vec<String>, + entry_path: &Path, + suffixes: Option<&Vec<String>>, +) -> Result<()> { + if !entry_path.exists() { + bail!("Not found: {:?}", entry_path); + } + if entry_path.is_file() { + add_file(files, suffixes, entry_path); + return Ok(()); + } + if !entry_path.is_dir() { + bail!("Not a directory: {:?}", entry_path); + } + let mut reader = fs::read_dir(entry_path).await?; + while let Some(entry) = reader.next_entry().await? { + let path = entry.path(); + if path.is_file() { + add_file(files, suffixes, &path); + } else if path.is_dir() { + list_files(files, &path, suffixes).await?; + } + } + Ok(()) +} + +fn add_file(files: &mut Vec<String>, suffixes: Option<&Vec<String>>, path: &Path) { + if is_valid_extension(suffixes, path) { + files.push(path.display().to_string()); + } +} + +fn is_valid_extension(suffixes: Option<&Vec<String>>, path: &Path) -> bool { + if let Some(suffixes) = suffixes { + if !suffixes.is_empty() { + if let Some(extension) = path.extension().map(|v| v.to_string_lossy().to_string()) { + return suffixes.contains(&extension); + } + return false; + } + } + true +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_parse_glob() { + assert_eq!(parse_glob("dir").unwrap(), ("dir".into(), vec![])); + assert_eq!( + parse_glob("dir/file.md").unwrap(), + ("dir/file.md".into(), vec![]) + ); + assert_eq!( + parse_glob("dir/**/*.md").unwrap(), + ("dir".into(), vec!["md".into()]) + ); + assert_eq!( + parse_glob("dir/**/*.{md,txt}").unwrap(), + ("dir".into(), vec!["md".into(), "txt".into()]) + ); + assert_eq!( + parse_glob("C:\\dir\\**\\*.{md,txt}").unwrap(), + ("C:\\dir".into(), vec!["md".into(), "txt".into()]) + ); + } +} diff --git a/src/rag/mod.rs b/src/rag/mod.rs new file mode 100644 index 0000000..387d3d9 --- /dev/null +++ b/src/rag/mod.rs @@ -0,0 +1,425 @@ +use self::loader::*; +use self::splitter::*; + +use crate::client::*; +use crate::config::*; +use crate::utils::*; + +mod loader; +mod splitter; + +use anyhow::bail; +use anyhow::{anyhow, Context, Result}; +use hnsw_rs::prelude::*; +use indexmap::IndexMap; +use inquire::{required, validator::Validation, Select, Text}; +use path_absolutize::Absolutize; +use serde::{Deserialize, Serialize}; +use serde_json::json; +use std::fmt::Debug; +use std::{io::BufReader, path::Path}; +use tokio::sync::mpsc; + +pub const TEMP_RAG_NAME: &str = "temp"; +pub const CHUNK_OVERLAP: usize = 20; +pub const SIMILARITY_THRESHOLD: f32 = 0.25; + +pub struct Rag { + client: Box<dyn Client>, + name: String, + path: String, + model: Model, + hnsw: Hnsw<'static, f32, DistCosine>, + data: RagData, +} + +impl Debug for Rag { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Rag") + .field("name", &self.name) + .field("path", &self.path) + .field("model", &self.model) + .field("data", &self.data) + .finish() + } +} + +impl Rag { + pub async fn init( + config: &GlobalConfig, + name: &str, + path: &Path, + abort_signal: AbortSignal, + ) -> Result<Self> { + debug!("init rag: {name}"); + let model = select_embedding_model(config)?; + let chunk_size = model.default_chunk_size(); + let chunk_size = set_chunk_size(chunk_size)?; + let data = RagData::new(&model.id(), chunk_size); + let mut rag = Self::create(config, name, path, data)?; + let paths = add_document_paths()?; + debug!("document paths: {paths:?}"); + let (stop_spinner_tx, set_spinner_message_tx) = run_spinner("Starting").await; + tokio::select! { + ret = rag.add_paths(&paths, Some(set_spinner_message_tx)) => { + let _ = stop_spinner_tx.send(()); + ret?; + } + _ = watch_abort_signal(abort_signal) => { + let _ = stop_spinner_tx.send(()); + bail!("Aborted!") + }, + }; + if !rag.is_temp() { + rag.save(path)?; + println!("⨠Saved rag to '{}'", path.display()); + } + Ok(rag) + } + + pub fn load(config: &GlobalConfig, name: &str, path: &Path) -> Result<Self> { + let err = || format!("Failed to load rag '{name}'"); + let file = std::fs::File::open(path).with_context(err)?; + let reader = BufReader::new(file); + let data: RagData = bincode::deserialize_from(reader).with_context(err)?; + Self::create(config, name, path, data) + } + + pub fn create(config: &GlobalConfig, name: &str, path: &Path, data: RagData) -> Result<Self> { + let hnsw = data.build_hnsw(); + let model = retrieve_embedding_model(&config.read(), &data.model)?; + let client = init_client(config, Some(model.clone()))?; + let rag = Rag { + client, + name: name.to_string(), + path: path.display().to_string(), + data, + model, + hnsw, + }; + Ok(rag) + } + + pub fn save(&self, path: &Path) -> Result<()> { + ensure_parent_exists(path)?; + let mut file = std::fs::File::create(path)?; + bincode::serialize_into(&mut file, &self.data) + .with_context(|| format!("Failed to save rag '{}'", self.name))?; + Ok(()) + } + + pub fn export(&self) -> Result<String> { + let files: Vec<_> = self.data.files.iter().map(|v| &v.path).collect(); + let data = json!({ + "path": self.path, + "model": self.model.id(), + "chunk_size": self.data.chunk_size, + "files": files, + }); + let output = serde_yaml::to_string(&data) + .with_context(|| format!("Unable to show info about rag '{}'", self.name))?; + Ok(output) + } + + pub fn name(&self) -> &str { + &self.name + } + + pub fn is_temp(&self) -> bool { + self.name == TEMP_RAG_NAME + } + + pub async fn search( + &self, + text: &str, + top_k: usize, + abort_signal: AbortSignal, + ) -> Result<String> { + let (stop_spinner_tx, _) = run_spinner("Embedding").await; + let ret = tokio::select! { + ret = self.search_impl(text, top_k) => { + ret + } + _ = watch_abort_signal(abort_signal) => { + bail!("Aborted!") + }, + }; + let _ = stop_spinner_tx.send(()); + let output = ret?.join("\n\n"); + Ok(output) + } + + pub async fn add_paths<T: AsRef<Path>>( + &mut self, + paths: &[T], + progress_tx: Option<mpsc::UnboundedSender<String>>, + ) -> Result<()> { + // List files + let mut file_paths = vec![]; + progress(&progress_tx, "Listing paths".into()); + for path in paths { + let path = path + .as_ref() + .absolutize() + .with_context(|| anyhow!("Invalid path '{}'", path.as_ref().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) + }; + list_files(&mut file_paths, Path::new(&path_str), suffixes).await?; + } + + // Load files + let mut rag_files = vec![]; + let file_paths_len = file_paths.len(); + progress(&progress_tx, format!("Loading files [1/{file_paths_len}]")); + for path in file_paths { + let extension = Path::new(&path) + .extension() + .map(|v| v.to_string_lossy().to_lowercase()) + .unwrap_or_default(); + let separator = autodetect_separator(&extension); + let splitter = Splitter::new(self.data.chunk_size, CHUNK_OVERLAP, separator); + let documents = load(&path, &extension) + .await + .with_context(|| format!("Failed to load text at '{path}'"))?; + let documents = + splitter.split_documents(&documents, &SplitterChunkHeaderOptions::default()); + rag_files.push(RagFile { path, documents }); + progress( + &progress_tx, + format!("Loading files [{}/{file_paths_len}]", rag_files.len()), + ); + } + + if rag_files.is_empty() { + return Ok(()); + } + + // Convert vectors + let mut vector_ids = vec![]; + let mut texts = vec![]; + for (file_index, file) in rag_files.iter().enumerate() { + for (document_index, doc) in file.documents.iter().enumerate() { + vector_ids.push(combine_vector_id(file_index, document_index)); + texts.push(doc.page_content.clone()) + } + } + + let embeddings_data = EmbeddingsData::new(texts, false); + let embeddings = self + .create_embeddings(embeddings_data, progress_tx.clone()) + .await?; + + self.data.add(rag_files, vector_ids, embeddings); + progress(&progress_tx, "Building vector store".into()); + self.hnsw = self.data.build_hnsw(); + + Ok(()) + } + + async fn search_impl(&self, text: &str, top_k: usize) -> Result<Vec<String>> { + let splitter = Splitter::new(self.data.chunk_size, CHUNK_OVERLAP, &DEFAULT_SEPARATES); + let texts = splitter.split_text(text); + let embeddings_data = EmbeddingsData::new(texts, true); + let embeddings = self.create_embeddings(embeddings_data, None).await?; + let output = self + .hnsw + .parallel_search(&embeddings, top_k, 30) + .into_iter() + .flat_map(|list| { + list.into_iter() + .filter_map(|v| { + if v.distance < SIMILARITY_THRESHOLD { + return None; + } + let (file_index, document_index) = split_vector_id(v.d_id); + let text = self.data.files[file_index].documents[document_index] + .page_content + .clone(); + Some(text) + }) + .collect::<Vec<_>>() + }) + .collect(); + Ok(output) + } + + async fn create_embeddings( + &self, + data: EmbeddingsData, + progress_tx: Option<mpsc::UnboundedSender<String>>, + ) -> Result<EmbeddingsOutput> { + let EmbeddingsData { texts, query } = data; + let mut output = vec![]; + let chunks = texts.chunks(self.model.max_concurrent_chunks()); + let chunks_len = chunks.len(); + progress( + &progress_tx, + format!("Creating embeddings [1/{chunks_len}]"), + ); + for (index, texts) in chunks.enumerate() { + let chunk_data = EmbeddingsData { + texts: texts.to_vec(), + query, + }; + let chunk_output = self + .client + .embeddings(chunk_data) + .await + .context("Failed to create embedding")?; + output.extend(chunk_output); + progress( + &progress_tx, + format!("Creating embeddings [{}/{chunks_len}]", index + 1), + ); + } + Ok(output) + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RagData { + pub model: String, + pub chunk_size: usize, + pub files: Vec<RagFile>, + pub vectors: IndexMap<VectorID, Vec<f32>>, +} + +impl RagData { + pub fn new(model: &str, chunk_size: usize) -> Self { + Self { + model: model.to_string(), + chunk_size, + files: Default::default(), + vectors: Default::default(), + } + } + + pub fn add( + &mut self, + files: Vec<RagFile>, + vector_ids: Vec<VectorID>, + embeddings: EmbeddingsOutput, + ) { + self.files.extend(files); + self.vectors.extend(vector_ids.into_iter().zip(embeddings)); + } + + pub fn build_hnsw(&self) -> Hnsw<'static, f32, DistCosine> { + let hnsw = Hnsw::new(32, self.vectors.len(), 16, 200, DistCosine {}); + let list: Vec<_> = self.vectors.iter().map(|(k, v)| (v, *k)).collect(); + hnsw.parallel_insert(&list); + hnsw + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RagFile { + path: String, + documents: Vec<RagDocument>, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RagDocument { + pub page_content: String, + pub metadata: RagMetadata, +} + +impl RagDocument { + pub fn new<S: Into<String>>(page_content: S) -> Self { + RagDocument { + page_content: page_content.into(), + metadata: IndexMap::new(), + } + } + + #[allow(unused)] + pub fn with_metadata(mut self, metadata: RagMetadata) -> Self { + self.metadata = metadata; + self + } +} + +impl Default for RagDocument { + fn default() -> Self { + RagDocument { + page_content: "".to_string(), + metadata: IndexMap::new(), + } + } +} + +pub type RagMetadata = IndexMap<String, String>; + +pub type VectorID = usize; + +pub fn combine_vector_id(file_index: usize, document_index: usize) -> VectorID { + file_index << (usize::BITS / 2) | document_index +} + +pub fn split_vector_id(value: VectorID) -> (usize, usize) { + let low_mask = (1 << (usize::BITS / 2)) - 1; + let low = value & low_mask; + let high = value >> (usize::BITS / 2); + (high, low) +} + +fn retrieve_embedding_model(config: &Config, model_id: &str) -> Result<Model> { + let models = list_embedding_models(config); + let model = + Model::find(&models, model_id).ok_or_else(|| anyhow!("No embedding model '{model_id}'"))?; + Ok(model) +} + +fn select_embedding_model(config: &GlobalConfig) -> Result<Model> { + let config = config.read(); + let model = match config.embedding_model.clone() { + Some(model_id) => retrieve_embedding_model(&config, &model_id)?, + None => { + let models = list_embedding_models(&config); + if models.is_empty() { + bail!("No embedding model"); + } + let model_ids: Vec<_> = models.iter().map(|v| v.id()).collect(); + let model_id = Select::new("Select embedding model:", model_ids).prompt()?; + retrieve_embedding_model(&config, &model_id)? + } + }; + Ok(model) +} + +fn set_chunk_size(chunk_size: usize) -> Result<usize> { + let value = Text::new("Set chunk size:") + .with_default(&chunk_size.to_string()) + .with_validator(move |text: &str| { + let out = match text.parse::<usize>() { + Ok(_) => Validation::Valid, + Err(_) => Validation::Invalid("Must be a integer".into()), + }; + Ok(out) + }) + .prompt()?; + value.parse().map_err(|_| anyhow!("Invalid chunk_size")) +} + +fn add_document_paths() -> Result<Vec<String>> { + let text = Text::new("Add document paths:") + .with_validator(required!("This field is required")) + .with_help_message("e.g. file1;dir2/;dir3/**/*.md") + .prompt()?; + let paths = text.split(';').map(|v| v.to_string()).collect(); + Ok(paths) +} + +fn progress(spinner_message_tx: &Option<mpsc::UnboundedSender<String>>, message: String) { + if let Some(tx) = spinner_message_tx { + let _ = tx.send(message); + } +} diff --git a/src/rag/splitter.rs b/src/rag/splitter.rs new file mode 100644 index 0000000..5fdacee --- /dev/null +++ b/src/rag/splitter.rs @@ -0,0 +1,564 @@ +use super::{RagDocument, RagMetadata}; + +use std::cmp::Ordering; + +pub const DEFAULT_SEPARATES: [&str; 4] = ["\n\n", "\n", " ", ""]; +pub const HTML_SEPARATES: [&str; 28] = [ + // First, try to split along HTML tags + "<body>", "<div>", "<p>", "<br>", "<li>", "<h1>", "<h2>", "<h3>", "<h4>", "<h5>", "<h6>", + "<span>", "<table>", "<tr>", "<td>", "<th>", "<ul>", "<ol>", "<header>", "<footer>", "<nav>", + // Head + "<head>", "<style>", "<script>", "<meta>", "<title>", // Normal type of lines + " ", "", +]; +pub const MARKDOWN_SEPARATES: [&str; 13] = [ + // First, try to split along Markdown headings (starting with level 2) + "\n## ", + "\n### ", + "\n#### ", + "\n##### ", + "\n###### ", + // Note the alternative syntax for headings (below) is not handled here + // Heading level 2 + // --------------- + // End of code block + "```\n\n", + // Horizontal lines + "\n\n***\n\n", + "\n\n---\n\n", + "\n\n___\n\n", + // Note that this splitter doesn't handle horizontal lines defined + // by *three or more* of ***, ---, or ___, but this is not handled + "\n\n", + "\n", + " ", + "", +]; +pub const LATEX_SEPARATES: [&str; 19] = [ + // First, try to split along Latex sections + "\n\\chapter{", + "\n\\section{", + "\n\\subsection{", + "\n\\subsubsection{", + // Now split by environments + "\n\\begin{enumerate}", + "\n\\begin{itemize}", + "\n\\begin{description}", + "\n\\begin{list}", + "\n\\begin{quote}", + "\n\\begin{quotation}", + "\n\\begin{verse}", + "\n\\begin{verbatim}", + // Now split by math environments + "\n\\begin{align}", + "$$", + "$", + // Now split by the normal type of lines + "\n\n", + "\n", + " ", + "", +]; + +pub fn autodetect_separator(extension: &str) -> &[&'static str] { + match extension { + "md" | "mkd" => &MARKDOWN_SEPARATES, + "htm" | "html" => &HTML_SEPARATES, + "tex" => &LATEX_SEPARATES, + _ => &DEFAULT_SEPARATES, + } +} + +pub struct Splitter { + pub chunk_size: usize, + pub chunk_overlap: usize, + pub separators: Vec<String>, + pub length_function: Box<dyn Fn(&str) -> usize + Send + Sync>, +} + +impl Default for Splitter { + fn default() -> Self { + Self { + chunk_size: 1000, + chunk_overlap: 20, + separators: DEFAULT_SEPARATES.iter().map(|v| v.to_string()).collect(), + length_function: Box::new(|text| text.len()), + } + } +} + +// Builder pattern for Options struct +impl Splitter { + pub fn new(chunk_size: usize, chunk_overlap: usize, separators: &[&str]) -> Self { + Self::default() + .with_chunk_size(chunk_size) + .with_chunk_overlap(chunk_overlap) + .with_separators(separators) + } + + pub fn with_chunk_size(mut self, chunk_size: usize) -> Self { + self.chunk_size = chunk_size; + self + } + + pub fn with_chunk_overlap(mut self, chunk_overlap: usize) -> Self { + self.chunk_overlap = chunk_overlap; + self + } + + pub fn with_separators(mut self, separators: &[&str]) -> Self { + self.separators = separators.iter().map(|v| v.to_string()).collect(); + self + } + + #[allow(unused)] + pub fn with_length_function<F>(mut self, length_function: F) -> Self + where + F: Fn(&str) -> usize + Send + Sync + 'static, + { + self.length_function = Box::new(length_function); + self + } + + pub fn split_documents( + &self, + documents: &[RagDocument], + chunk_header_options: &SplitterChunkHeaderOptions, + ) -> Vec<RagDocument> { + let mut texts: Vec<String> = Vec::new(); + let mut metadatas: Vec<RagMetadata> = Vec::new(); + documents.iter().for_each(|d| { + if !d.page_content.is_empty() { + texts.push(d.page_content.clone()); + metadatas.push(d.metadata.clone()); + } + }); + + self.create_documents(&texts, &metadatas, chunk_header_options) + } + + pub fn create_documents( + &self, + texts: &[String], + metadatas: &[RagMetadata], + chunk_header_options: &SplitterChunkHeaderOptions, + ) -> Vec<RagDocument> { + let SplitterChunkHeaderOptions { + chunk_header, + chunk_overlap_header, + append_chunk_overlap_header, + } = chunk_header_options; + + let mut documents = Vec::new(); + for (i, text) in texts.iter().enumerate() { + let mut line_counter_index = 1; + let mut prev_chunk = None; + let mut index_prev_chunk = None; + + for chunk in self.split_text(text) { + let mut page_content = chunk_header.clone(); + + let index_chunk = { + let idx = match index_prev_chunk { + Some(v) => v + 1, + None => 0, + }; + text[idx..].find(&chunk).map(|i| i + idx).unwrap_or(0) + }; + if prev_chunk.is_none() { + line_counter_index += self.number_of_newlines(text, 0, index_chunk); + } else { + let index_end_prev_chunk: usize = index_prev_chunk.unwrap_or_default() + + (self.length_function)(prev_chunk.as_deref().unwrap_or_default()); + + match index_end_prev_chunk.cmp(&index_chunk) { + Ordering::Less => { + line_counter_index += + self.number_of_newlines(text, index_end_prev_chunk, index_chunk); + } + Ordering::Greater => { + let number = + self.number_of_newlines(text, index_chunk, index_end_prev_chunk); + line_counter_index = line_counter_index.saturating_sub(number); + } + Ordering::Equal => {} + } + + if *append_chunk_overlap_header { + page_content += chunk_overlap_header; + } + } + + let newlines_count = self.number_of_newlines(&chunk, 0, chunk.len()); + + let mut metadata = metadatas[i].clone(); + metadata.insert( + "loc".into(), + format!( + "{}:{}", + line_counter_index, + line_counter_index + newlines_count + ), + ); + page_content += &chunk; + documents.push(RagDocument { + page_content, + metadata, + }); + + line_counter_index += newlines_count; + prev_chunk = Some(chunk); + index_prev_chunk = Some(index_chunk); + } + } + + documents + } + + fn number_of_newlines(&self, text: &str, start: usize, end: usize) -> usize { + text[start..end].matches('\n').count() + } + + pub fn split_text(&self, text: &str) -> Vec<String> { + let keep_separator = self + .separators + .iter() + .any(|v| v.chars().any(|v| !v.is_whitespace())); + self.split_text_impl(text, &self.separators, keep_separator) + } + + fn split_text_impl( + &self, + text: &str, + separators: &[String], + keep_separator: bool, + ) -> Vec<String> { + let mut final_chunks = Vec::new(); + + let mut separator: String = separators.last().cloned().unwrap_or_default(); + let mut new_separators: Vec<String> = vec![]; + for (i, s) in separators.iter().enumerate() { + if s.is_empty() { + separator.clone_from(s); + break; + } + if text.contains(s) { + separator.clone_from(s); + new_separators = separators[i + 1..].to_vec(); + break; + } + } + + // Now that we have the separator, split the text + let splits = split_on_separator(text, &separator, keep_separator); + + // Now go merging things, recursively splitting longer texts. + let mut good_splits = Vec::new(); + let _separator = if keep_separator { "" } else { &separator }; + for s in splits { + if (self.length_function)(s) < self.chunk_size { + good_splits.push(s.to_string()); + } else { + if !good_splits.is_empty() { + let merged_text = self.merge_splits(&good_splits, _separator); + final_chunks.extend(merged_text); + good_splits.clear(); + } + if new_separators.is_empty() { + final_chunks.push(s.to_string()); + } else { + let other_info = self.split_text_impl(s, &new_separators, keep_separator); + final_chunks.extend(other_info); + } + } + } + if !good_splits.is_empty() { + let merged_text = self.merge_splits(&good_splits, _separator); + final_chunks.extend(merged_text); + } + final_chunks + } + + fn merge_splits(&self, splits: &[String], separator: &str) -> Vec<String> { + let mut docs = Vec::new(); + let mut current_doc = Vec::new(); + let mut total = 0; + for d in splits { + let _len = (self.length_function)(d); + if total + _len + current_doc.len() * separator.len() > self.chunk_size { + if total > self.chunk_size { + // warn!("Warning: Created a chunk of size {}, which is longer than the specified {}", total, self.chunk_size); + } + if !current_doc.is_empty() { + let doc = self.join_docs(¤t_doc, separator); + if let Some(doc) = doc { + docs.push(doc); + } + // Keep on popping if: + // - we have a larger chunk than in the chunk overlap + // - or if we still have any chunks and the length is long + while total > self.chunk_overlap + || (total + _len + current_doc.len() * separator.len() > self.chunk_size + && total > 0) + { + total -= (self.length_function)(¤t_doc[0]); + current_doc.remove(0); + } + } + } + current_doc.push(d.to_string()); + total += _len; + } + let doc = self.join_docs(¤t_doc, separator); + if let Some(doc) = doc { + docs.push(doc); + } + docs + } + + fn join_docs(&self, docs: &[String], separator: &str) -> Option<String> { + let text = docs.join(separator).trim().to_string(); + if text.is_empty() { + None + } else { + Some(text) + } + } +} + +pub struct SplitterChunkHeaderOptions { + pub chunk_header: String, + pub chunk_overlap_header: String, + pub append_chunk_overlap_header: bool, +} + +impl Default for SplitterChunkHeaderOptions { + fn default() -> Self { + Self { + chunk_header: "".into(), + chunk_overlap_header: "(cont'd) ".into(), + append_chunk_overlap_header: false, + } + } +} + +impl SplitterChunkHeaderOptions { + // Set the value of chunk_header + #[allow(unused)] + pub fn with_chunk_header(mut self, header: &str) -> Self { + self.chunk_header = header.to_string(); + self + } + + // Set the value of chunk_overlap_header + #[allow(unused)] + pub fn with_chunk_overlap_header(mut self, overlap_header: &str) -> Self { + self.chunk_overlap_header = overlap_header.to_string(); + self + } + + // Set the value of append_chunk_overlap_header + #[allow(unused)] + pub fn with_append_chunk_overlap_header(mut self, value: bool) -> Self { + self.append_chunk_overlap_header = value; + self + } +} + +fn split_on_separator<'a>(text: &'a str, separator: &str, keep_separator: bool) -> Vec<&'a str> { + let splits: Vec<&str> = if !separator.is_empty() { + if keep_separator { + let mut splits = Vec::new(); + let mut prev_idx = 0; + let sep_len = separator.len(); + + while let Some(idx) = text[prev_idx..].find(separator) { + splits.push(&text[prev_idx.saturating_sub(sep_len)..prev_idx + idx]); + prev_idx += idx + sep_len; + } + + if prev_idx < text.len() { + splits.push(&text[prev_idx.saturating_sub(sep_len)..]); + } + + splits + } else { + text.split(separator).collect() + } + } else { + text.split("").collect() + }; + splits.into_iter().filter(|s| !s.is_empty()).collect() +} + +#[cfg(test)] +mod tests { + use super::*; + use indexmap::IndexMap; + use pretty_assertions::assert_eq; + use serde_json::{json, Value}; + + fn build_metadata(source: &str, loc_from_line: usize, loc_to_line: usize) -> Value { + json!({ + "source": source, + "loc": format!("{loc_from_line}:{loc_to_line}"), + }) + } + #[test] + fn test_split_text() { + let splitter = Splitter { + chunk_size: 7, + chunk_overlap: 3, + separators: vec![" ".into()], + ..Default::default() + }; + let output = splitter.split_text("foo bar baz 123"); + assert_eq!(output, vec!["foo bar", "bar baz", "baz 123"]); + } + + #[test] + fn test_create_document() { + let splitter = Splitter::new(3, 0, &[" "]); + let chunk_header_options = SplitterChunkHeaderOptions::default(); + let mut metadata1 = IndexMap::new(); + metadata1.insert("source".into(), "1".into()); + let mut metadata2 = IndexMap::new(); + metadata2.insert("source".into(), "2".into()); + let output = splitter.create_documents( + &["foo bar".into(), "baz".into()], + &[metadata1, metadata2], + &chunk_header_options, + ); + let output = json!(output); + assert_eq!( + output, + json!([ + { + "page_content": "foo", + "metadata": build_metadata("1", 1, 1), + }, + { + "page_content": "bar", + "metadata": build_metadata("1", 1, 1), + }, + { + "page_content": "baz", + "metadata": build_metadata("2", 1, 1), + }, + ]) + ); + } + + #[test] + fn test_chunk_header() { + let splitter = Splitter::new(3, 0, &[" "]); + let chunk_header_options = SplitterChunkHeaderOptions::default() + .with_chunk_header("SOURCE NAME: testing\n-----\n") + .with_append_chunk_overlap_header(true); + let mut metadata1 = IndexMap::new(); + metadata1.insert("source".into(), "1".into()); + let mut metadata2 = IndexMap::new(); + metadata2.insert("source".into(), "2".into()); + let output = splitter.create_documents( + &["foo bar".into(), "baz".into()], + &[metadata1, metadata2], + &chunk_header_options, + ); + let output = json!(output); + assert_eq!( + output, + json!([ + { + "page_content": "SOURCE NAME: testing\n-----\nfoo", + "metadata": build_metadata("1", 1, 1), + }, + { + "page_content": "SOURCE NAME: testing\n-----\n(cont'd) bar", + "metadata": build_metadata("1", 1, 1), + }, + { + "page_content": "SOURCE NAME: testing\n-----\nbaz", + "metadata": build_metadata("2", 1, 1), + }, + ]) + ); + } + + #[test] + fn test_markdown_splitter() { + let text = r#"# š¦ļøš LangChain + +ā” Building applications with LLMs through composability ā” + +## Quick Install + +```bash +# Hopefully this code block isn't split +pip install langchain +``` + +As an open source project in a rapidly developing field, we are extremely open to contributions."#; + let splitter = Splitter::new(100, 0, &MARKDOWN_SEPARATES); + let output = splitter.split_text(text); + let expected_output = vec![ + "# š¦ļøš LangChain\n\nā” Building applications with LLMs through composability ā”", + "## Quick Install\n\n```bash\n# Hopefully this code block isn't split\npip install langchain", + "```", + "As an open source project in a rapidly developing field, we are extremely open to contributions.", + ]; + assert_eq!(output, expected_output); + } + + #[test] + fn test_html_splitter() { + let text = r#"<!DOCTYPE html> +<html> + <head> + <title>š¦ļøš LangChain</title> + <style> + body { + font-family: Arial, sans-serif; + } + h1 { + color: darkblue; + } + </style> + </head> + <body> + <div> + <h1>š¦ļøš LangChain</h1> + <p>ā” Building applications with LLMs through composability ā”</p> + </div> + <div> + As an open source project in a rapidly developing field, we are extremely open to contributions. + </div> + </body> +</html>"#; + let splitter = Splitter::new(175, 20, &HTML_SEPARATES); + let output = splitter.split_text(text); + let expected_output = vec![ + "<!DOCTYPE html>\n<html>", + "<head>\n <title>š¦ļøš LangChain</title>", + r#"<style> + body { + font-family: Arial, sans-serif; + } + h1 { + color: darkblue; + } + </style> + </head>"#, + r#"<body> + <div> + <h1>š¦ļøš LangChain</h1> + <p>ā” Building applications with LLMs through composability ā”</p> + </div>"#, + r#"<div> + As an open source project in a rapidly developing field, we are extremely open to contributions. + </div> + </body> +</html>"#, + ]; + assert_eq!(output, expected_output); + } +} |
