summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/claude.rs3
-rw-r--r--src/client/cohere.rs51
-rw-r--r--src/client/common.rs101
-rw-r--r--src/client/gemini.rs2
-rw-r--r--src/client/model.rs13
-rw-r--r--src/client/ollama.rs2
-rw-r--r--src/client/openai.rs2
-rw-r--r--src/client/openai_compatible.rs29
-rw-r--r--src/client/qianwen.rs2
-rw-r--r--src/client/reka.rs10
-rw-r--r--src/client/vertexai.rs2
-rw-r--r--src/config/input.rs21
-rw-r--r--src/config/mod.rs29
-rw-r--r--src/rag/mod.rs118
14 files changed, 311 insertions, 74 deletions
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>,