From f5e8e872b08d2aa9b3c063b0a98c59f842f94ecc Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 27 Jul 2024 21:33:04 +0800 Subject: feat: support patching request url, headers and body (#756) --- src/client/cohere.rs | 45 +++++++++++++++------------------------------ 1 file changed, 15 insertions(+), 30 deletions(-) (limited to 'src/client/cohere.rs') diff --git a/src/client/cohere.rs b/src/client/cohere.rs index 4ab6f38..b6e9755 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -2,7 +2,7 @@ use super::rag_dedicated::*; use super::*; use anyhow::{bail, Context, Result}; -use reqwest::{Client as ReqwestClient, RequestBuilder}; +use reqwest::RequestBuilder; use serde::Deserialize; use serde_json::{json, Value}; @@ -16,7 +16,7 @@ pub struct CohereConfig { pub api_key: Option, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -26,30 +26,19 @@ impl CohereClient { pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key:", true, PromptKind::String)]; - fn chat_completions_builder( - &self, - client: &ReqwestClient, - data: ChatCompletionsData, - ) -> Result { + fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result { let api_key = self.get_api_key()?; - let mut body = build_chat_completions_body(data, &self.model)?; - self.patch_chat_completions_body(&mut body); + let body = build_chat_completions_body(data, &self.model)?; - let url = CHAT_COMPLETIONS_API_URL; + let mut request_data = RequestData::new(CHAT_COMPLETIONS_API_URL, body); - debug!("Cohere Chat Completions Request: {url} {body}"); + request_data.bearer_auth(api_key); - let builder = client.post(url).bearer_auth(api_key).json(&body); - - Ok(builder) + Ok(request_data) } - fn embeddings_builder( - &self, - client: &ReqwestClient, - data: EmbeddingsData, - ) -> Result { + fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { let api_key = self.get_api_key()?; let input_type = match data.query { @@ -63,27 +52,23 @@ impl CohereClient { "input_type": input_type, }); - let url = EMBEDDINGS_API_URL; - - debug!("Cohere Embeddings Request: {url} {body}"); + let mut request_data = RequestData::new(EMBEDDINGS_API_URL, body); - let builder = client.post(url).bearer_auth(api_key).json(&body); + request_data.bearer_auth(api_key); - Ok(builder) + Ok(request_data) } - fn rerank_builder(&self, client: &ReqwestClient, data: RerankData) -> Result { + fn prepare_rerank(&self, data: RerankData) -> Result { let api_key = self.get_api_key()?; let body = rag_dedicated_build_rerank_body(data, &self.model); - let url = RERANK_API_URL; - - debug!("Cohere Rerank Request: {url} {body}"); + let mut request_data = RequestData::new(RERANK_API_URL, body); - let builder = client.post(url).bearer_auth(api_key).json(&body); + request_data.bearer_auth(api_key); - Ok(builder) + Ok(request_data) } } -- cgit v1.2.3