summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-09-14 20:10:49 +0800
committerGitHub <noreply@github.com>2024-09-14 20:10:49 +0800
commitbeaf5946a98f1552763ff974942cc71ee531fe62 (patch)
tree2588545daec3169cc93c389de4bbcb421a1cd60d
parent6211d01a648e941fc69954d0855bcdcef98f27b9 (diff)
downloadaichat-beaf5946a98f1552763ff974942cc71ee531fe62.tar.gz
feat: add `.source rag` repl command (#871)
-rw-r--r--src/config/mod.rs13
-rw-r--r--src/rag/mod.rs36
-rw-r--r--src/repl/mod.rs21
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);