diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-22 21:29:23 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-22 21:29:23 +0800 |
| commit | ba3bcfd67c1d6fea5d3d3c5908c975682ee7909b (patch) | |
| tree | 29811c719f131f946cd37634507885ee8c63534d /src/client/vertexai.rs | |
| parent | 91a06543b24733cf578f3f2d4cb0884e2b18cf2f (diff) | |
| download | aichat-ba3bcfd67c1d6fea5d3d3c5908c975682ee7909b.tar.gz | |
feat: allow patching req body with client config (#534)
Diffstat (limited to 'src/client/vertexai.rs')
| -rw-r--r-- | src/client/vertexai.rs | 13 |
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(); } |
