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) --- config.example.yaml | 2 +- src/client/azure.rs | 73 ---------------------------------------------- src/client/azure_openai.rs | 73 ++++++++++++++++++++++++++++++++++++++++++++++ src/client/common.rs | 10 +++---- src/client/mod.rs | 11 +++++-- src/config/mod.rs | 2 +- 6 files changed, 88 insertions(+), 83 deletions(-) delete mode 100644 src/client/azure.rs create mode 100644 src/client/azure_openai.rs diff --git a/config.example.yaml b/config.example.yaml index 9c3a1b0..42c5f6e 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -22,7 +22,7 @@ clients: organization_id: # See https://learn.microsoft.com/en-us/azure/ai-services/openai/chatgpt-quickstart - - type: azure + - type: azure-openai api_base: https://RESOURCE.openai.azure.com api_key: xxx models: 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, - pub api_base: Option, - pub api_key: Option, - pub models: Vec, - pub extra: Option, -} - -#[derive(Debug, Clone, Deserialize)] -pub struct AzureModel { - name: String, - max_tokens: Option, -} - -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 { - 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) - } -} 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) + } +} diff --git a/src/client/common.rs b/src/client/common.rs index 464e46a..d43f1b6 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -20,7 +20,7 @@ use tokio::time::sleep; #[macro_export] macro_rules! register_client { ( - $(($module:ident, $name:literal, $config_key:ident, $config:ident, $client:ident),)+ + $(($module:ident, $name:literal, $config:ident, $client:ident),)+ ) => { $( mod $module; @@ -34,7 +34,7 @@ macro_rules! register_client { pub enum ClientConfig { $( #[serde(rename = $name)] - $config_key($config), + $config($config), )+ #[serde(other)] Unknown, @@ -55,7 +55,7 @@ macro_rules! register_client { pub fn init(global_config: $crate::config::GlobalConfig) -> Option> { let model = global_config.read().model.clone(); let config = { - if let ClientConfig::$config_key(c) = &global_config.read().clients[model.client_index] { + if let ClientConfig::$config(c) = &global_config.read().clients[model.client_index] { c.clone() } else { return None; @@ -107,7 +107,7 @@ macro_rules! register_client { .iter() .enumerate() .flat_map(|(i, v)| match v { - $(ClientConfig::$config_key(c) => $client::list_models(c, i),)+ + $(ClientConfig::$config(c) => $client::list_models(c, i),)+ ClientConfig::Unknown => vec![], }) .collect() @@ -258,7 +258,7 @@ pub trait Client { impl Default for ClientConfig { fn default() -> Self { - Self::OpenAI(OpenAIConfig::default()) + Self::OpenAIConfig(OpenAIConfig::default()) } } diff --git a/src/client/mod.rs b/src/client/mod.rs index 55f0ed0..f124b62 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -8,7 +8,12 @@ pub use message::*; pub use model::*; register_client!( - (openai, "openai", OpenAI, OpenAIConfig, OpenAIClient), - (localai, "localai", LocalAI, LocalAIConfig, LocalAIClient), - (azure, "azure", Azure, AzureConfig, AzureClient), + (openai, "openai", OpenAIConfig, OpenAIClient), + (localai, "localai", LocalAIConfig, LocalAIClient), + ( + azure_openai, + "azure-openai", + AzureOpenAIConfig, + AzureOpenAIClient + ), ); diff --git a/src/config/mod.rs b/src/config/mod.rs index 4749c2a..1720772 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -717,7 +717,7 @@ impl Config { } } - if let Some(ClientConfig::OpenAI(client_config)) = self.clients.get_mut(0) { + if let Some(ClientConfig::OpenAIConfig(client_config)) = self.clients.get_mut(0) { if let Some(api_key) = value.get("api_key").and_then(|v| v.as_str()) { client_config.api_key = Some(api_key.to_string()) } -- cgit v1.2.3