summaryrefslogtreecommitdiffstats
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
parent887bf0a744646c6d9cfb361ba0c929d4a4342337 (diff)
downloadaichat-4380b4f20bed88221a3dcf90d2f6e8b23b539794.tar.gz
refactor: rename azure to azure_openai, improve register_client! (#208)
-rw-r--r--config.example.yaml2
-rw-r--r--src/client/azure_openai.rs (renamed from src/client/azure.rs)14
-rw-r--r--src/client/common.rs10
-rw-r--r--src/client/mod.rs11
-rw-r--r--src/config/mod.rs2
5 files changed, 22 insertions, 17 deletions
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_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())
}