summaryrefslogtreecommitdiffstats
path: root/src/client/vertexai.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/client/vertexai.rs')
-rw-r--r--src/client/vertexai.rs13
1 files changed, 5 insertions, 8 deletions
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<Value>,
#[serde(default)]
pub models: Vec<ModelData>,
+ pub patches: Option<ModelPatches>,
pub extra: Option<ExtraConfig>,
}
@@ -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<CompletionOutput> {
pub(crate) fn gemini_build_body(
data: SendData,
model: &Model,
- safety_settings: Option<Value>,
) -> Result<Value> {
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();
}