summaryrefslogtreecommitdiffstats
path: root/src/client/openai.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/openai.rs
parent1cc89eff514669d6459f62cc4eed8e0b34d8ef0c (diff)
downloadaichat-d1aafa11153ab689c21c2c57c47da52337d8e8d1.tar.gz
feat: customize model's max_output_tokens (#428)
Diffstat (limited to 'src/client/openai.rs')
-rw-r--r--src/client/openai.rs13
1 files changed, 7 insertions, 6 deletions
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 797ee98..535e739 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -58,7 +58,7 @@ impl OpenAIClient {
let api_key = self.get_api_key()?;
let api_base = self.get_api_base().unwrap_or_else(|_| API_BASE.to_string());
- let body = openai_build_body(data, self.model.name.clone());
+ let body = openai_build_body(data, &self.model);
let url = format!("{api_base}/chat/completions");
@@ -139,7 +139,7 @@ pub async fn openai_send_message_streaming(
Ok(())
}
-pub fn openai_build_body(data: SendData, model: String) -> Value {
+pub fn openai_build_body(data: SendData, model: &Model) -> Value {
let SendData {
messages,
temperature,
@@ -147,15 +147,16 @@ pub fn openai_build_body(data: SendData, model: String) -> Value {
} = data;
let mut body = json!({
- "model": model,
+ "model": &model.name,
"messages": messages,
});
- // The default max_tokens of gpt-4-vision-preview is only 16, we need to make it larger
- if model == "gpt-4-vision-preview" {
+ if let Some(max_tokens) = model.max_output_tokens {
+ body["max_tokens"] = json!(max_tokens);
+ } else if model.name == "gpt-4-vision-preview" {
+ // The default max_tokens of gpt-4-vision-preview is only 16, we need to make it larger
body["max_tokens"] = json!(4096);
}
-
if let Some(v) = temperature {
body["temperature"] = v.into();
}