summaryrefslogtreecommitdiffstats
path: root/src/rag
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-05 09:02:23 +0800
committerGitHub <noreply@github.com>2024-06-05 09:02:23 +0800
commit1ec6abfaee2fdc189b348b7e3a8145bd9a84da74 (patch)
tree6f68a860a39b0fbe87784de925b5a31e00e74e33 /src/rag
parent71f2e94579511d7524f5534377001ab3f02a9597 (diff)
downloadaichat-1ec6abfaee2fdc189b348b7e3a8145bd9a84da74.tar.gz
feat: support RAG (#560)
* feat: support RAG * support more embeddings models and implement concurrent embedding api * show the progress of addings paths * ignore embedding context when saving message * embedding model max_chunk_size => default_chunk_size * support pdf and pandoc formats (docx, epub, ipynb)
Diffstat (limited to 'src/rag')
-rw-r--r--src/rag/loader.rs146
-rw-r--r--src/rag/mod.rs425
-rw-r--r--src/rag/splitter.rs564
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(&current_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)(&current_doc[0]);
+ current_doc.remove(0);
+ }
+ }
+ }
+ current_doc.push(d.to_string());
+ total += _len;
+ }
+ let doc = self.join_docs(&current_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);
+ }
+}