From 4380b4f20bed88221a3dcf90d2f6e8b23b539794 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 4 Nov 2023 06:28:40 +0800 Subject: refactor: rename azure to azure_openai, improve register_client! (#208) --- src/client/azure_openai.rs | 73 ++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 73 insertions(+) create mode 100644 src/client/azure_openai.rs (limited to 'src/client/azure_openai.rs') diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs new file mode 100644 index 0000000..815413e --- /dev/null +++ b/src/client/azure_openai.rs @@ -0,0 +1,73 @@ +use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS}; +use super::{AzureOpenAIClient, 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 AzureOpenAIConfig { + pub name: Option, + pub api_base: Option, + pub api_key: Option, + pub models: Vec, + pub extra: Option, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct AzureOpenAIModel { + name: String, + max_tokens: Option, +} + +openai_compatible_client!(AzureOpenAIClient); + +impl AzureOpenAIClient { + 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: &AzureOpenAIConfig, client_index: usize) -> Vec { + 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 { + 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) + } +} -- cgit v1.2.3