From 0e740d81e94505bd57036755abaaecb12c3b26e3 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sun, 28 Jul 2024 06:04:36 +0800 Subject: feat: abandon rag_dedicated client and improve (#757) --- src/client/cohere.rs | 90 +++++++++++++++++++++++++++++----------------------- 1 file changed, 50 insertions(+), 40 deletions(-) (limited to 'src/client/cohere.rs') diff --git a/src/client/cohere.rs b/src/client/cohere.rs index b6e9755..aff919e 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -1,5 +1,5 @@ -use super::rag_dedicated::*; use super::*; +use super::openai_compatible::*; use anyhow::{bail, Context, Result}; use reqwest::RequestBuilder; @@ -25,62 +25,71 @@ impl CohereClient { pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key:", true, PromptKind::String)]; +} - fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result { - let api_key = self.get_api_key()?; +impl_client_trait!( + CohereClient, + ( + prepare_chat_completions, + chat_completions, + chat_completions_streaming + ), + (prepare_embeddings, embeddings), + (prepare_rerank, generic_rerank), +); - let body = build_chat_completions_body(data, &self.model)?; +fn prepare_chat_completions( + self_: &CohereClient, + data: ChatCompletionsData, +) -> Result { + let api_key = self_.get_api_key()?; - let mut request_data = RequestData::new(CHAT_COMPLETIONS_API_URL, body); + let body = build_chat_completions_body(data, &self_.model)?; - request_data.bearer_auth(api_key); + let mut request_data = RequestData::new(CHAT_COMPLETIONS_API_URL, body); - Ok(request_data) - } + request_data.bearer_auth(api_key); - fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { - let api_key = self.get_api_key()?; + Ok(request_data) +} - let input_type = match data.query { - true => "search_query", - false => "search_document", - }; +fn prepare_embeddings(self_: &CohereClient, data: EmbeddingsData) -> Result { + let api_key = self_.get_api_key()?; + + let input_type = match data.query { + true => "search_query", + false => "search_document", + }; - let body = json!({ - "model": self.model.name(), - "texts": data.texts, - "input_type": input_type, - }); + let body = json!({ + "model": self_.model.name(), + "texts": data.texts, + "input_type": input_type, + }); - let mut request_data = RequestData::new(EMBEDDINGS_API_URL, body); + let mut request_data = RequestData::new(EMBEDDINGS_API_URL, body); - request_data.bearer_auth(api_key); + request_data.bearer_auth(api_key); - Ok(request_data) - } + Ok(request_data) +} - fn prepare_rerank(&self, data: RerankData) -> Result { - let api_key = self.get_api_key()?; +fn prepare_rerank(self_: &CohereClient, data: RerankData) -> Result { + let api_key = self_.get_api_key()?; - let body = rag_dedicated_build_rerank_body(data, &self.model); + let body = generic_build_rerank_body(data, &self_.model); - let mut request_data = RequestData::new(RERANK_API_URL, body); + let mut request_data = RequestData::new(RERANK_API_URL, body); - request_data.bearer_auth(api_key); + request_data.bearer_auth(api_key); - Ok(request_data) - } + Ok(request_data) } -impl_client_trait!( - CohereClient, - chat_completions, - chat_completions_streaming, - embeddings, - rag_dedicated_rerank -); - -async fn chat_completions(builder: RequestBuilder) -> Result { +async fn chat_completions( + builder: RequestBuilder, + _model: &Model, +) -> Result { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; @@ -95,6 +104,7 @@ async fn chat_completions(builder: RequestBuilder) -> Result Result<()> { let res = builder.send().await?; let status = res.status(); @@ -131,7 +141,7 @@ async fn chat_completions_streaming( Ok(()) } -async fn embeddings(builder: RequestBuilder) -> Result { +async fn embeddings(builder: RequestBuilder, _model: &Model) -> Result { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; -- cgit v1.2.3