diff options
| author | sigoden <sigoden@gmail.com> | 2024-07-27 21:33:04 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-07-27 21:33:04 +0800 |
| commit | f5e8e872b08d2aa9b3c063b0a98c59f842f94ecc (patch) | |
| tree | 6f6a6ce299d7383d290a7015b52ce777566d1e8e /src/client/qianwen.rs | |
| parent | adf6716c8436b53854eb55588d258f4326580a4d (diff) | |
| download | aichat-f5e8e872b08d2aa9b3c063b0a98c59f842f94ecc.tar.gz | |
feat: support patching request url, headers and body (#756)
Diffstat (limited to 'src/client/qianwen.rs')
| -rw-r--r-- | src/client/qianwen.rs | 46 |
1 files changed, 20 insertions, 26 deletions
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs index e3aea31..4aa67ae 100644 --- a/src/client/qianwen.rs +++ b/src/client/qianwen.rs @@ -27,7 +27,7 @@ pub struct QianwenConfig { pub api_key: Option<String>, #[serde(default)] pub models: Vec<ModelData>, - pub patch: Option<ModelPatch>, + pub patch: Option<RequestPatch>, pub extra: Option<ExtraConfig>, } @@ -37,11 +37,7 @@ impl QianwenClient { pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key:", true, PromptKind::String)]; - fn chat_completions_builder( - &self, - client: &ReqwestClient, - data: ChatCompletionsData, - ) -> Result<RequestBuilder> { + fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result<RequestData> { let api_key = self.get_api_key()?; let stream = data.stream; @@ -50,27 +46,24 @@ impl QianwenClient { true => CHAT_COMPLETIONS_API_URL_VL, false => CHAT_COMPLETIONS_API_URL, }; - let (mut body, has_upload) = build_chat_completions_body(data, &self.model)?; - self.patch_chat_completions_body(&mut body); - debug!("Qianwen Chat Completions Request: {url} {body}"); + let (body, has_upload) = build_chat_completions_body(data, &self.model)?; + + let mut request_data = RequestData::new(url, body); + + request_data.bearer_auth(api_key); - let mut builder = client.post(url).bearer_auth(api_key).json(&body); if stream { - builder = builder.header("X-DashScope-SSE", "enable"); + request_data.header("X-DashScope-SSE", "enable"); } if has_upload { - builder = builder.header("X-DashScope-OssResourceResolve", "enable"); + request_data.header("X-DashScope-OssResourceResolve", "enable"); } - Ok(builder) + Ok(request_data) } - fn embeddings_builder( - &self, - client: &ReqwestClient, - data: EmbeddingsData, - ) -> Result<RequestBuilder> { + fn prepare_embeddings(&self, data: EmbeddingsData) -> Result<RequestData> { let api_key = self.get_api_key()?; let text_type = match data.query { @@ -88,13 +81,11 @@ impl QianwenClient { } }); - let url = EMBEDDINGS_API_URL; - - debug!("Qianwen 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) } } @@ -109,7 +100,8 @@ impl Client for QianwenClient { ) -> Result<ChatCompletionsOutput> { let api_key = self.get_api_key()?; patch_messages(self.model.name(), &api_key, &mut data.messages).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, &self.model).await } @@ -121,7 +113,8 @@ impl Client for QianwenClient { ) -> Result<()> { let api_key = self.get_api_key()?; patch_messages(self.model.name(), &api_key, &mut data.messages).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, &self.model).await } @@ -130,7 +123,8 @@ impl Client for QianwenClient { client: &ReqwestClient, data: EmbeddingsData, ) -> Result<Vec<Vec<f32>>> { - 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 } } |
