summaryrefslogtreecommitdiffstats
path: root/src/client/gemini.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-08-17 16:01:39 +0800
committerGitHub <noreply@github.com>2024-08-17 16:01:39 +0800
commit669f2c602c4631db1c91fd7a27098b7685027f9a (patch)
tree3f743d4c2fa6e4eb847c6f9f2618b140f75b95b4 /src/client/gemini.rs
parent580ed6bea370345f76ca69ecb4c1cc30afa689c5 (diff)
downloadaichat-669f2c602c4631db1c91fd7a27098b7685027f9a.tar.gz
feat: enable custom `api_base` for most clients (#793)
Diffstat (limited to 'src/client/gemini.rs')
-rw-r--r--src/client/gemini.rs21
1 files changed, 18 insertions, 3 deletions
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<String>,
pub api_key: Option<String>,
+ pub api_base: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
pub patch: Option<RequestPatch>,
@@ -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<RequestData> {
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<RequestData> {
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
);