diff options
| author | sigoden <sigoden@gmail.com> | 2024-09-14 20:10:49 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-09-14 20:10:49 +0800 |
| commit | beaf5946a98f1552763ff974942cc71ee531fe62 (patch) | |
| tree | 2588545daec3169cc93c389de4bbcb421a1cd60d | |
| parent | 6211d01a648e941fc69954d0855bcdcef98f27b9 (diff) | |
| download | aichat-beaf5946a98f1552763ff974942cc71ee531fe62.tar.gz | |
feat: add `.source rag` repl command (#871)
| -rw-r--r-- | src/config/mod.rs | 13 | ||||
| -rw-r--r-- | src/rag/mod.rs | 36 | ||||
| -rw-r--r-- | src/repl/mod.rs | 21 |
3 files changed, 63 insertions, 7 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs index 4f7cfd7..cd2d4ad 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1199,6 +1199,16 @@ impl Config { Ok(()) } + pub fn rag_sources(config: &GlobalConfig) -> Result<String> { + match config.read().rag.as_ref() { + Some(rag) => match rag.get_last_sources() { + Some(v) => Ok(v), + None => bail!("No sources"), + }, + None => bail!("No RAG"), + } + } + pub fn rag_info(&self) -> Result<String> { if let Some(rag) = &self.rag { rag.export() @@ -1226,7 +1236,7 @@ impl Config { config.rag_min_score_keyword_search, ) }; - let embeddings = rag + let (embeddings, ids) = rag .search( text, top_k, @@ -1237,6 +1247,7 @@ impl Config { ) .await?; let text = config.read().rag_template(&embeddings, text); + rag.set_last_sources(&ids); Ok(text) } diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 42e6582..973bc1c 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -15,6 +15,7 @@ use anyhow::{anyhow, bail, Context, Result}; use hnsw_rs::prelude::*; use indexmap::{IndexMap, IndexSet}; use inquire::{required, validator::Validation, Confirm, Select, Text}; +use parking_lot::RwLock; use path_absolutize::Absolutize; use serde::{Deserialize, Serialize}; use serde_json::json; @@ -28,6 +29,7 @@ pub struct Rag { hnsw: Hnsw<'static, f32, DistCosine>, bm25: BM25<DocumentId>, data: RagData, + last_sources: RwLock<Option<String>>, } impl Debug for Rag { @@ -51,6 +53,7 @@ impl Clone for Rag { hnsw: self.data.build_hnsw(), bm25: self.bm25.clone(), data: self.data.clone(), + last_sources: RwLock::new(None), } } } @@ -119,6 +122,7 @@ impl Rag { embedding_model, hnsw, bm25, + last_sources: RwLock::new(None), }; Ok(rag) } @@ -216,6 +220,27 @@ impl Rag { (self.data.reranker_model.clone(), self.data.top_k) } + pub fn get_last_sources(&self) -> Option<String> { + self.last_sources.read().clone() + } + + pub fn set_last_sources(&self, ids: &[DocumentId]) { + let sources: IndexSet<_> = ids + .iter() + .filter_map(|id| { + let (file_index, _) = split_document_id(*id); + let file = self.data.files.get(&file_index)?; + Some(file.path.clone()) + }) + .collect(); + let sources = if sources.is_empty() { + None + } else { + Some(sources.into_iter().collect::<Vec<_>>().join("\n")) + }; + *self.last_sources.write() = sources; + } + pub fn set_reranker_model(&mut self, reranker_model: Option<String>) -> Result<()> { self.data.reranker_model = reranker_model; self.save()?; @@ -287,7 +312,7 @@ impl Rag { min_score_keyword_search: f32, rerank_model: Option<&str>, abort_signal: AbortSignal, - ) -> Result<String> { + ) -> Result<(String, Vec<DocumentId>)> { let spinner = create_spinner("Searching").await; let ret = tokio::select! { ret = self.hybird_search(text, top_k, min_score_vector_search, min_score_keyword_search, rerank_model) => { @@ -298,8 +323,9 @@ impl Rag { }, }; spinner.stop(); - let output = ret?.join("\n\n"); - Ok(output) + let (ids, documents): (Vec<_>, Vec<_>) = ret?.into_iter().unzip(); + let embeddings = documents.join("\n\n"); + Ok((embeddings, ids)) } pub async fn sync_documents<T: AsRef<str>>( @@ -426,7 +452,7 @@ impl Rag { min_score_vector_search: f32, min_score_keyword_search: f32, rerank_model: Option<&str>, - ) -> Result<Vec<String>> { + ) -> Result<Vec<(DocumentId, String)>> { let (vector_search_result, text_search_result) = tokio::join!( self.vector_search(query, top_k, min_score_vector_search), self.keyword_search(query, top_k, min_score_keyword_search) @@ -478,7 +504,7 @@ impl Rag { .into_iter() .filter_map(|id| { let document = self.data.get(id)?; - Some(document.page_content.clone()) + Some((id, document.page_content.clone())) }) .collect(); Ok(output) diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 09fa74d..57e0376 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; 32] = [ + static ref REPL_COMMANDS: [ReplCommand; 33] = [ 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()), @@ -106,6 +106,11 @@ lazy_static::lazy_static! { AssertState::True(StateFlags::RAG), ), ReplCommand::new( + ".sources rag", + "View the RAG sources in the last query", + AssertState::True(StateFlags::RAG), + ), + ReplCommand::new( ".info rag", "View RAG info", AssertState::True(StateFlags::RAG), @@ -368,6 +373,20 @@ impl Repl { } } } + ".sources" => { + match args.map(|v| match v.split_once(' ') { + Some((subcmd, args)) => (subcmd, Some(args.trim())), + None => (v, None), + }) { + Some(("rag", _)) => { + let output = Config::rag_sources(&self.config)?; + println!("{}", output); + } + _ => { + println!(r#"Usage: .sources rag"#) + } + } + } ".file" => match args { Some(args) => { let (files, text) = split_files_text(args); |
