From ba3bcfd67c1d6fea5d3d3c5908c975682ee7909b Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 22 May 2024 21:29:23 +0800 Subject: feat: allow patching req body with client config (#534) --- src/client/vertexai.rs | 13 +++++-------- 1 file changed, 5 insertions(+), 8 deletions(-) (limited to 'src/client/vertexai.rs') diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index 21fc154..67a4b21 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -1,7 +1,7 @@ use super::{ access_token::*, catch_error, json_stream, message::*, patch_system_message, Client, - CompletionOutput, ExtraConfig, Model, ModelData, PromptAction, PromptKind, SendData, - SseHandler, ToolCall, VertexAIClient, + CompletionOutput, ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, + SendData, SseHandler, ToolCall, VertexAIClient, }; use anyhow::{anyhow, bail, Context, Result}; @@ -22,6 +22,7 @@ pub struct VertexAIConfig { pub safety_settings: Option, #[serde(default)] pub models: Vec, + pub patches: Option, pub extra: Option, } @@ -47,7 +48,8 @@ impl VertexAIClient { }; let url = format!("{base_url}/google/models/{}:{func}", self.model.name()); - let body = gemini_build_body(data, &self.model, self.config.safety_settings.clone())?; + let mut body = gemini_build_body(data, &self.model)?; + self.patch_request_body(&mut body); debug!("VertexAI Request: {url} {body}"); @@ -178,7 +180,6 @@ fn gemini_extract_completion_text(data: &Value) -> Result { pub(crate) fn gemini_build_body( data: SendData, model: &Model, - safety_settings: Option, ) -> Result { let SendData { mut messages, @@ -259,10 +260,6 @@ pub(crate) fn gemini_build_body( let mut body = json!({ "contents": contents, "generationConfig": {} }); - if let Some(safety_settings) = safety_settings { - body["safetySettings"] = safety_settings; - } - if let Some(v) = model.max_tokens_param() { body["generationConfig"]["maxOutputTokens"] = v.into(); } -- cgit v1.2.3