diff options
| author | sigoden <sigoden@gmail.com> | 2023-11-04 06:28:40 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-11-04 06:28:40 +0800 |
| commit | 4380b4f20bed88221a3dcf90d2f6e8b23b539794 (patch) | |
| tree | 16e609e84dcacbb9f8c15ba23027a136bf5df5f4 /src | |
| parent | 887bf0a744646c6d9cfb361ba0c929d4a4342337 (diff) | |
| download | aichat-4380b4f20bed88221a3dcf90d2f6e8b23b539794.tar.gz | |
refactor: rename azure to azure_openai, improve register_client! (#208)
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/azure_openai.rs (renamed from src/client/azure.rs) | 14 | ||||
| -rw-r--r-- | src/client/common.rs | 10 | ||||
| -rw-r--r-- | src/client/mod.rs | 11 | ||||
| -rw-r--r-- | src/config/mod.rs | 2 |
4 files changed, 21 insertions, 16 deletions
diff --git a/src/client/azure.rs b/src/client/azure_openai.rs index bae7851..815413e 100644 --- a/src/client/azure.rs +++ b/src/client/azure_openai.rs @@ -1,5 +1,5 @@ use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS}; -use super::{AzureClient, ExtraConfig, PromptType, SendData, Model}; +use super::{AzureOpenAIClient, ExtraConfig, PromptType, SendData, Model}; use crate::utils::PromptKind; @@ -9,23 +9,23 @@ use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; #[derive(Debug, Clone, Deserialize)] -pub struct AzureConfig { +pub struct AzureOpenAIConfig { pub name: Option<String>, pub api_base: Option<String>, pub api_key: Option<String>, - pub models: Vec<AzureModel>, + pub models: Vec<AzureOpenAIModel>, pub extra: Option<ExtraConfig>, } #[derive(Debug, Clone, Deserialize)] -pub struct AzureModel { +pub struct AzureOpenAIModel { name: String, max_tokens: Option<usize>, } -openai_compatible_client!(AzureClient); +openai_compatible_client!(AzureOpenAIClient); -impl AzureClient { +impl AzureOpenAIClient { config_get_fn!(api_base, get_api_base); config_get_fn!(api_key, get_api_key); @@ -41,7 +41,7 @@ impl AzureClient { ), ]; - pub fn list_models(local_config: &AzureConfig, client_index: usize) -> Vec<Model> { + pub fn list_models(local_config: &AzureOpenAIConfig, client_index: usize) -> Vec<Model> { let client_name = Self::name(local_config); local_config 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<Box<dyn Client>> { 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()) } |
