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/ernie.rs | 65 ++++++++++++++++++++++------------------------------- 1 file changed, 27 insertions(+), 38 deletions(-) (limited to 'src/client/ernie.rs') diff --git a/src/client/ernie.rs b/src/client/ernie.rs index f1d432d..e508978 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -19,7 +19,7 @@ pub struct ErnieConfig { pub secret_key: Option, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -29,54 +29,46 @@ impl ErnieClient { ("secret_key", "Secret Key:", true, PromptKind::String), ]; - fn chat_completions_builder( - &self, - client: &ReqwestClient, - data: ChatCompletionsData, - ) -> Result { + fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result { let access_token = get_access_token(self.name())?; - let mut body = build_chat_completions_body(data, &self.model); - self.patch_chat_completions_body(&mut body); - let url = format!( "{API_BASE}/wenxinworkshop/chat/{}?access_token={access_token}", &self.model.name(), ); - debug!("Ernie Chat Completions Request: {url} {body}"); + let body = build_chat_completions_body(data, &self.model); - let builder = client.post(url).json(&body); + let request_data = RequestData::new(url, body); - Ok(builder) + Ok(request_data) } - fn embeddings_builder( - &self, - client: &ReqwestClient, - data: EmbeddingsData, - ) -> Result { + fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { let access_token = get_access_token(self.name())?; - let body = json!({ - "input": data.texts, - }); - let url = format!( "{API_BASE}/wenxinworkshop/embeddings/{}?access_token={access_token}", &self.model.name(), ); - debug!("Ernie Embeddings Request: {url} {body}"); + let body = json!({ + "input": data.texts, + }); - let builder = client.post(url).json(&body); + let request_data = RequestData::new(url, body); - Ok(builder) + Ok(request_data) } - fn rerank_builder(&self, client: &ReqwestClient, data: RerankData) -> Result { + fn prepare_rerank(&self, data: RerankData) -> Result { let access_token = get_access_token(self.name())?; + let url = format!( + "{API_BASE}/wenxinworkshop/reranker/{}?access_token={access_token}", + &self.model.name(), + ); + let RerankData { query, documents, @@ -89,16 +81,9 @@ impl ErnieClient { "top_n": top_n }); - let url = format!( - "{API_BASE}/wenxinworkshop/reranker/{}?access_token={access_token}", - &self.model.name(), - ); - - debug!("Ernie Rerank Request: {url} {body}"); - - let builder = client.post(url).json(&body); + let request_data = RequestData::new(url, body); - Ok(builder) + Ok(request_data) } async fn prepare_access_token(&self) -> Result<()> { @@ -135,7 +120,8 @@ impl Client for ErnieClient { data: ChatCompletionsData, ) -> Result { self.prepare_access_token().await?; - let builder = self.chat_completions_builder(client, data)?; + let request_data = self.prepare_chat_completions(data)?; + let builder = self.request_builder(client, request_data, ApiType::ChatCompletions); chat_completions(builder).await } @@ -146,7 +132,8 @@ impl Client for ErnieClient { data: ChatCompletionsData, ) -> Result<()> { self.prepare_access_token().await?; - let builder = self.chat_completions_builder(client, data)?; + let request_data = self.prepare_chat_completions(data)?; + let builder = self.request_builder(client, request_data, ApiType::ChatCompletions); chat_completions_streaming(builder, handler).await } @@ -156,12 +143,14 @@ impl Client for ErnieClient { data: EmbeddingsData, ) -> Result { self.prepare_access_token().await?; - let builder = self.embeddings_builder(client, data)?; + let request_data = self.prepare_embeddings(data)?; + let builder = self.request_builder(client, request_data, ApiType::Embeddings); embeddings(builder).await } async fn rerank_inner(&self, client: &ReqwestClient, data: RerankData) -> Result { - let builder = self.rerank_builder(client, data)?; + let request_data = self.prepare_rerank(data)?; + let builder = self.request_builder(client, request_data, ApiType::Rerank); rerank(builder).await } } -- cgit v1.2.3