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/cohere.rs | 51 +++++++++++++++++++++++++++++++++++++++++++++++++-- 1 file changed, 49 insertions(+), 2 deletions(-) (limited to 'src/client/cohere.rs') 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(); -- cgit v1.2.3