summaryrefslogtreecommitdiffstats
path: root/src/client/azure_openai.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-01-13 19:52:07 +0800
committerGitHub <noreply@github.com>2024-01-13 19:52:07 +0800
commitfe35cfd9419302f01baf9672493c0b0a4b41d889 (patch)
tree94e763745fb7989c97af39cc1dfb44250440a5eb /src/client/azure_openai.rs
parent4e99df4c1bd4028a77251bdb00ff23c664372b5f (diff)
downloadaichat-fe35cfd9419302f01baf9672493c0b0a4b41d889.tar.gz
feat: supports model capabilities (#297)
1. automatically switch to the model that has the necessary capabilities. 2. throw an error if the client does not have a model with the necessary capabilities
Diffstat (limited to 'src/client/azure_openai.rs')
-rw-r--r--src/client/azure_openai.rs11
1 files changed, 3 insertions, 8 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index 4700ad7..5c9f3ef 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -1,5 +1,5 @@
use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS};
-use super::{AzureOpenAIClient, ExtraConfig, PromptType, SendData, Model};
+use super::{AzureOpenAIClient, ExtraConfig, Model, ModelConfig, PromptType, SendData};
use crate::utils::PromptKind;
@@ -13,16 +13,10 @@ pub struct AzureOpenAIConfig {
pub name: Option<String>,
pub api_base: Option<String>,
pub api_key: Option<String>,
- pub models: Vec<AzureOpenAIModel>,
+ pub models: Vec<ModelConfig>,
pub extra: Option<ExtraConfig>,
}
-#[derive(Debug, Clone, Deserialize)]
-pub struct AzureOpenAIModel {
- name: String,
- max_tokens: Option<usize>,
-}
-
openai_compatible_client!(AzureOpenAIClient);
impl AzureOpenAIClient {
@@ -50,6 +44,7 @@ impl AzureOpenAIClient {
.map(|v| {
Model::new(client_name, &v.name)
.set_max_tokens(v.max_tokens)
+ .set_capabilities(v.capabilities)
.set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS)
})
.collect()