summaryrefslogtreecommitdiffstats
path: root/src/client/qianwen.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/qianwen.rs
parent1cc89eff514669d6459f62cc4eed8e0b34d8ef0c (diff)
downloadaichat-d1aafa11153ab689c21c2c57c47da52337d8e8d1.tar.gz
feat: customize model's max_output_tokens (#428)
Diffstat (limited to 'src/client/qianwen.rs')
-rw-r--r--src/client/qianwen.rs10
1 files changed, 7 insertions, 3 deletions
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index 2034736..6225b96 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -97,7 +97,7 @@ impl QianwenClient {
true => API_URL_VL,
false => API_URL,
};
- let (body, has_upload) = build_body(data, self.model.name.clone(), is_vl)?;
+ let (body, has_upload) = build_body(data, &self.model, is_vl)?;
debug!("Qianwen Request: {url} {body}");
@@ -180,7 +180,7 @@ fn check_error(data: &Value) -> Result<()> {
Ok(())
}
-fn build_body(data: SendData, model: String, is_vl: bool) -> Result<(Value, bool)> {
+fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool)> {
let SendData {
messages,
temperature,
@@ -233,6 +233,10 @@ fn build_body(data: SendData, model: String, is_vl: bool) -> Result<(Value, bool
parameters["incremental_output"] = true.into();
}
+ if let Some(max_tokens) = model.max_output_tokens {
+ parameters["max_tokens"] = max_tokens.into();
+ }
+
if let Some(v) = temperature {
parameters["temperature"] = v.into();
}
@@ -240,7 +244,7 @@ fn build_body(data: SendData, model: String, is_vl: bool) -> Result<(Value, bool
};
let body = json!({
- "model": model,
+ "model": &model.name,
"input": input,
"parameters": parameters
});