summaryrefslogtreecommitdiffstats
path: root/src/client/azure_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/azure_openai.rs
parent1cc89eff514669d6459f62cc4eed8e0b34d8ef0c (diff)
downloadaichat-d1aafa11153ab689c21c2c57c47da52337d8e8d1.tar.gz
feat: customize model's max_output_tokens (#428)
Diffstat (limited to 'src/client/azure_openai.rs')
-rw-r--r--src/client/azure_openai.rs15
1 files changed, 3 insertions, 12 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index 3d0cf3f..1726bbe 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -1,5 +1,5 @@
use super::openai::openai_build_body;
-use super::{AzureOpenAIClient, ExtraConfig, Model, ModelConfig, PromptType, SendData};
+use super::{convert_models, AzureOpenAIClient, ExtraConfig, Model, ModelConfig, PromptType, SendData};
use crate::utils::PromptKind;
@@ -37,23 +37,14 @@ impl AzureOpenAIClient {
pub fn list_models(local_config: &AzureOpenAIConfig) -> Vec<Model> {
let client_name = Self::name(local_config);
-
- local_config
- .models
- .iter()
- .map(|v| {
- Model::new(client_name, &v.name)
- .set_max_input_tokens(v.max_input_tokens)
- .set_capabilities(v.capabilities)
- })
- .collect()
+ convert_models(client_name, &local_config.models)
}
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
let api_base = self.get_api_base()?;
let api_key = self.get_api_key()?;
- let mut body = openai_build_body(data, self.model.name.clone());
+ let mut body = openai_build_body(data, &self.model);
self.model.merge_extra_fields(&mut body);
let url = format!(