summaryrefslogtreecommitdiffstats
path: root/src/client/azure.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-11-04 06:28:40 +0800
committerGitHub <noreply@github.com>2023-11-04 06:28:40 +0800
commit4380b4f20bed88221a3dcf90d2f6e8b23b539794 (patch)
tree16e609e84dcacbb9f8c15ba23027a136bf5df5f4 /src/client/azure.rs
parent887bf0a744646c6d9cfb361ba0c929d4a4342337 (diff)
downloadaichat-4380b4f20bed88221a3dcf90d2f6e8b23b539794.tar.gz
refactor: rename azure to azure_openai, improve register_client! (#208)
Diffstat (limited to 'src/client/azure.rs')
-rw-r--r--src/client/azure.rs73
1 files changed, 0 insertions, 73 deletions
diff --git a/src/client/azure.rs b/src/client/azure.rs
deleted file mode 100644
index bae7851..0000000
--- a/src/client/azure.rs
+++ /dev/null
@@ -1,73 +0,0 @@
-use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS};
-use super::{AzureClient, ExtraConfig, PromptType, SendData, Model};
-
-use crate::utils::PromptKind;
-
-use anyhow::Result;
-use async_trait::async_trait;
-use reqwest::{Client as ReqwestClient, RequestBuilder};
-use serde::Deserialize;
-
-#[derive(Debug, Clone, Deserialize)]
-pub struct AzureConfig {
- pub name: Option<String>,
- pub api_base: Option<String>,
- pub api_key: Option<String>,
- pub models: Vec<AzureModel>,
- pub extra: Option<ExtraConfig>,
-}
-
-#[derive(Debug, Clone, Deserialize)]
-pub struct AzureModel {
- name: String,
- max_tokens: Option<usize>,
-}
-
-openai_compatible_client!(AzureClient);
-
-impl AzureClient {
- config_get_fn!(api_base, get_api_base);
- config_get_fn!(api_key, get_api_key);
-
- pub const PROMPTS: [PromptType<'static>; 4] = [
- ("api_base", "API Base:", true, PromptKind::String),
- ("api_key", "API Key:", true, PromptKind::String),
- ("models[].name", "Model Name:", true, PromptKind::String),
- (
- "models[].max_tokens",
- "Max Tokens:",
- true,
- PromptKind::Integer,
- ),
- ];
-
- pub fn list_models(local_config: &AzureConfig, client_index: usize) -> Vec<Model> {
- let client_name = Self::name(local_config);
-
- local_config
- .models
- .iter()
- .map(|v| {
- Model::new(client_index, client_name, &v.name)
- .set_max_tokens(v.max_tokens)
- .set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS)
- })
- .collect()
- }
-
- 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 body = openai_build_body(data, self.model.llm_name.clone());
-
- let url = format!(
- "{}/openai/deployments/{}/chat/completions?api-version=2023-05-15",
- &api_base, self.model.llm_name
- );
-
- let builder = client.post(url).header("api-key", api_key).json(&body);
-
- Ok(builder)
- }
-}