From 7f05dc1a4af868b64906666c65020838cf2aa985 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 25 Mar 2024 21:06:35 +0800 Subject: feat: support customizing gemini safeSettings (#375) --- src/client/gemini.rs | 5 ++++- src/client/vertexai.rs | 42 ++++++++++++++++++++++++++---------------- 2 files changed, 30 insertions(+), 17 deletions(-) (limited to 'src/client') diff --git a/src/client/gemini.rs b/src/client/gemini.rs index 60c1e3b..39371fd 100644 --- a/src/client/gemini.rs +++ b/src/client/gemini.rs @@ -23,6 +23,7 @@ const TOKENS_COUNT_FACTORS: TokensCountFactors = (5, 2); pub struct GeminiConfig { pub name: Option, pub api_key: Option, + pub block_threshold: Option, pub extra: Option, } @@ -73,7 +74,9 @@ impl GeminiClient { false => "generateContent", }; - let body = build_body(data, self.model.name.clone())?; + let block_threshold = self.config.block_threshold.clone(); + + let body = build_body(data, self.model.name.clone(), block_threshold)?; let model = self.model.name.clone(); diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index ba43796..80e58e7 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -14,13 +14,13 @@ use serde::Deserialize; use serde_json::{json, Value}; use std::path::PathBuf; -const MODELS: [(&str, usize, &str); 2] = [ +const MODELS: [(&str, usize, &str); 5] = [ // 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.0-ultra", 8192, "text"), - // ("gemini-1.0-ultra-vision", 8192, "text,vision"), - // ("gemini-1.5-pro", 1000000, "text"), + ("gemini-1.0-ultra", 8192, "text"), + ("gemini-1.0-ultra-vision", 8192, "text,vision"), + ("gemini-1.5-pro", 1000000, "text"), ]; const TOKENS_COUNT_FACTORS: TokensCountFactors = (5, 2); @@ -32,6 +32,7 @@ pub struct VertexAIConfig { pub name: Option, pub api_base: Option, pub adc_file: Option, + pub block_threshold: Option, pub extra: Option, } @@ -84,7 +85,9 @@ impl VertexAIClient { false => "generateContent", }; - let body = build_body(data, self.model.name.clone())?; + let block_threshold = self.config.block_threshold.clone(); + + let body = build_body(data, self.model.name.clone(), block_threshold)?; let model = self.model.name.clone(); @@ -106,7 +109,9 @@ impl VertexAIClient { let (token, expires_in) = fetch_access_token(&client, &self.config.adc_file) .await .with_context(|| "Failed to fetch access token")?; - let expires_at = Utc::now() + Duration::try_seconds(expires_in).ok_or_else(|| anyhow!("Failed to parse expires_in of access_token"))?; + let expires_at = Utc::now() + + Duration::try_seconds(expires_in) + .ok_or_else(|| anyhow!("Failed to parse expires_in of access_token"))?; unsafe { ACCESS_TOKEN = (token, expires_at.timestamp()) }; } Ok(()) @@ -208,7 +213,11 @@ fn check_error(data: &Value) -> Result<()> { } } -pub(crate) fn build_body(data: SendData, _model: String) -> Result { +pub(crate) fn build_body( + data: SendData, + _model: String, + block_threshold: Option, +) -> Result { let SendData { mut messages, temperature, @@ -258,15 +267,16 @@ pub(crate) fn build_body(data: SendData, _model: String) -> Result { ); } - let mut body = json!({ - "contents": contents, - "safetySettings":[ - {"category":"HARM_CATEGORY_HARASSMENT","threshold":"BLOCK_ONLY_HIGH"}, - {"category":"HARM_CATEGORY_HATE_SPEECH","threshold":"BLOCK_ONLY_HIGH"}, - {"category":"HARM_CATEGORY_SEXUALLY_EXPLICIT","threshold":"BLOCK_ONLY_HIGH"}, - {"category":"HARM_CATEGORY_DANGEROUS_CONTENT","threshold":"BLOCK_ONLY_HIGH"} - ] - }); + let mut body = json!({ "contents": contents }); + + if let Some(block_threshold) = block_threshold { + body["safetySettings"] = json!([ + {"category":"HARM_CATEGORY_HARASSMENT","threshold":block_threshold}, + {"category":"HARM_CATEGORY_HATE_SPEECH","threshold":block_threshold}, + {"category":"HARM_CATEGORY_SEXUALLY_EXPLICIT","threshold":block_threshold}, + {"category":"HARM_CATEGORY_DANGEROUS_CONTENT","threshold":block_threshold} + ]); + } if let Some(temperature) = temperature { body["generationConfig"] = json!({ -- cgit v1.2.3