summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-07-01 20:19:41 +0800
committerGitHub <noreply@github.com>2024-07-01 20:19:41 +0800
commit7c6dac061b2a33aae334c08b725e02459c5e9ade (patch)
tree3c2067aaf80a24f988600b42aa12b647675bc383 /src
parentd193950d204ed82bdb3d6a3a111511d189f084bb (diff)
downloadaichat-7c6dac061b2a33aae334c08b725e02459c5e9ade.tar.gz
feat: support `.rebuild rag` repl command(#672)
Diffstat (limited to 'src')
-rw-r--r--src/config/mod.rs14
-rw-r--r--src/rag/loader.rs240
-rw-r--r--src/rag/mod.rs265
-rw-r--r--src/repl/mod.rs26
-rw-r--r--src/utils/request.rs4
-rw-r--r--src/utils/spinner.rs2
6 files changed, 289 insertions, 262 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 879f511..25891bd 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -885,11 +885,23 @@ impl Config {
Ok(())
}
+ pub async fn rebuild_rag(config: &GlobalConfig, abort_signal: AbortSignal) -> Result<()> {
+ let rag_name = match config.read().rag.clone() {
+ Some(v) => v.name().to_string(),
+ None => bail!("No RAG"),
+ };
+ let rag_path = config.read().rag_file(&rag_name)?;
+ let mut rag = Rag::load(config, &rag_name, &rag_path)?;
+ rag.rebuild(config, &rag_path, abort_signal).await?;
+ config.write().rag = Some(Arc::new(rag));
+ Ok(())
+ }
+
pub fn rag_info(&self) -> Result<String> {
if let Some(rag) = &self.rag {
rag.export()
} else {
- bail!("No rag")
+ bail!("No RAG")
}
}
diff --git a/src/rag/loader.rs b/src/rag/loader.rs
index b9fe298..23f75cb 100644
--- a/src/rag/loader.rs
+++ b/src/rag/loader.rs
@@ -2,140 +2,108 @@ use super::*;
use anyhow::{bail, Context, Result};
use async_recursion::async_recursion;
-use serde_json::Value;
use std::{collections::HashMap, path::Path};
pub const EXTENSION_METADATA: &str = "__extension__";
+pub const PATH_METADATA: &str = "__path__";
-pub async fn load(
+pub async fn load_recrusive_url(
loaders: &HashMap<String, String>,
path: &str,
- extension: &str,
-) -> Result<Vec<RagDocument>> {
- if extension == RECURSIVE_URL_LOADER {
- let loader_command = loaders
- .get(extension)
- .with_context(|| format!("RAG document loader '{extension}' not configured"))?;
- let contents = run_loader_command(path, extension, loader_command)?;
- let output = match parse_json_documents(&contents) {
- Some(v) => v,
- None => vec![RagDocument::new(contents)],
- };
- Ok(output)
- } else if extension == URL_LOADER {
- let (contents, extension) = fetch(loaders, path).await?;
- let mut metadata: RagMetadata = Default::default();
- metadata.insert("path".into(), path.into());
- metadata.insert(EXTENSION_METADATA.into(), extension);
- Ok(vec![RagDocument::new(contents).with_metadata(metadata)])
+) -> Result<Vec<(String, RagMetadata)>> {
+ let extension = RECURSIVE_URL_LOADER;
+ let loader_command = loaders
+ .get(extension)
+ .with_context(|| format!("RAG document loader '{extension}' not configured"))?;
+ let contents = run_loader_command(path, extension, loader_command)?;
+ let pages: Vec<WebPage> = serde_json::from_str(&contents).context(r#"The crawler response is invalid. It should follow the JSON format: `[{"path":"...", "text":"..."}]`."#)?;
+ let output = pages
+ .into_iter()
+ .map(|v| {
+ let WebPage { 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)
+}
+
+#[derive(Debug, Deserialize)]
+struct WebPage {
+ path: String,
+ text: String,
+}
+
+pub async fn load_path(
+ loaders: &HashMap<String, String>,
+ path: &str,
+) -> Result<Vec<(String, RagMetadata)>> {
+ let (path_str, suffixes) = parse_glob(path)?;
+ let suffixes = if suffixes.is_empty() {
+ None
} else {
- match loaders.get(extension) {
- Some(loader_command) => load_with_command(path, extension, loader_command),
- None => load_plain(path, extension).await,
+ Some(&suffixes)
+ };
+ let mut file_paths = vec![];
+ list_files(&mut file_paths, Path::new(&path_str), suffixes).await?;
+ let mut output = vec![];
+ let file_paths_len = file_paths.len();
+ match file_paths_len {
+ 0 => {}
+ 1 => output.push(load_file(loaders, &file_paths[0]).await?),
+ _ => {
+ for path in file_paths {
+ println!("🚀 Loading file {path}");
+ output.push(load_file(loaders, &path).await?)
+ }
+ println!("✨ Load directory completed");
}
}
+ Ok(output)
}
-async fn load_plain(path: &str, extension: &str) -> Result<Vec<RagDocument>> {
- let contents = tokio::fs::read_to_string(path).await?;
- if extension == "json" {
- if let Some(documents) = parse_json_documents(&contents) {
- return Ok(documents);
- }
+pub async fn load_file(
+ loaders: &HashMap<String, String>,
+ path: &str,
+) -> Result<(String, RagMetadata)> {
+ let extension = get_extension(path);
+ match loaders.get(&extension) {
+ Some(loader_command) => load_with_command(path, &extension, loader_command),
+ None => load_plain(path, &extension).await,
}
- let mut document = RagDocument::new(contents);
- document.metadata.insert("path".into(), path.to_string());
- Ok(vec![document])
+}
+
+pub async fn load_url(
+ loaders: &HashMap<String, String>,
+ path: &str,
+) -> Result<(String, RagMetadata)> {
+ let (contents, extension) = fetch(loaders, path).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<Vec<RagDocument>> {
+) -> Result<(String, RagMetadata)> {
let contents = run_loader_command(path, extension, loader_command)?;
- let mut document = RagDocument::new(contents);
- document.metadata.insert("path".into(), path.to_string());
- document
- .metadata
- .insert(EXTENSION_METADATA.into(), DEFAULT_EXTENSION.to_string());
- Ok(vec![document])
-}
-
-fn parse_json_documents(data: &str) -> Option<Vec<RagDocument>> {
- let value: Value = serde_json::from_str(data).ok()?;
- let items = match value {
- Value::Array(v) => v,
- _ => return None,
- };
- if items.is_empty() {
- return None;
- }
- match &items[0] {
- Value::String(_) => {
- let documents: Vec<_> = items
- .into_iter()
- .flat_map(|item| {
- if let Value::String(content) = item {
- Some(RagDocument::new(content))
- } else {
- None
- }
- })
- .collect();
- Some(documents)
- }
- Value::Object(obj) => {
- let key = [
- "page_content",
- "pageContent",
- "content",
- "html",
- "markdown",
- "text",
- ]
- .into_iter()
- .map(|v| v.to_string())
- .find(|key| obj.get(key).and_then(|v| v.as_str()).is_some())?;
- let documents: Vec<_> = items
- .into_iter()
- .flat_map(|item| {
- if let Value::Object(mut obj) = item {
- if let Some(page_content) = obj.get(&key).and_then(|v| v.as_str()) {
- let page_content = page_content.to_string();
- obj.remove(&key);
- let mut metadata: IndexMap<_, _> = obj
- .into_iter()
- .map(|(k, v)| {
- if let Value::String(v) = v {
- (k, v)
- } else {
- (k, v.to_string())
- }
- })
- .collect();
- if key == "markdown" {
- metadata.insert(EXTENSION_METADATA.into(), "md".into());
- } else if key == "html" {
- metadata.insert(EXTENSION_METADATA.into(), "html".into());
- }
- return Some(RagDocument {
- page_content,
- metadata,
- });
- }
- }
- None
- })
- .collect();
- if documents.is_empty() {
- None
- } else {
- Some(documents)
- }
- }
- _ => None,
- }
+ 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))
}
pub fn parse_glob(path_str: &str) -> Result<(String, Vec<String>)> {
@@ -158,6 +126,8 @@ pub fn parse_glob(path_str: &str) -> Result<(String, Vec<String>)> {
let extensions = vec![extensions_str.to_string()];
Ok((base_path, extensions))
}
+ } else if path_str.ends_with("/**") || path_str.ends_with(r"\**") {
+ Ok((path_str[0..path_str.len() - 3].to_string(), vec![]))
} else {
Ok((path_str.to_string(), vec![]))
}
@@ -209,6 +179,13 @@ fn is_valid_extension(suffixes: Option<&Vec<String>>, path: &Path) -> bool {
true
}
+fn get_extension(path: &str) -> String {
+ Path::new(&path)
+ .extension()
+ .map(|v| v.to_string_lossy().to_lowercase())
+ .unwrap_or_else(|| DEFAULT_EXTENSION.to_string())
+}
+
#[cfg(test)]
mod tests {
use super::*;
@@ -216,6 +193,7 @@ mod tests {
#[test]
fn test_parse_glob() {
assert_eq!(parse_glob("dir").unwrap(), ("dir".into(), vec![]));
+ assert_eq!(parse_glob("dir/**").unwrap(), ("dir".into(), vec![]));
assert_eq!(
parse_glob("dir/file.md").unwrap(),
("dir/file.md".into(), vec![])
@@ -233,36 +211,4 @@ mod tests {
("C:\\dir".into(), vec!["md".into(), "txt".into()])
);
}
-
- #[test]
- fn test_parse_json_documents() {
- let data = r#"["foo", "bar"]"#;
- assert_eq!(
- parse_json_documents(data).unwrap(),
- vec![RagDocument::new("foo"), RagDocument::new("bar")]
- );
-
- let data = r#"[{"content": "foo"}, {"content": "bar"}]"#;
- assert_eq!(
- parse_json_documents(data).unwrap(),
- vec![RagDocument::new("foo"), RagDocument::new("bar")]
- );
-
- let mut metadata = IndexMap::new();
- metadata.insert("k1".into(), "1".into());
- let data = r#"[{"k1": 1, "text": "foo" }]"#;
- assert_eq!(
- parse_json_documents(data).unwrap(),
- vec![RagDocument::new("foo").with_metadata(metadata.clone())]
- );
-
- let data = r#""hello""#;
- assert!(parse_json_documents(data).is_none());
-
- let data = r#"{"key":"value"}"#;
- assert!(parse_json_documents(data).is_none());
-
- let data = r#"[{"key":"value"}]"#;
- assert!(parse_json_documents(data).is_none());
- }
}
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/**")
diff --git a/src/repl/mod.rs b/src/repl/mod.rs
index 6c15002..2111c48 100644
--- a/src/repl/mod.rs
+++ b/src/repl/mod.rs
@@ -33,7 +33,7 @@ lazy_static! {
const MENU_NAME: &str = "completion_menu";
lazy_static! {
- static ref REPL_COMMANDS: [ReplCommand; 26] = [
+ static ref REPL_COMMANDS: [ReplCommand; 27] = [
ReplCommand::new(".help", "Show this help message", AssertState::pass()),
ReplCommand::new(".info", "View system info", AssertState::pass()),
ReplCommand::new(".model", "Change the current LLM", AssertState::pass()),
@@ -89,17 +89,22 @@ lazy_static! {
),
ReplCommand::new(
".rag",
- "Init or use a rag",
+ "Init or use the RAG",
AssertState::False(StateFlags::AGENT)
),
ReplCommand::new(
".info rag",
- "View rag info",
+ "View RAG info",
+ AssertState::True(StateFlags::RAG),
+ ),
+ ReplCommand::new(
+ ".rebuild rag",
+ "Rebuild the RAG to sync document changes",
AssertState::True(StateFlags::RAG),
),
ReplCommand::new(
".exit rag",
- "Leave the rag",
+ "Leave the RAG",
AssertState::TrueFalse(StateFlags::RAG, StateFlags::AGENT),
),
ReplCommand::new(".agent", "Use a agent", AssertState::bare()),
@@ -314,6 +319,19 @@ Tips: use <tab> to autocomplete conversation starter text.
}
}
}
+ ".rebuild" => {
+ match args.map(|v| match v.split_once(' ') {
+ Some((subcmd, args)) => (subcmd, Some(args.trim())),
+ None => (v, None),
+ }) {
+ Some(("rag", _)) => {
+ Config::rebuild_rag(&self.config, self.abort_signal.clone()).await?;
+ }
+ _ => {
+ println!(r#"Usage: .rebuild rag"#)
+ }
+ }
+ }
".file" => match args {
Some(args) => {
let (files, text) = split_files_text(args);
diff --git a/src/utils/request.rs b/src/utils/request.rs
index fadd3b8..678e2bb 100644
--- a/src/utils/request.rs
+++ b/src/utils/request.rs
@@ -48,7 +48,9 @@ pub async fn fetch(loaders: &HashMap<String, String>, path: &str) -> Result<(Str
}
let result = match loaders.get(&extension) {
Some(loader_command) => {
- let save_path = temp_file("-download-", "").display().to_string();
+ let save_path = temp_file("-download-", &format!(".{extension}"))
+ .display()
+ .to_string();
let mut save_file = tokio::fs::File::create(&save_path).await?;
let mut size = 0;
while let Some(chunk) = res.chunk().await? {
diff --git a/src/utils/spinner.rs b/src/utils/spinner.rs
index a04a18f..8f386db 100644
--- a/src/utils/spinner.rs
+++ b/src/utils/spinner.rs
@@ -78,11 +78,13 @@ impl Drop for Spinner {
impl Spinner {
pub fn set_message(&self, message: String) -> Result<()> {
self.0.send(SpinnerEvent::SetMessage(message))?;
+ std::thread::sleep(Duration::from_millis(10));
Ok(())
}
pub fn stop(&self) {
let _ = self.0.send(SpinnerEvent::Stop);
+ std::thread::sleep(Duration::from_millis(10));
}
}