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 | |
| parent | d193950d204ed82bdb3d6a3a111511d189f084bb (diff) | |
| download | aichat-7c6dac061b2a33aae334c08b725e02459c5e9ade.tar.gz | |
feat: support `.rebuild rag` repl command(#672)
| -rw-r--r-- | src/config/mod.rs | 14 | ||||
| -rw-r--r-- | src/rag/loader.rs | 240 | ||||
| -rw-r--r-- | src/rag/mod.rs | 265 | ||||
| -rw-r--r-- | src/repl/mod.rs | 26 | ||||
| -rw-r--r-- | src/utils/request.rs | 4 | ||||
| -rw-r--r-- | src/utils/spinner.rs | 2 |
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)); } } |
