From 669f2c602c4631db1c91fd7a27098b7685027f9a Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 17 Aug 2024 16:01:39 +0800 Subject: feat: enable custom `api_base` for most clients (#793) --- src/client/gemini.rs | 21 ++++++++++++++++++--- 1 file changed, 18 insertions(+), 3 deletions(-) (limited to 'src/client/gemini.rs') diff --git a/src/client/gemini.rs b/src/client/gemini.rs index 2616218..572e082 100644 --- a/src/client/gemini.rs +++ b/src/client/gemini.rs @@ -6,12 +6,13 @@ use reqwest::RequestBuilder; use serde::Deserialize; use serde_json::{json, Value}; -const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta/models/"; +const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta"; #[derive(Debug, Clone, Deserialize, Default)] pub struct GeminiConfig { pub name: Option, pub api_key: Option, + pub api_base: Option, #[serde(default)] pub models: Vec, pub patch: Option, @@ -20,6 +21,7 @@ pub struct GeminiConfig { impl GeminiClient { config_get_fn!(api_key, get_api_key); + config_get_fn!(api_base, get_api_base); pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key:", true, PromptKind::String)]; @@ -41,13 +43,22 @@ fn prepare_chat_completions( data: ChatCompletionsData, ) -> Result { let api_key = self_.get_api_key()?; + let api_base = self_ + .get_api_base() + .unwrap_or_else(|_| API_BASE.to_string()); let func = match data.stream { true => "streamGenerateContent", false => "generateContent", }; - let url = format!("{API_BASE}{}:{}?key={}", self_.model.name(), func, api_key); + let url = format!( + "{}/models/{}:{}?key={}", + api_base.trim_end_matches('/'), + self_.model.name(), + func, + api_key + ); let body = gemini_build_chat_completions_body(data, &self_.model)?; @@ -58,9 +69,13 @@ fn prepare_chat_completions( fn prepare_embeddings(self_: &GeminiClient, data: EmbeddingsData) -> Result { let api_key = self_.get_api_key()?; + let api_base = self_ + .get_api_base() + .unwrap_or_else(|_| API_BASE.to_string()); let url = format!( - "{API_BASE}{}:embedContent?key={}", + "{}/models/{}:embedContent?key={}", + api_base.trim_end_matches('/'), self_.model.name(), api_key ); -- cgit v1.2.3