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/openai_compatible.rs | 29 ++++++++++++++++++++++++++--- 1 file changed, 26 insertions(+), 3 deletions(-) (limited to 'src/client/openai_compatible.rs') 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 { - 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 ); -- cgit v1.2.3