From 55e36c7e9da2e1c93ebeaabdc8355d0a22361f03 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 31 Aug 2024 22:02:39 +0800 Subject: feat: webui support RAG (#815) --- src/config/input.rs | 33 +++------------------------------ src/config/mod.rs | 45 +++++++++++++++++++++++++++++++++++++++++---- 2 files changed, 44 insertions(+), 34 deletions(-) (limited to 'src/config') diff --git a/src/config/input.rs b/src/config/input.rs index 9c7a666..829d9ed 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -145,36 +145,9 @@ impl Input { if !self.text.is_empty() { let rag = self.config.read().rag.clone(); if let Some(rag) = rag { - let (top_k, min_score_vector_search, min_score_keyword_search) = { - let config = self.config.read(); - ( - config.rag_top_k, - config.rag_min_score_vector_search, - config.rag_min_score_keyword_search, - ) - }; - let rerank = match self.config.read().rag_reranker_model.clone() { - Some(reranker_model_id) => { - let min_score = self.config.read().rag_min_score_rerank; - let rerank_model = - Model::retrieve_reranker(&self.config.read(), &reranker_model_id)?; - let rerank_client = init_client(&self.config, Some(rerank_model))?; - Some((rerank_client, min_score)) - } - None => None, - }; - let embeddings = rag - .search( - &self.text, - top_k, - min_score_vector_search, - min_score_keyword_search, - rerank, - abort_signal, - ) - .await?; - let text = self.config.read().rag_template(&embeddings, &self.text); - self.patched_text = Some(text); + let result = + Config::search_rag(&self.config, &rag, &self.text, abort_signal).await?; + self.patched_text = Some(result); self.rag_name = Some(rag.name().to_string()); } } diff --git a/src/config/mod.rs b/src/config/mod.rs index 1a0b1d6..5d1172b 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -9,8 +9,8 @@ pub use self::role::{Role, RoleLike, BUILTIN_ROLES, CODE_ROLE, EXPLAIN_SHELL_ROL use self::session::Session; use crate::client::{ - create_client_config, list_chat_models, list_client_types, list_reranker_models, ClientConfig, - Model, OPENAI_COMPATIBLE_PLATFORMS, + create_client_config, init_client, list_chat_models, list_client_types, list_reranker_models, + ClientConfig, Model, OPENAI_COMPATIBLE_PLATFORMS, }; use crate::function::{FunctionDeclaration, Functions, ToolResult}; use crate::rag::Rag; @@ -1095,7 +1095,44 @@ impl Config { Ok(()) } - pub fn list_rags(&self) -> Vec { + pub async fn search_rag( + config: &GlobalConfig, + rag: &Rag, + text: &str, + abort_signal: AbortSignal, + ) -> Result { + let (top_k, min_score_vector_search, min_score_keyword_search) = { + let config = config.read(); + ( + config.rag_top_k, + config.rag_min_score_vector_search, + config.rag_min_score_keyword_search, + ) + }; + let rerank = match config.read().rag_reranker_model.clone() { + Some(reranker_model_id) => { + let min_score = config.read().rag_min_score_rerank; + let rerank_model = Model::retrieve_reranker(&config.read(), &reranker_model_id)?; + let rerank_client = init_client(config, Some(rerank_model))?; + Some((rerank_client, min_score)) + } + None => None, + }; + let embeddings = rag + .search( + text, + top_k, + min_score_vector_search, + min_score_keyword_search, + rerank, + abort_signal, + ) + .await?; + let text = config.read().rag_template(&embeddings, text); + Ok(text) + } + + pub fn list_rags() -> Vec { let rags_dir = match Self::rags_dir() { Ok(dir) => dir, Err(_) => return vec![], @@ -1327,7 +1364,7 @@ impl Config { .into_iter() .map(|v| (v, None)) .collect(), - ".rag" => self.list_rags().into_iter().map(|v| (v, None)).collect(), + ".rag" => Self::list_rags().into_iter().map(|v| (v, None)).collect(), ".agent" => list_agents().into_iter().map(|v| (v, None)).collect(), ".starter" => match &self.agent { Some(agent) => agent -- cgit v1.2.3