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/gemini.rs | 43 +++++++++++++++---------------------------- 1 file changed, 15 insertions(+), 28 deletions(-) (limited to 'src/client/gemini.rs') diff --git a/src/client/gemini.rs b/src/client/gemini.rs index 37f7c12..aa1a5b1 100644 --- a/src/client/gemini.rs +++ b/src/client/gemini.rs @@ -2,7 +2,7 @@ use super::vertexai::*; use super::*; use anyhow::{Context, Result}; -use reqwest::{Client as ReqwestClient, RequestBuilder}; +use reqwest::RequestBuilder; use serde::Deserialize; use serde_json::{json, Value}; @@ -14,7 +14,7 @@ pub struct GeminiConfig { pub api_key: Option, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -24,11 +24,7 @@ impl GeminiClient { 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 func = match data.stream { @@ -36,25 +32,24 @@ impl GeminiClient { false => "generateContent", }; - let mut body = gemini_build_chat_completions_body(data, &self.model)?; - self.patch_chat_completions_body(&mut body); - let url = format!("{API_BASE}{}:{}?key={}", &self.model.name(), func, api_key); - debug!("Gemini Chat Completions Request: {url} {body}"); + let body = gemini_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 api_key = self.get_api_key()?; + let url = format!( + "{API_BASE}{}:embedContent?key={}", + &self.model.name(), + api_key + ); + let body = json!({ "content": { "parts": [ @@ -65,17 +60,9 @@ impl GeminiClient { } }); - let url = format!( - "{API_BASE}{}:embedContent?key={}", - &self.model.name(), - api_key - ); - - debug!("Gemini Embeddings Request: {url} {body}"); - - let builder = client.post(url).json(&body); + let request_data = RequestData::new(url, body); - Ok(builder) + Ok(request_data) } } -- cgit v1.2.3