From abc588daac6053ec2edbdcde3f5a2dc5eb7d50b8 Mon Sep 17 00:00:00 2001 From: sigoden Date: Fri, 21 Jun 2024 06:00:26 +0800 Subject: feat: support rerank (#620) --- src/client/claude.rs | 3 +- src/client/cohere.rs | 51 +++++++++++++++++++- src/client/common.rs | 101 +++++++++++++++++++++++++++++++++++++--- src/client/gemini.rs | 2 +- src/client/model.rs | 13 ++++-- src/client/ollama.rs | 2 +- src/client/openai.rs | 2 +- src/client/openai_compatible.rs | 29 ++++++++++-- src/client/qianwen.rs | 2 +- src/client/reka.rs | 10 ++-- src/client/vertexai.rs | 2 +- 11 files changed, 189 insertions(+), 28 deletions(-) (limited to 'src/client') 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 { + 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 { @@ -137,7 +153,7 @@ async fn embeddings(builder: RequestBuilder) -> Result { 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>, } +pub async fn cohere_rerank(builder: RequestBuilder) -> Result { + 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 { 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 { 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>> { + 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 { + 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 { + 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>> { + ) -> Result { bail!("No embeddings api") } + + async fn rerank_inner( + &self, + _client: &ReqwestClient, + _data: RerankData, + ) -> Result { + bail!("No rerank api") + } } impl Default for ClientConfig { @@ -459,6 +521,31 @@ impl EmbeddingsData { pub type EmbeddingsOutput = Vec>; +#[derive(Debug)] +pub struct RerankData { + pub query: String, + pub documents: Vec, + pub top_n: usize, +} + +impl RerankData { + pub fn new(query: String, documents: Vec, top_n: usize) -> Self { + Self { + query, + documents, + top_n, + } + } +} + +pub type RerankOutput = Vec; + +#[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 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 { 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 { 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 { + 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 { 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 Result { - 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 { + 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 { 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 { 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 { - 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 { 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() -- cgit v1.2.3