summaryrefslogtreecommitdiffstats
path: root/src/client/vertexai.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-23 16:46:48 +0800
committerGitHub <noreply@github.com>2024-04-23 16:46:48 +0800
commitd1aafa11153ab689c21c2c57c47da52337d8e8d1 (patch)
treedc20dc033e9d376aab09941835a842b22fe32c02 /src/client/vertexai.rs
parent1cc89eff514669d6459f62cc4eed8e0b34d8ef0c (diff)
downloadaichat-d1aafa11153ab689c21c2c57c47da52337d8e8d1.tar.gz
feat: customize model's max_output_tokens (#428)
Diffstat (limited to 'src/client/vertexai.rs')
-rw-r--r--src/client/vertexai.rs16
1 files changed, 9 insertions, 7 deletions
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 88035ec..eceeb2c 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -81,9 +81,9 @@ impl VertexAIClient {
let block_threshold = self.config.block_threshold.clone();
- let body = build_body(data, self.model.name.clone(), block_threshold)?;
+ let body = build_body(data, &self.model, block_threshold)?;
- let model = self.model.name.clone();
+ let model = &self.model.name;
let url = format!("{api_base}/{}:{}", model, func);
@@ -176,7 +176,7 @@ fn check_error(data: &Value) -> Result<()> {
pub(crate) fn build_body(
data: SendData,
- _model: String,
+ model: &Model,
block_threshold: Option<String>,
) -> Result<Value> {
let SendData {
@@ -228,7 +228,7 @@ pub(crate) fn build_body(
);
}
- let mut body = json!({ "contents": contents });
+ let mut body = json!({ "contents": contents, "generationConfig": {} });
if let Some(block_threshold) = block_threshold {
body["safetySettings"] = json!([
@@ -239,10 +239,12 @@ 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(temperature) = temperature {
- body["generationConfig"] = json!({
- "temperature": temperature,
- });
+ body["generationConfig"]["temperature"] = temperature.into();
}
Ok(body)