summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
Diffstat (limited to 'src/client')
-rw-r--r--src/client/claude.rs38
-rw-r--r--src/client/cohere.rs34
-rw-r--r--src/client/common.rs34
-rw-r--r--src/client/ernie.rs106
-rw-r--r--src/client/gemini.rs17
-rw-r--r--src/client/message.rs21
-rw-r--r--src/client/mistral.rs17
-rw-r--r--src/client/model.rs11
-rw-r--r--src/client/ollama.rs13
-rw-r--r--src/client/openai.rs42
-rw-r--r--src/client/qianwen.rs38
-rw-r--r--src/client/vertexai.rs31
12 files changed, 176 insertions, 226 deletions
diff --git a/src/client/claude.rs b/src/client/claude.rs
index 054731c..68ab509 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -15,13 +15,6 @@ use serde_json::{json, Value};
const API_BASE: &str = "https://api.anthropic.com/v1/messages";
-const MODELS: [(&str, usize, &str); 3] = [
- // https://docs.anthropic.com/claude/docs/models-overview
- ("claude-3-opus-20240229", 200000, "text,vision"),
- ("claude-3-sonnet-20240229", 200000, "text,vision"),
- ("claude-3-haiku-20240307", 200000, "text,vision"),
-];
-
#[derive(Debug, Clone, Deserialize)]
pub struct ClaudeConfig {
pub name: Option<String>,
@@ -52,7 +45,16 @@ impl Client for ClaudeClient {
}
impl ClaudeClient {
- list_models_fn!(ClaudeConfig, &MODELS);
+ list_models_fn!(
+ ClaudeConfig,
+ [
+ // https://docs.anthropic.com/claude/docs/models-overview
+ ("claude-3-opus-20240229", "text,vision", 200000, 4096),
+ ("claude-3-sonnet-20240229", "text,vision", 200000, 4096),
+ ("claude-3-haiku-20240307", "text,vision", 200000, 4096),
+ ]
+ );
+
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 1] =
@@ -61,7 +63,7 @@ impl ClaudeClient {
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
let api_key = self.get_api_key().ok();
- let body = build_body(data, &self.model)?;
+ let body = claude_build_body(data, &self.model)?;
let url = API_BASE;
@@ -136,7 +138,7 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHand
Ok(())
}
-fn build_body(data: SendData, model: &Model) -> Result<Value> {
+pub fn claude_build_body(data: SendData, model: &Model) -> Result<Value> {
let SendData {
mut messages,
temperature,
@@ -191,18 +193,16 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
);
}
- let max_tokens = model.max_output_tokens.unwrap_or(4096);
-
let mut body = json!({
"model": &model.name,
- "max_tokens": max_tokens,
"messages": messages,
});
-
- if let Some(system) = system_message {
- body["system"] = system.into();
+ if let Some(v) = system_message {
+ body["system"] = v.into();
+ }
+ if let Some(v) = model.max_output_tokens {
+ body["max_tokens"] = v.into();
}
-
if let Some(v) = temperature {
body["temperature"] = v.into();
}
@@ -218,8 +218,8 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
fn catch_error(data: &Value, status: u16) -> Result<()> {
debug!("Invalid response, status: {status}, data: {data}");
if let Some(error) = data["error"].as_object() {
- if let (Some(type_), Some(message)) = (error["type"].as_str(), error["message"].as_str()) {
- bail!("{message} (type: {type_})");
+ if let (Some(typ), Some(message)) = (error["type"].as_str(), error["message"].as_str()) {
+ bail!("{message} (type: {typ})");
}
}
bail!("Invalid response, status: {status}, data: {data}");
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index cfab0fa..828b799 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -13,12 +13,6 @@ use serde_json::{json, Value};
const API_URL: &str = "https://api.cohere.ai/v1/chat";
-const MODELS: [(&str, usize, &str); 2] = [
- // https://docs.cohere.com/docs/command-r
- ("command-r", 128000, "text"),
- ("command-r-plus", 128000, "text"),
-];
-
#[derive(Debug, Clone, Deserialize, Default)]
pub struct CohereConfig {
pub name: Option<String>,
@@ -49,7 +43,14 @@ impl Client for CohereClient {
}
impl CohereClient {
- list_models_fn!(CohereConfig, &MODELS);
+ list_models_fn!(
+ CohereConfig,
+ [
+ // https://docs.cohere.com/docs/command-r
+ ("command-r", "text", 128000),
+ ("command-r-plus", "text", 128000),
+ ]
+ );
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 1] =
@@ -159,23 +160,22 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
"message": message,
});
- if let Some(preamble) = system_message {
- body["preamble"] = preamble.into();
- }
-
- if let Some(max_tokens) = model.max_output_tokens {
- body["max_tokens"] = max_tokens.into();
+ if let Some(v) = system_message {
+ body["preamble"] = v.into();
}
if !messages.is_empty() {
body["chat_history"] = messages.into();
}
- if let Some(temperature) = temperature {
- body["temperature"] = temperature.into();
+ if let Some(v) = model.max_output_tokens {
+ body["max_tokens"] = v.into();
+ }
+ if let Some(v) = temperature {
+ body["temperature"] = v.into();
}
- if let Some(top_p) = top_p {
- body["p"] = top_p.into();
+ if let Some(v) = top_p {
+ body["p"] = v.into();
}
if stream {
body["stream"] = true.into();
diff --git a/src/client/common.rs b/src/client/common.rs
index 593140d..e003e7e 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -1,4 +1,4 @@
-use super::{openai::OpenAIConfig, ClientConfig, Message, MessageContent, Model, ReplyHandler};
+use super::{openai::OpenAIConfig, ClientConfig, Message, Model, ReplyHandler};
use crate::{
config::{GlobalConfig, Input},
@@ -205,11 +205,18 @@ macro_rules! list_models_fn {
Model::from_config(client_name, &local_config.models)
}
};
- ($config:ident, $models:expr) => {
+ ($config:ident, [$(($name:literal, $capabilities:literal, $max_input_tokens:literal $(, $max_output_tokens:literal)? )),+$(,)?]) => {
pub fn list_models(local_config: &$config) -> Vec<Model> {
let client_name = Self::name(local_config);
if local_config.models.is_empty() {
- Model::from_static(client_name, $models)
+ vec![
+ $(
+ Model::new(client_name, $name)
+ .set_capabilities($capabilities.into())
+ .set_max_input_tokens(Some($max_input_tokens))
+ $(.set_max_output_tokens(Some($max_output_tokens)))?
+ ),+
+ ]
} else {
Model::from_config(client_name, &local_config.models)
}
@@ -402,27 +409,6 @@ where
Ok(())
}
-pub fn patch_system_message(messages: &mut Vec<Message>) {
- if messages[0].role.is_system() {
- let system_message = messages.remove(0);
- if let (Some(message), MessageContent::Text(system_text)) =
- (messages.get_mut(0), system_message.content)
- {
- if let MessageContent::Text(text) = message.content.clone() {
- message.content = MessageContent::Text(format!("{}\n\n{}", system_text, text))
- }
- }
- }
-}
-
-pub fn extract_sytem_message(messages: &mut Vec<Message>) -> Option<String> {
- if messages[0].role.is_system() {
- let system_message = messages.remove(0);
- return Some(system_message.content.to_text());
- }
- None
-}
-
pub async fn json_stream<S, F>(mut stream: S, mut handle: F) -> Result<()>
where
S: Stream<Item = Result<bytes::Bytes, reqwest::Error>> + Unpin,
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index ffe10f0..31ec536 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -18,52 +18,6 @@ use std::{env, sync::Mutex};
const API_BASE: &str = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1";
const ACCESS_TOKEN_URL: &str = "https://aip.baidubce.com/oauth/2.0/token";
-const MODELS: [(&str, &str, usize, isize); 7] = [
- // https://cloud.baidu.com/doc/WENXINWORKSHOP/s/clntwmv7t
- (
- "ernie-4.0-8k",
- "/wenxinworkshop/chat/completions_pro",
- 5120,
- 2048,
- ),
- (
- "ernie-3.5-8k",
- "/wenxinworkshop/chat/ernie-3.5-8k-0205",
- 5120,
- 2048,
- ),
- (
- "ernie-3.5-4k",
- "/wenxinworkshop/chat/ernie-3.5-4k-0205",
- 2048,
- 2048,
- ),
- (
- "ernie-speed-8k",
- "/wenxinworkshop/chat/ernie_speed",
- 7168,
- 2048,
- ),
- (
- "ernie-speed-128k",
- "/wenxinworkshop/chat/ernie-speed-128k",
- 124000,
- 4096,
- ),
- (
- "ernie-lite-8k",
- "/wenxinworkshop/chat/ernie-lite-8k",
- 7168,
- 2048,
- ),
- (
- "ernie-tiny-8k",
- "/wenxinworkshop/chat/ernie-tiny-8k",
- 7168,
- 2048,
- ),
-];
-
lazy_static! {
static ref ACCESS_TOKEN: Mutex<Option<String>> = Mutex::new(None);
}
@@ -101,35 +55,38 @@ impl Client for ErnieClient {
}
impl ErnieClient {
+ list_models_fn!(
+ ErnieConfig,
+ [
+ // https://cloud.baidu.com/doc/WENXINWORKSHOP/s/clntwmv7t
+ ("ernie-4.0-8k", "text", 5120, 2048),
+ ("ernie-3.5-8k", "text", 5120, 2048),
+ ("ernie-3.5-4k", "text", 2048, 2048),
+ ("ernie-speed-8k", "text", 7168, 2048),
+ ("ernie-speed-128k", "text", 124000, 4096),
+ ("ernie-lite-8k", "text", 7168, 2048),
+ ("ernie-tiny-8k", "text", 7168, 2048),
+ ]
+ );
+
pub const PROMPTS: [PromptType<'static>; 2] = [
("api_key", "API Key:", true, PromptKind::String),
("secret_key", "Secret Key:", true, PromptKind::String),
];
- pub fn list_models(local_config: &ErnieConfig) -> Vec<Model> {
- let client_name = Self::name(local_config);
- if local_config.models.is_empty() {
- MODELS
- .into_iter()
- .map(|(name, _, max_input_tokens, max_output_tokens)| {
- Model::new(client_name, name)
- .set_max_input_tokens(Some(max_input_tokens))
- .set_max_output_tokens(Some(max_output_tokens))
- }) // ERNIE tokenizer is different from cl100k_base
- .collect()
- } else {
- Model::from_config(client_name, &local_config.models)
- }
- }
-
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
let body = build_body(data, &self.model);
- let model = &self.model.name;
- let (_, chat_endpoint, _, _) = MODELS
- .iter()
- .find(|(v, _, _, _)| v == model)
- .ok_or_else(|| anyhow!("Miss Model '{}'", self.model.id()))?;
+ let endpoint = match self.model.name.as_str() {
+ "ernie-4.0-8k" => "/wenxinworkshop/chat/completions_pro",
+ "ernie-3.5-8k" => "/wenxinworkshop/chat/ernie-3.5-8k-0205",
+ "ernie-3.5-4k" => "/wenxinworkshop/chat/ernie-3.5-4k-0205",
+ "ernie-speed-8k" => "/wenxinworkshop/chat/ernie_speed",
+ "ernie-speed-128k" => "/wenxinworkshop/chat/ernie-speed-128k",
+ "ernie-lite-8k" => "/wenxinworkshop/chat/ernie-lite-8k",
+ "ernie-tiny-8k" => "/wenxinworkshop/chat/ernie-tiny-8k",
+ _ => bail!("Miss Model '{}'", self.model.id()),
+ };
let access_token = ACCESS_TOKEN
.lock()
@@ -137,7 +94,7 @@ impl ErnieClient {
.clone()
.ok_or_else(|| anyhow!("Failed to load access token"))?;
- let url = format!("{API_BASE}{chat_endpoint}?access_token={access_token}");
+ let url = format!("{API_BASE}{endpoint}?access_token={access_token}");
debug!("Ernie Request: {url} {body}");
@@ -240,15 +197,14 @@ fn build_body(data: SendData, model: &Model) -> Value {
"messages": messages,
});
- if let Some(temperature) = temperature {
- body["temperature"] = temperature.into();
+ if let Some(v) = model.max_output_tokens {
+ body["max_output_tokens"] = v.into();
}
- if let Some(top_p) = top_p {
- body["top_p"] = top_p.into();
+ if let Some(v) = temperature {
+ body["temperature"] = v.into();
}
-
- if let Some(max_output_tokens) = model.max_output_tokens {
- body["max_output_tokens"] = max_output_tokens.into();
+ if let Some(v) = top_p {
+ body["top_p"] = v.into();
}
if stream {
diff --git a/src/client/gemini.rs b/src/client/gemini.rs
index 22f32c5..b930d79 100644
--- a/src/client/gemini.rs
+++ b/src/client/gemini.rs
@@ -12,13 +12,6 @@ use serde::Deserialize;
const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta/models/";
-const MODELS: [(&str, usize, &str); 3] = [
- // https://ai.google.dev/models/gemini
- ("gemini-1.0-pro-latest", 30720, "text"),
- ("gemini-1.0-pro-vision-latest", 12288, "text,vision"),
- ("gemini-1.5-pro-latest", 1048576, "text,vision"),
-];
-
#[derive(Debug, Clone, Deserialize, Default)]
pub struct GeminiConfig {
pub name: Option<String>,
@@ -50,7 +43,15 @@ impl Client for GeminiClient {
}
impl GeminiClient {
- list_models_fn!(GeminiConfig, &MODELS);
+ list_models_fn!(
+ GeminiConfig,
+ [
+ // https://ai.google.dev/models/gemini
+ ("gemini-1.0-pro-latest", "text", 30720),
+ ("gemini-1.0-pro-vision-latest", "text,vision", 12288),
+ ("gemini-1.5-pro-latest", "text,vision", 1048576),
+ ]
+ );
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 1] =
diff --git a/src/client/message.rs b/src/client/message.rs
index f9978c5..9aff64a 100644
--- a/src/client/message.rs
+++ b/src/client/message.rs
@@ -114,6 +114,27 @@ pub struct ImageUrl {
pub url: String,
}
+pub fn patch_system_message(messages: &mut Vec<Message>) {
+ if messages[0].role.is_system() {
+ let system_message = messages.remove(0);
+ if let (Some(message), MessageContent::Text(system_text)) =
+ (messages.get_mut(0), system_message.content)
+ {
+ if let MessageContent::Text(text) = message.content.clone() {
+ message.content = MessageContent::Text(format!("{}\n\n{}", system_text, text))
+ }
+ }
+ }
+}
+
+pub fn extract_sytem_message(messages: &mut Vec<Message>) -> Option<String> {
+ if messages[0].role.is_system() {
+ let system_message = messages.remove(0);
+ return Some(system_message.content.to_text());
+ }
+ None
+}
+
#[cfg(test)]
mod tests {
use super::*;
diff --git a/src/client/mistral.rs b/src/client/mistral.rs
index 9dbd831..23bf1f2 100644
--- a/src/client/mistral.rs
+++ b/src/client/mistral.rs
@@ -10,13 +10,6 @@ use serde::Deserialize;
const API_URL: &str = "https://api.mistral.ai/v1/chat/completions";
-const MODELS: [(&str, usize, &str); 3] = [
- // https://docs.mistral.ai/platform/endpoints/
- ("open-mixtral-8x22b", 64000, "text"),
- ("mistral-small-latest", 32000, "text"),
- ("mistral-large-latest", 32000, "text"),
-];
-
#[derive(Debug, Clone, Deserialize)]
pub struct MistralConfig {
pub name: Option<String>,
@@ -29,7 +22,15 @@ pub struct MistralConfig {
openai_compatible_client!(MistralClient);
impl MistralClient {
- list_models_fn!(MistralConfig, &MODELS);
+ list_models_fn!(
+ MistralConfig,
+ [
+ // https://docs.mistral.ai/platform/endpoints/
+ ("open-mixtral-8x22b", "text", 64000),
+ ("mistral-small-latest", "text", 32000),
+ ("mistral-large-latest", "text", 32000),
+ ]
+ );
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 1] =
diff --git a/src/client/model.rs b/src/client/model.rs
index b97c244..3f2cbdd 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -49,17 +49,6 @@ impl Model {
.collect()
}
- pub fn from_static(client_name: &str, models: &[(&str, usize, &str)]) -> Vec<Self> {
- models
- .iter()
- .map(|(name, max_input_tokens, capabilities)| {
- Model::new(client_name, name)
- .set_capabilities((*capabilities).into())
- .set_max_input_tokens(Some(*max_input_tokens))
- })
- .collect()
- }
-
pub fn find(models: &[Self], value: &str) -> Option<Self> {
let mut model = None;
let (client_name, model_name) = match value.split_once(':') {
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index 434e487..1394ced 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -179,15 +179,14 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
"options": {},
});
- if let Some(num_predict) = model.max_output_tokens {
- body["options"]["num_predict"] = num_predict.into();
+ if let Some(v) = model.max_output_tokens {
+ body["options"]["num_predict"] = v.into();
}
-
- if let Some(temperature) = temperature {
- body["options"]["temperature"] = temperature.into();
+ if let Some(v) = temperature {
+ body["options"]["temperature"] = v.into();
}
- if let Some(top_p) = top_p {
- body["options"]["top_p"] = top_p.into();
+ if let Some(v) = top_p {
+ body["options"]["top_p"] = v.into();
}
Ok(body)
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 2c5d99b..37da878 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -12,18 +12,6 @@ use serde_json::{json, Value};
const API_BASE: &str = "https://api.openai.com/v1";
-const MODELS: [(&str, usize, &str); 8] = [
- // https://platform.openai.com/docs/models
- ("gpt-3.5-turbo", 16385, "text"),
- ("gpt-3.5-turbo-1106", 16385, "text"),
- ("gpt-4-turbo", 128000, "text,vision"),
- ("gpt-4-turbo-preview", 128000, "text"),
- ("gpt-4-1106-preview", 128000, "text"),
- ("gpt-4-vision-preview", 128000, "text,vision"),
- ("gpt-4", 8192, "text"),
- ("gpt-4-32k", 32768, "text"),
-];
-
#[derive(Debug, Clone, Deserialize, Default)]
pub struct OpenAIConfig {
pub name: Option<String>,
@@ -38,7 +26,20 @@ pub struct OpenAIConfig {
openai_compatible_client!(OpenAIClient);
impl OpenAIClient {
- list_models_fn!(OpenAIConfig, &MODELS);
+ list_models_fn!(
+ OpenAIConfig,
+ [
+ // https://platform.openai.com/docs/models
+ ("gpt-3.5-turbo", "text", 16385),
+ ("gpt-3.5-turbo-1106", "text", 16385),
+ ("gpt-4-turbo", "text,vision", 128000),
+ ("gpt-4-turbo-preview", "text", 128000),
+ ("gpt-4-1106-preview", "text", 128000),
+ ("gpt-4-vision-preview", "text,vision", 128000, 4096),
+ ("gpt-4", "text", 8192),
+ ("gpt-4-32k", "text", 32768),
+ ]
+ );
config_get_fn!(api_key, get_api_key);
config_get_fn!(api_base, get_api_base);
@@ -139,17 +140,14 @@ pub fn openai_build_body(data: SendData, model: &Model) -> Value {
"messages": messages,
});
- if let Some(max_tokens) = model.max_output_tokens {
- body["max_tokens"] = max_tokens.into();
- } else if model.name == "gpt-4-vision-preview" {
- // The default max_tokens of gpt-4-vision-preview is only 16, we need to make it larger
- body["max_tokens"] = 4096.into();
+ if let Some(v) = model.max_output_tokens {
+ body["max_tokens"] = v.into();
}
- if let Some(temperature) = temperature {
- body["temperature"] = temperature.into();
+ if let Some(v) = temperature {
+ body["temperature"] = v.into();
}
- if let Some(top_p) = top_p {
- body["top_p"] = top_p.into();
+ if let Some(v) = top_p {
+ body["top_p"] = v.into();
}
if stream {
body["stream"] = true.into();
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index c5a72e0..52548d8 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -24,17 +24,6 @@ const API_URL: &str =
const API_URL_VL: &str =
"https://dashscope.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation";
-const MODELS: [(&str, usize, &str); 6] = [
- // https://help.aliyun.com/zh/dashscope/developer-reference/api-details
- ("qwen-turbo", 6000, "text"),
- ("qwen-plus", 30000, "text"),
- ("qwen-max", 6000, "text"),
- ("qwen-max-longcontext", 28000, "text"),
- // https://help.aliyun.com/zh/dashscope/developer-reference/tongyi-qianwen-vl-plus-api
- ("qwen-vl-plus", 0, "text,vision"),
- ("qwen-vl-max", 0, "text,vision"),
-];
-
#[derive(Debug, Clone, Deserialize, Default)]
pub struct QianwenConfig {
pub name: Option<String>,
@@ -73,7 +62,19 @@ impl Client for QianwenClient {
}
impl QianwenClient {
- list_models_fn!(QianwenConfig, &MODELS);
+ list_models_fn!(
+ QianwenConfig,
+ [
+ // https://help.aliyun.com/zh/dashscope/developer-reference/api-details
+ ("qwen-turbo", "text", 6000),
+ ("qwen-plus", "text", 30000),
+ ("qwen-max", "text", 6000),
+ ("qwen-max-longcontext", "text", 28000),
+ // https://help.aliyun.com/zh/dashscope/developer-reference/tongyi-qianwen-vl-plus-api
+ ("qwen-vl-plus", "text,vision", 0),
+ ("qwen-vl-max", "text,vision", 0),
+ ]
+ );
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 1] =
@@ -211,15 +212,14 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
parameters["incremental_output"] = true.into();
}
- if let Some(max_tokens) = model.max_output_tokens {
- parameters["max_tokens"] = max_tokens.into();
+ if let Some(v) = model.max_output_tokens {
+ parameters["max_tokens"] = v.into();
}
-
- if let Some(temperature) = temperature {
- parameters["temperature"] = temperature.into();
+ if let Some(v) = temperature {
+ parameters["temperature"] = v.into();
}
- if let Some(top_p) = top_p {
- parameters["top_p"] = top_p.into();
+ if let Some(v) = top_p {
+ parameters["top_p"] = v.into();
}
let body = json!({
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)