diff options
| -rw-r--r-- | config.example.yaml | 15 | ||||
| -rw-r--r-- | models.yaml | 6 | ||||
| -rw-r--r-- | src/client/claude.rs | 3 | ||||
| -rw-r--r-- | src/client/cohere.rs | 51 | ||||
| -rw-r--r-- | src/client/common.rs | 101 | ||||
| -rw-r--r-- | src/client/gemini.rs | 2 | ||||
| -rw-r--r-- | src/client/model.rs | 13 | ||||
| -rw-r--r-- | src/client/ollama.rs | 2 | ||||
| -rw-r--r-- | src/client/openai.rs | 2 | ||||
| -rw-r--r-- | src/client/openai_compatible.rs | 29 | ||||
| -rw-r--r-- | src/client/qianwen.rs | 2 | ||||
| -rw-r--r-- | src/client/reka.rs | 10 | ||||
| -rw-r--r-- | src/client/vertexai.rs | 2 | ||||
| -rw-r--r-- | src/config/input.rs | 21 | ||||
| -rw-r--r-- | src/config/mod.rs | 29 | ||||
| -rw-r--r-- | src/rag/mod.rs | 118 |
16 files changed, 328 insertions, 78 deletions
diff --git a/config.example.yaml b/config.example.yaml index ead96f4..15877b3 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -35,16 +35,20 @@ bots: # Specifies the embedding model to use rag_embedding_model: null +# Specifies the rerank model to use +rag_rerank_model: null # Specifies the chunk size rag_chunk_size: null # Specifies the chunk overlap rag_chunk_overlap: null # Specifies the number of documents to retrieve rag_top_k: 4 -# Specifies the minimum relevance score for vector search -rag_min_score_vector: 0 -# Specifies the minimum relevance score for full-text search -rag_min_score_text: 0 +# Specifies the minimum relevance score for vector searching +rag_min_score_vector_search: 0 +# Specifies the minimum relevance score for full-text searching +rag_min_score_fulltext_search: 0 +# Specifies the minimum relevance score for reranking +rag_min_score_rerank: 0 # Defines the query structure using variables like __CONTEXT__ and __INPUT__ to tailor searches to specific needs rag_template: | @@ -88,6 +92,9 @@ clients: # max_input_tokens: 2048 # default_chunk_size: 2000 # max_concurrent_chunks: 100 + # - name: xxxx + # mode: rerank # Rerank model + # max_input_tokens: 2048 # patches: # <regex>: # The regex to match model names, e.g. '.*' 'gpt-4o' 'gpt-4o|gpt-4-.*' # chat_completions_body: # The JSON to be merged with the chat completions request body. diff --git a/models.yaml b/models.yaml index 1724036..46bf1f6 100644 --- a/models.yaml +++ b/models.yaml @@ -200,6 +200,12 @@ max_input_tokens: 512 default_chunk_size: 1000 max_concurrent_chunks: 96 + - name: rerank-english-v3.0 + mode: rerank + max_input_tokens: 4096 + - name: rerank-multilingual-v3.0 + mode: rerank + max_input_tokens: 4096 - platform: reka docs: diff --git a/src/client/claude.rs b/src/client/claude.rs index f95fafc..2b836ee 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -38,8 +38,7 @@ impl ClaudeClient { debug!("Claude Request: {url} {body}"); let mut builder = client.post(url).json(&body); - builder = builder - .header("anthropic-version", "2023-06-01"); + builder = builder.header("anthropic-version", "2023-06-01"); if let Some(api_key) = api_key { builder = builder.header("x-api-key", api_key) } diff --git a/src/client/cohere.rs b/src/client/cohere.rs index 8745347..698c2a6 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -7,6 +7,7 @@ use serde_json::{json, Value}; const CHAT_COMPLETIONS_API_URL: &str = "https://api.cohere.ai/v1/chat"; const EMBEDDINGS_API_URL: &str = "https://api.cohere.ai/v1/embed"; +const RERANK_API_URL: &str = "https://api.cohere.ai/v1/rerank"; #[derive(Debug, Clone, Deserialize, Default)] pub struct CohereConfig { @@ -69,13 +70,28 @@ impl CohereClient { Ok(builder) } + + fn rerank_builder(&self, client: &ReqwestClient, data: RerankData) -> Result<RequestBuilder> { + let api_key = self.get_api_key()?; + + let body = cohere_build_rerank_body(data, &self.model); + + let url = RERANK_API_URL; + + debug!("Cohere Rerank Request: {url} {body}"); + + let builder = client.post(url).bearer_auth(api_key).json(&body); + + Ok(builder) + } } impl_client_trait!( CohereClient, chat_completions, chat_completions_streaming, - embeddings + embeddings, + cohere_rerank ); async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> { @@ -137,7 +153,7 @@ async fn embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> { catch_error(&data, status.as_u16())?; } let res_body: EmbeddingsResBody = - serde_json::from_value(data).context("Invalid request data")?; + serde_json::from_value(data).context("Invalid embeddings data")?; Ok(res_body.embeddings) } @@ -146,6 +162,22 @@ struct EmbeddingsResBody { embeddings: Vec<Vec<f32>>, } +pub async fn cohere_rerank(builder: RequestBuilder) -> Result<RerankOutput> { + let res = builder.send().await?; + let status = res.status(); + let data: Value = res.json().await?; + if !status.is_success() { + catch_error(&data, status.as_u16())?; + } + let res_body: RerankResBody = serde_json::from_value(data).context("Invalid rerank data")?; + Ok(res_body.results) +} + +#[derive(Deserialize)] +struct RerankResBody { + results: RerankOutput, +} + fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<Value> { let ChatCompletionsData { mut messages, @@ -277,6 +309,21 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu Ok(body) } +pub fn cohere_build_rerank_body(data: RerankData, model: &Model) -> Value { + let RerankData { + query, + documents, + top_n, + } = data; + + json!({ + "model": model.name(), + "query": query, + "documents": documents, + "top_n": top_n + }) +} + fn extract_chat_completions(data: &Value) -> Result<ChatCompletionsOutput> { let text = data["text"].as_str().unwrap_or_default(); diff --git a/src/client/common.rs b/src/client/common.rs index bc84fe8..6c86acd 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -151,6 +151,10 @@ macro_rules! register_client { pub fn list_embedding_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> { list_models(config).into_iter().filter(|v| v.mode() == "embedding").collect() } + + pub fn list_rerank_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> { + list_models(config).into_iter().filter(|v| v.mode() == "rerank").collect() + } }; } @@ -236,12 +240,55 @@ macro_rules! impl_client_trait { async fn embeddings_inner( &self, - client: &ReqwestClient, - data: EmbeddingsData, - ) -> Result<Vec<Vec<f32>>> { + client: &reqwest::Client, + data: $crate::client::EmbeddingsData, + ) -> Result<$crate::client::EmbeddingsOutput> { + let builder = self.embeddings_builder(client, data)?; + $embeddings(builder).await + } + } + }; + ($client:ident, $chat_completions:path, $chat_completions_streaming:path, $embeddings:path, $rerank:path) => { + #[async_trait::async_trait] + impl $crate::client::Client for $crate::client::$client { + client_common_fns!(); + + async fn chat_completions_inner( + &self, + client: &reqwest::Client, + data: $crate::client::ChatCompletionsData, + ) -> anyhow::Result<$crate::client::ChatCompletionsOutput> { + let builder = self.chat_completions_builder(client, data)?; + $chat_completions(builder).await + } + + async fn chat_completions_streaming_inner( + &self, + client: &reqwest::Client, + handler: &mut $crate::client::SseHandler, + data: $crate::client::ChatCompletionsData, + ) -> Result<()> { + let builder = self.chat_completions_builder(client, data)?; + $chat_completions_streaming(builder, handler).await + } + + async fn embeddings_inner( + &self, + client: &reqwest::Client, + data: $crate::client::EmbeddingsData, + ) -> Result<$crate::client::EmbeddingsOutput> { let builder = self.embeddings_builder(client, data)?; $embeddings(builder).await } + + async fn rerank_inner( + &self, + client: &ReqwestClient, + data: RerankData, + ) -> Result<RerankOutput> { + let builder = self.rerank_builder(client, data)?; + $rerank(builder).await + } } }; } @@ -308,7 +355,7 @@ pub trait Client: Sync + Send { let data = input.prepare_completion_data(self.model(), false)?; self.chat_completions_inner(&client, data) .await - .with_context(|| "Failed to get chat completions") + .with_context(|| "Failed to fetch chat completions") } async fn chat_completions_streaming( @@ -334,7 +381,7 @@ pub trait Client: Sync + Send { self.chat_completions_streaming_inner(&client, handler, data).await } => { handler.done()?; - ret.with_context(|| "Failed to get chat completions") + ret.with_context(|| "Failed to fetch chat completions") } _ = watch_abort_signal(abort_signal) => { handler.done()?; @@ -348,7 +395,14 @@ pub trait Client: Sync + Send { self.model().guard_max_concurrent_chunks(&data)?; self.embeddings_inner(&client, data) .await - .with_context(|| "Failed to get embeddings") + .context("Failed to fetch embeddings") + } + + async fn rerank(&self, data: RerankData) -> Result<RerankOutput> { + let client = self.build_client()?; + self.rerank_inner(&client, data) + .await + .context("Failed to fetch rerank") } fn patch_chat_completions_body(&self, body: &mut Value) { @@ -377,9 +431,17 @@ pub trait Client: Sync + Send { &self, _client: &ReqwestClient, _data: EmbeddingsData, - ) -> Result<Vec<Vec<f32>>> { + ) -> Result<EmbeddingsOutput> { bail!("No embeddings api") } + + async fn rerank_inner( + &self, + _client: &ReqwestClient, + _data: RerankData, + ) -> Result<RerankOutput> { + bail!("No rerank api") + } } impl Default for ClientConfig { @@ -459,6 +521,31 @@ impl EmbeddingsData { pub type EmbeddingsOutput = Vec<Vec<f32>>; +#[derive(Debug)] +pub struct RerankData { + pub query: String, + pub documents: Vec<String>, + pub top_n: usize, +} + +impl RerankData { + pub fn new(query: String, documents: Vec<String>, top_n: usize) -> Self { + Self { + query, + documents, + top_n, + } + } +} + +pub type RerankOutput = Vec<RerankResult>; + +#[derive(Debug, Deserialize)] +pub struct RerankResult { + pub index: usize, + pub relevance_score: f64, +} + pub type PromptAction<'a> = (&'a str, &'a str, bool, PromptKind); pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String, Value)> { diff --git a/src/client/gemini.rs b/src/client/gemini.rs index 03eef7a..6382fb8 100644 --- a/src/client/gemini.rs +++ b/src/client/gemini.rs @@ -94,7 +94,7 @@ async fn gemini_embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> catch_error(&data, status.as_u16())?; } let res_body: EmbeddingsResBody = - serde_json::from_value(data).context("Invalid request data")?; + serde_json::from_value(data).context("Invalid embeddings data")?; let output = vec![res_body.embedding.values]; Ok(output) } diff --git a/src/client/model.rs b/src/client/model.rs index 56421bf..ebf1264 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -1,5 +1,5 @@ use super::{ - list_chat_models, list_embedding_models, + list_chat_models, list_embedding_models, list_rerank_models, message::{Message, MessageContent}, EmbeddingsData, }; @@ -46,14 +46,21 @@ impl Model { pub fn retrieve_chat(config: &Config, model_id: &str) -> Result<Self> { match Self::find(&list_chat_models(config), model_id) { Some(v) => Ok(v), - None => bail!("Invalid model '{model_id}'"), + None => bail!("Invalid chat model '{model_id}'"), } } pub fn retrieve_embedding(config: &Config, model_id: &str) -> Result<Self> { match Self::find(&list_embedding_models(config), model_id) { Some(v) => Ok(v), - None => bail!("Invalid model '{model_id}'"), + None => bail!("Invalid embedding model '{model_id}'"), + } + } + + pub fn retrieve_rerank(config: &Config, model_id: &str) -> Result<Self> { + match Self::find(&list_rerank_models(config), model_id) { + Some(v) => Ok(v), + None => bail!("Invalid rerank model '{model_id}'"), } } diff --git a/src/client/ollama.rs b/src/client/ollama.rs index 9bc8978..be065c1 100644 --- a/src/client/ollama.rs +++ b/src/client/ollama.rs @@ -141,7 +141,7 @@ async fn embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> { catch_error(&data, status.as_u16())?; } let res_body: EmbeddingsResBody = - serde_json::from_value(data).context("Invalid request data")?; + serde_json::from_value(data).context("Invalid embeddings data")?; let output = vec![res_body.embedding]; Ok(output) } diff --git a/src/client/openai.rs b/src/client/openai.rs index 0c51b33..40058e6 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -148,7 +148,7 @@ pub async fn openai_embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutp catch_error(&data, status.as_u16())?; } let res_body: EmbeddingsResBody = - serde_json::from_value(data).context("Invalid request data")?; + serde_json::from_value(data).context("Invalid embeddings data")?; let output = res_body.data.into_iter().map(|v| v.embedding).collect(); Ok(output) } diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs index f5f446a..789604f 100644 --- a/src/client/openai_compatible.rs +++ b/src/client/openai_compatible.rs @@ -1,3 +1,4 @@ +use super::cohere::*; use super::openai::*; use super::*; @@ -68,7 +69,7 @@ impl OpenAICompatibleClient { client: &ReqwestClient, data: EmbeddingsData, ) -> Result<RequestBuilder> { - let api_key = self.get_api_key()?; + let api_key = self.get_api_key().ok(); let api_base = self.get_api_base_ext()?; let body = openai_build_embeddings_body(data, &self.model); @@ -77,7 +78,28 @@ impl OpenAICompatibleClient { debug!("OpenAICompatible Embeddings Request: {url} {body}"); - let builder = client.post(url).bearer_auth(api_key).json(&body); + let mut builder = client.post(url).json(&body); + if let Some(api_key) = api_key { + builder = builder.bearer_auth(api_key); + } + + Ok(builder) + } + + fn rerank_builder(&self, client: &ReqwestClient, data: RerankData) -> Result<RequestBuilder> { + let api_key = self.get_api_key().ok(); + let api_base = self.get_api_base_ext()?; + + let body = cohere_build_rerank_body(data, &self.model); + + let url = format!("{api_base}/rerank"); + + debug!("OpenAICompatible Rerank Request: {url} {body}"); + + let mut builder = client.post(url).json(&body); + if let Some(api_key) = api_key { + builder = builder.bearer_auth(api_key); + } Ok(builder) } @@ -108,5 +130,6 @@ impl_client_trait!( OpenAICompatibleClient, openai_chat_completions, openai_chat_completions_streaming, - openai_embeddings + openai_embeddings, + cohere_rerank ); diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs index b0d0b58..ec84011 100644 --- a/src/client/qianwen.rs +++ b/src/client/qianwen.rs @@ -326,7 +326,7 @@ async fn embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> { let data: Value = builder.send().await?.json().await?; maybe_catch_error(&data)?; let res_body: EmbeddingsResBody = - serde_json::from_value(data).context("Invalid request data")?; + serde_json::from_value(data).context("Invalid embeddings data")?; let output = res_body .output .embeddings diff --git a/src/client/reka.rs b/src/client/reka.rs index 46f9118..2e9b88f 100644 --- a/src/client/reka.rs +++ b/src/client/reka.rs @@ -43,11 +43,7 @@ impl RekaClient { } } -impl_client_trait!( - RekaClient, - chat_completions, - chat_completions_streaming -); +impl_client_trait!(RekaClient, chat_completions, chat_completions_streaming); async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> { let res = builder.send().await?; @@ -113,7 +109,9 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Valu } fn extract_chat_completions(data: &Value) -> Result<ChatCompletionsOutput> { - let text = data["responses"][0]["message"]["content"].as_str().ok_or_else(|| anyhow!("Invalid response data: {data}"))?; + let text = data["responses"][0]["message"]["content"] + .as_str() + .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; let output = ChatCompletionsOutput { text: text.to_string(), diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index dc75c9f..69d1bc4 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -185,7 +185,7 @@ async fn embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> { catch_error(&data, status.as_u16())?; } let res_body: EmbeddingsResBody = - serde_json::from_value(data).context("Invalid request data")?; + serde_json::from_value(data).context("Invalid embeddings data")?; let output = res_body .predictions .into_iter() diff --git a/src/config/input.rs b/src/config/input.rs index 892a41f..03ec999 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -169,20 +169,31 @@ 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, min_score_text) = { + let (top_k, min_score_vector_search, min_score_fulltext_search) = { let config = self.config.read(); ( config.rag_top_k, - config.rag_min_score_vector, - config.rag_min_score_text, + config.rag_min_score_vector_search, + config.rag_min_score_fulltext_search, ) }; + let rerank = match self.config.read().rag_rerank_model.clone() { + Some(rerank_model_id) => { + let min_score = self.config.read().rag_min_score_rerank; + let rerank_model = + Model::retrieve_rerank(&self.config.read(), &rerank_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, - min_score_text, + min_score_vector_search, + min_score_fulltext_search, + rerank, abort_signal, ) .await?; diff --git a/src/config/mod.rs b/src/config/mod.rs index 509e4de..e07dca6 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -9,8 +9,8 @@ pub use self::role::{Role, RoleLike, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; use self::session::Session; use crate::client::{ - create_client_config, list_chat_models, list_client_types, ClientConfig, Model, - OPENAI_COMPATIBLE_PLATFORMS, + create_client_config, list_chat_models, list_client_types, list_rerank_models, ClientConfig, + Model, OPENAI_COMPATIBLE_PLATFORMS, }; use crate::function::{FunctionDeclaration, Functions, FunctionsFilter, ToolResult}; use crate::rag::Rag; @@ -102,11 +102,13 @@ pub struct Config { pub bot_prelude: Option<String>, pub bots: Vec<BotConfig>, pub rag_embedding_model: Option<String>, + pub rag_rerank_model: Option<String>, pub rag_chunk_size: Option<usize>, pub rag_chunk_overlap: Option<usize>, pub rag_top_k: usize, - pub rag_min_score_vector: f32, - pub rag_min_score_text: f32, + pub rag_min_score_vector_search: f32, + pub rag_min_score_fulltext_search: f32, + pub rag_min_score_rerank: f32, pub rag_template: Option<String>, pub compress_threshold: usize, pub summarize_prompt: Option<String>, @@ -156,11 +158,13 @@ impl Default for Config { bot_prelude: None, bots: vec![], rag_embedding_model: None, + rag_rerank_model: None, rag_chunk_size: None, rag_chunk_overlap: None, rag_top_k: 4, - rag_min_score_vector: 0.0, - rag_min_score_text: 0.0, + rag_min_score_vector_search: 0.0, + rag_min_score_fulltext_search: 0.0, + rag_min_score_rerank: 0.0, rag_template: None, compress_threshold: 4000, summarize_prompt: None, @@ -442,6 +446,10 @@ impl Config { ), ("temperature", format_option_value(&role.temperature())), ("top_p", format_option_value(&role.top_p())), + ( + "rag_rerank_model", + format_option_value(&self.rag_rerank_model), + ), ("rag_top_k", self.rag_top_k.to_string()), ("function_calling", self.function_calling.to_string()), ("compress_threshold", self.compress_threshold.to_string()), @@ -490,6 +498,13 @@ impl Config { let value = parse_value(value)?; self.set_top_p(value); } + "rag_rerank_model" => { + self.rag_rerank_model = if value == "null" { + None + } else { + Some(value.to_string()) + } + } "rag_top_k" => { if let Some(value) = parse_value(value)? { self.rag_top_k = value; @@ -1052,6 +1067,7 @@ impl Config { "max_output_tokens", "temperature", "top_p", + "rag_rerank_model", "rag_top_k", "function_calling", "compress_threshold", @@ -1072,6 +1088,7 @@ impl Config { Some(v) => vec![v.to_string()], None => vec![], }, + "rag_rerank_model" => list_rerank_models(self).iter().map(|v| v.id()).collect(), "function_calling" => complete_bool(self.function_calling), "save" => complete_bool(self.save), "save_session" => { diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 7939d98..16c3d99 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -14,6 +14,7 @@ use anyhow::bail; use anyhow::{anyhow, Context, Result}; use hnsw_rs::prelude::*; use indexmap::IndexMap; +use indexmap::IndexSet; use inquire::{required, validator::Validation, Select, Text}; use path_absolutize::Absolutize; use serde::{Deserialize, Serialize}; @@ -22,13 +23,13 @@ use std::{collections::HashMap, fmt::Debug, io::BufReader, path::Path}; use tokio::sync::mpsc; pub struct Rag { - client: Box<dyn Client>, name: String, path: String, - model: Model, + embedding_model: Model, hnsw: Hnsw<'static, f32, DistCosine>, bm25: BM25<VectorID>, data: RagData, + embedding_client: Box<dyn Client>, } impl Debug for Rag { @@ -36,7 +37,7 @@ impl Debug for Rag { f.debug_struct("Rag") .field("name", &self.name) .field("path", &self.path) - .field("model", &self.model) + .field("embedding_model", &self.embedding_model) .field("data", &self.data) .finish() } @@ -51,8 +52,8 @@ impl Rag { abort_signal: AbortSignal, ) -> Result<Self> { debug!("init rag: {name}"); - let (model, chunk_size, chunk_overlap) = Self::config(config)?; - let data = RagData::new(&model.id(), chunk_size, chunk_overlap); + let (embedding_model, chunk_size, chunk_overlap) = Self::config(config)?; + let data = RagData::new(embedding_model.id(), chunk_size, chunk_overlap); let mut rag = Self::create(config, name, save_path, data)?; let mut paths = doc_paths.to_vec(); if paths.is_empty() { @@ -88,22 +89,22 @@ impl Rag { pub fn create(config: &GlobalConfig, name: &str, path: &Path, data: RagData) -> Result<Self> { let hnsw = data.build_hnsw(); let bm25 = data.build_bm25(); - let model = Model::retrieve_embedding(&config.read(), &data.model)?; - let client = init_client(config, Some(model.clone()))?; + let embedding_model = Model::retrieve_embedding(&config.read(), &data.embedding_model)?; + let embedding_client = init_client(config, Some(embedding_model.clone()))?; let rag = Rag { - client, name: name.to_string(), path: path.display().to_string(), data, - model, + embedding_model, hnsw, bm25, + embedding_client, }; Ok(rag) } pub fn config(config: &GlobalConfig) -> Result<(Model, usize, usize)> { - let (embedding_model, chunk_size, chunk_overlap) = { + let (embedding_model_id, chunk_size, chunk_overlap) = { let config = config.read(); ( config.rag_embedding_model.clone(), @@ -111,7 +112,7 @@ impl Rag { config.rag_chunk_overlap, ) }; - let model_id = match embedding_model { + let embedding_model_id = match embedding_model_id { Some(value) => { println!("Select embedding model: {value}"); value @@ -130,7 +131,8 @@ impl Rag { } } }; - let model = Model::retrieve_embedding(&config.read(), &model_id)?; + let embedding_model = Model::retrieve_embedding(&config.read(), &embedding_model_id)?; + let chunk_size = match chunk_size { Some(value) => { println!("Set chunk size: {value}"); @@ -138,9 +140,9 @@ impl Rag { } None => { if *IS_STDOUT_TERMINAL { - set_chunk_size(&model)? + set_chunk_size(&embedding_model)? } else { - let value = model.default_chunk_size(); + let value = embedding_model.default_chunk_size(); println!("Set chunk size: {value}"); value } @@ -161,7 +163,8 @@ impl Rag { } } }; - Ok((model, chunk_size, chunk_overlap)) + + Ok((embedding_model, chunk_size, chunk_overlap)) } pub fn save(&self, path: &Path) -> Result<()> { @@ -176,7 +179,7 @@ impl Rag { let files: Vec<_> = self.data.files.iter().map(|v| &v.path).collect(); let data = json!({ "path": self.path, - "model": self.model.id(), + "embedding_model": self.embedding_model.id(), "chunk_size": self.data.chunk_size, "chunk_overlap": self.data.chunk_overlap, "files": files, @@ -198,13 +201,14 @@ impl Rag { &self, text: &str, top_k: usize, - min_score_vector: f32, - min_score_text: f32, + min_score_vector_search: f32, + min_score_fulltext_search: f32, + rerank: Option<(Box<dyn Client>, f32)>, abort_signal: AbortSignal, ) -> Result<String> { let (stop_spinner_tx, _) = run_spinner("Searching").await; let ret = tokio::select! { - ret = self.hybird_search(text, top_k, min_score_vector, min_score_text) => { + ret = self.hybird_search(text, top_k, min_score_vector_search, min_score_fulltext_search, rerank) => { ret } _ = watch_abort_signal(abort_signal) => { @@ -290,6 +294,7 @@ impl Rag { self.data.add(rag_files, vector_ids, embeddings); progress(&progress_tx, "Building vector store".into()); self.hnsw = self.data.build_hnsw(); + self.bm25 = self.data.build_bm25(); Ok(()) } @@ -298,22 +303,58 @@ impl Rag { &self, query: &str, top_k: usize, - min_score_vector: f32, - min_score_text: f32, + min_score_vector_search: f32, + min_score_fulltext_search: f32, + rerank: Option<(Box<dyn Client>, f32)>, ) -> Result<Vec<String>> { let (vector_search_result, text_search_result) = tokio::join!( - self.vector_search(query, top_k, min_score_vector), - self.text_search(query, top_k, min_score_text) + self.vector_search(query, top_k, min_score_vector_search), + self.fulltext_search(query, top_k, min_score_fulltext_search) ); let vector_search_ids = vector_search_result?; - let text_search_ids = text_search_result?; - let ids = reciprocal_rank_fusion(vector_search_ids, text_search_ids, 1.0, 1.0, top_k); - let output: Vec<_> = ids + let fulltext_search_ids = text_search_result?; + debug!("vector_search_ids: {vector_search_ids:?}, fulltext_search_ids: {fulltext_search_ids:?}"); + let ids = match rerank { + Some((client, min_score)) => { + let min_score = min_score as f64; + let ids: IndexSet<VectorID> = [vector_search_ids, fulltext_search_ids] + .concat() + .into_iter() + .collect(); + let mut documents = vec![]; + let mut documents_ids = vec![]; + for id in ids { + if let Some(document) = self.data.get(id) { + documents_ids.push(id); + documents.push(document.page_content.to_string()); + } + } + let data = RerankData::new(query.to_string(), documents, top_k); + let list = client.rerank(data).await?; + let ids = list + .into_iter() + .filter_map(|item| { + if item.relevance_score < min_score { + None + } else { + documents_ids.get(item.index).cloned() + } + }) + .collect(); + debug!("rerank_ids: {ids:?}"); + ids + } + None => { + let ids = + reciprocal_rank_fusion(vector_search_ids, fulltext_search_ids, 1.0, 1.0, top_k); + debug!("rrf_ids: {ids:?}"); + ids + } + }; + let output = ids .into_iter() .filter_map(|id| { - let (file_index, document_index) = split_vector_id(id); - let file = self.data.files.get(file_index)?; - let document = file.documents.get(document_index)?; + let document = self.data.get(id)?; Some(document.page_content.clone()) }) .collect(); @@ -352,7 +393,7 @@ impl Rag { Ok(output) } - async fn text_search( + async fn fulltext_search( &self, query: &str, top_k: usize, @@ -369,7 +410,7 @@ impl Rag { ) -> Result<EmbeddingsOutput> { let EmbeddingsData { texts, query } = data; let mut output = vec![]; - let chunks = texts.chunks(self.model.max_concurrent_chunks()); + let chunks = texts.chunks(self.embedding_model.max_concurrent_chunks()); let chunks_len = chunks.len(); progress( &progress_tx, @@ -381,7 +422,7 @@ impl Rag { query, }; let chunk_output = self - .client + .embedding_client .embeddings(chunk_data) .await .context("Failed to create embedding")?; @@ -397,7 +438,7 @@ impl Rag { #[derive(Debug, Clone, Serialize, Deserialize)] pub struct RagData { - pub model: String, + pub embedding_model: String, pub chunk_size: usize, pub chunk_overlap: usize, pub files: Vec<RagFile>, @@ -405,9 +446,9 @@ pub struct RagData { } impl RagData { - pub fn new(model: &str, chunk_size: usize, chunk_overlap: usize) -> Self { + pub fn new(embedding_model: String, chunk_size: usize, chunk_overlap: usize) -> Self { Self { - model: model.to_string(), + embedding_model, chunk_size, chunk_overlap, files: Default::default(), @@ -415,6 +456,13 @@ impl RagData { } } + pub fn get(&self, id: VectorID) -> Option<&RagDocument> { + let (file_index, document_index) = split_vector_id(id); + let file = self.files.get(file_index)?; + let document = file.documents.get(document_index)?; + Some(document) + } + pub fn add( &mut self, files: Vec<RagFile>, |
