From 6b77890ec400090caa5d9fa3bffee0db91c9adfc Mon Sep 17 00:00:00 2001 From: sigoden Date: Sun, 3 Nov 2024 08:07:46 +0800 Subject: feat: support `.edit rag-docs` (#964) --- src/config/mod.rs | 51 ++++++++++++++++++++++++++++++++++++++++++++++++++- src/rag/mod.rs | 16 +++++++++++----- src/repl/mod.rs | 12 ++++++++++-- 3 files changed, 71 insertions(+), 8 deletions(-) (limited to 'src') diff --git a/src/config/mod.rs b/src/config/mod.rs index 8f3d96c..b38dcd9 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1222,12 +1222,61 @@ impl Config { Ok(()) } + pub async fn edit_rag_docs(config: &GlobalConfig, abort_signal: AbortSignal) -> Result<()> { + let mut rag = match config.read().rag.clone() { + Some(v) => v.as_ref().clone(), + None => bail!("No RAG"), + }; + + let document_paths = rag.document_paths(); + let temp_file = temp_file(&format!("-rag-{}", rag.name()), ".txt"); + tokio::fs::write(&temp_file, &document_paths.join("\n")) + .await + .with_context(|| { + format!( + "Failed to write current document paths to '{}'", + temp_file.display() + ) + })?; + let editor = config.read().editor()?; + edit_file(&editor, &temp_file)?; + let new_document_paths = + tokio::fs::read_to_string(&temp_file) + .await + .with_context(|| { + format!( + "Failed to read new document paths from '{}'", + temp_file.display() + ) + })?; + let new_document_paths = new_document_paths + .split('\n') + .filter_map(|v| { + let v = v.trim(); + if v.is_empty() { + None + } else { + Some(v.to_string()) + } + }) + .collect::>(); + if new_document_paths.is_empty() || new_document_paths == document_paths { + bail!("No changes") + } + rag.refresh_document_paths(&new_document_paths, config, abort_signal) + .await?; + config.write().rag = Some(Arc::new(rag)); + Ok(()) + } + pub async fn rebuild_rag(config: &GlobalConfig, abort_signal: AbortSignal) -> Result<()> { let mut rag = match config.read().rag.clone() { Some(v) => v.as_ref().clone(), None => bail!("No RAG"), }; - rag.rebuild(config, abort_signal).await?; + let document_paths = rag.document_paths().to_vec(); + rag.refresh_document_paths(&document_paths, config, abort_signal) + .await?; config.write().rag = Some(Arc::new(rag)); Ok(()) } diff --git a/src/rag/mod.rs b/src/rag/mod.rs index fd734e2..5bb4ea4 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -128,17 +128,23 @@ impl Rag { Ok(rag) } - pub async fn rebuild( + pub fn document_paths(&self) -> &[String] { + &self.data.document_paths + } + + pub async fn refresh_document_paths( &mut self, + document_paths: &[T], config: &GlobalConfig, abort_signal: AbortSignal, - ) -> Result<()> { - debug!("rebuild rag: {}", self.name); + ) -> Result<()> + where + T: AsRef, + { let loaders = config.read().document_loaders.clone(); let spinner = create_spinner("Starting").await; - let paths = self.data.document_paths.clone(); tokio::select! { - ret = self.sync_documents(loaders, &paths, Some(spinner.clone())) => { + ret = self.sync_documents(loaders, document_paths, Some(spinner.clone())) => { spinner.stop(); ret?; } diff --git a/src/repl/mod.rs b/src/repl/mod.rs index c08afc3..7bf70ea 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -31,7 +31,7 @@ lazy_static::lazy_static! { const MENU_NAME: &str = "completion_menu"; lazy_static::lazy_static! { - static ref REPL_COMMANDS: [ReplCommand; 34] = [ + static ref REPL_COMMANDS: [ReplCommand; 35] = [ 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()), @@ -105,6 +105,11 @@ lazy_static::lazy_static! { "Init or use the RAG", AssertState::False(StateFlags::AGENT) ), + ReplCommand::new( + ".edit rag-docs", + "Edit the RAG documents", + AssertState::True(StateFlags::RAG), + ), ReplCommand::new( ".rebuild rag", "Rebuild the RAG to sync document changes", @@ -360,8 +365,11 @@ impl Repl { Some(("session", _)) => { self.config.write().edit_session()?; } + Some(("rag-docs", _)) => { + Config::edit_rag_docs(&self.config, self.abort_signal.clone()).await?; + } _ => { - println!(r#"Usage: .edit "#) + println!(r#"Usage: .edit "#) } } } -- cgit v1.2.3