summaryrefslogtreecommitdiffstats
path: root/src/client/vertexai.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-25 10:15:54 +0800
committerGitHub <noreply@github.com>2024-04-25 10:15:54 +0800
commit4db9b309803796bc5f996d0b3713344eb44207ec (patch)
treea0ba8f3a2da32e5392e22272c641e3939de49be6 /src/client/vertexai.rs
parent8dacea4deb0fe67c0cdbe15bd47e6e50c1dd12a1 (diff)
downloadaichat-4db9b309803796bc5f996d0b3713344eb44207ec.tar.gz
refactor: rewrite list models of all clients (#436)
Diffstat (limited to 'src/client/vertexai.rs')
-rw-r--r--src/client/vertexai.rs31
1 files changed, 15 insertions, 16 deletions
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index e0ae567..30aba75 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -13,13 +13,6 @@ use serde::Deserialize;
use serde_json::{json, Value};
use std::path::PathBuf;
-const MODELS: [(&str, usize, &str); 3] = [
- // https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models
- ("gemini-1.0-pro", 24568, "text"),
- ("gemini-1.0-pro-vision", 14336, "text,vision"),
- ("gemini-1.5-pro-preview-0409", 1000000, "text,vision"),
-];
-
static mut ACCESS_TOKEN: (String, i64) = (String::new(), 0); // safe under linear operation
#[derive(Debug, Clone, Deserialize, Default)]
@@ -56,7 +49,15 @@ impl Client for VertexAIClient {
}
impl VertexAIClient {
- list_models_fn!(VertexAIConfig, &MODELS);
+ list_models_fn!(
+ VertexAIConfig,
+ [
+ // https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models
+ ("gemini-1.0-pro", "text", 24568),
+ ("gemini-1.0-pro-vision", "text,vision", 14336),
+ ("gemini-1.5-pro-preview-0409", "text,vision", 1000000),
+ ]
+ );
config_get_fn!(api_base, get_api_base);
pub const PROMPTS: [PromptType<'static>; 1] =
@@ -216,16 +217,14 @@ pub(crate) fn build_body(
]);
}
- if let Some(max_output_tokens) = model.max_output_tokens {
- body["generationConfig"]["maxOutputTokens"] = max_output_tokens.into();
+ if let Some(v) = model.max_output_tokens {
+ body["generationConfig"]["maxOutputTokens"] = v.into();
}
-
- if let Some(temperature) = temperature {
- body["generationConfig"]["temperature"] = temperature.into();
+ if let Some(v) = temperature {
+ body["generationConfig"]["temperature"] = v.into();
}
-
- if let Some(top_p) = top_p {
- body["generationConfig"]["topP"] = top_p.into();
+ if let Some(v) = top_p {
+ body["generationConfig"]["topP"] = v.into();
}
Ok(body)