diff options
| author | sigoden <sigoden@gmail.com> | 2023-11-01 10:28:54 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-11-01 10:28:54 +0800 |
| commit | 7f2210dbcab017ee30e8f95771e68494b4860c59 (patch) | |
| tree | 7a0c7bdc8f6930e262bf8cfdc8ffbd57aa843e40 /src/client/azure_openai.rs | |
| parent | 64202758ec8693dd18951c554728bb5c6b7e5eeb (diff) | |
| download | aichat-7f2210dbcab017ee30e8f95771e68494b4860c59.tar.gz | |
refactor: clients/* and config.rs (#193)
- add register_clients macro to make it easier to add a new client
- no create_client_config, just add const PROMPTS
- move ModelInfo from clients/ to config/
- model's max_tokens are optional
- improve code quanity on config/mod.rs
- add/use macro config_get_fn
Diffstat (limited to 'src/client/azure_openai.rs')
| -rw-r--r-- | src/client/azure_openai.rs | 82 |
1 files changed, 21 insertions, 61 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs index 45a5ecc..cabcebd 100644 --- a/src/client/azure_openai.rs +++ b/src/client/azure_openai.rs @@ -1,8 +1,5 @@ use super::openai::{openai_build_body, openai_send_message, openai_send_message_streaming}; -use super::{ - prompt_input_api_base, prompt_input_api_key, prompt_input_max_token, prompt_input_model_name, - Client, ClientConfig, ExtraConfig, ModelInfo, SendData, -}; +use super::{AzureOpenAIClient, Client, ExtraConfig, ModelInfo, PromptKind, PromptType, SendData}; use crate::config::SharedConfig; use crate::repl::ReplyStreamHandler; @@ -14,17 +11,10 @@ use serde::Deserialize; use std::env; -#[derive(Debug)] -pub struct AzureOpenAIClient { - global_config: SharedConfig, - config: AzureOpenAIConfig, - model_info: ModelInfo, -} - #[derive(Debug, Clone, Deserialize)] pub struct AzureOpenAIConfig { pub name: Option<String>, - pub api_base: String, + pub api_base: Option<String>, pub api_key: Option<String>, pub models: Vec<AzureOpenAIModel>, pub extra: Option<ExtraConfig>, @@ -33,17 +23,13 @@ pub struct AzureOpenAIConfig { #[derive(Debug, Clone, Deserialize)] pub struct AzureOpenAIModel { name: String, - max_tokens: usize, + max_tokens: Option<usize>, } #[async_trait] impl Client for AzureOpenAIClient { - fn config(&self) -> &SharedConfig { - &self.global_config - } - - fn extra_config(&self) -> &Option<ExtraConfig> { - &self.config.extra + fn config(&self) -> (&SharedConfig, &Option<ExtraConfig>) { + (&self.global_config, &self.config.extra) } async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> { @@ -63,27 +49,19 @@ impl Client for AzureOpenAIClient { } impl AzureOpenAIClient { - pub const NAME: &str = "azure-openai"; - - pub fn init(global_config: SharedConfig) -> Option<Box<dyn Client>> { - let model_info = global_config.read().model_info.clone(); - let config = { - if let ClientConfig::AzureOpenAI(c) = &global_config.read().clients[model_info.index] { - c.clone() - } else { - return None; - } - }; - Some(Box::new(Self { - global_config, - config, - model_info, - })) - } - - pub fn name(local_config: &AzureOpenAIConfig) -> &str { - local_config.name.as_deref().unwrap_or(Self::NAME) - } + config_get_fn!(api_base, get_api_base); + + 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, index: usize) -> Vec<ModelInfo> { let client = Self::name(local_config); @@ -95,26 +73,6 @@ impl AzureOpenAIClient { .collect() } - pub fn create_config() -> Result<String> { - let mut client_config = format!("clients:\n - type: {}\n", Self::NAME); - - let api_base = prompt_input_api_base()?; - client_config.push_str(&format!(" api_base: {api_base}\n")); - - let api_key = prompt_input_api_key()?; - client_config.push_str(&format!(" api_key: {api_key}\n")); - - let model_name = prompt_input_model_name()?; - - let max_tokens = prompt_input_max_token()?; - - client_config.push_str(&format!( - " models:\n - name: {model_name}\n max_tokens: {max_tokens}\n" - )); - - Ok(client_config) - } - fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> { let api_key = self.config.api_key.clone(); let api_key = api_key @@ -127,11 +85,13 @@ impl AzureOpenAIClient { }) .ok_or_else(|| anyhow!("Miss api_key"))?; + let api_base = self.get_api_base()?; + let body = openai_build_body(data, self.model_info.name.clone()); let url = format!( "{}/openai/deployments/{}/chat/completions?api-version=2023-05-15", - self.config.api_base, self.model_info.name + &api_base, self.model_info.name ); let builder = client.post(url).header("api-key", api_key).json(&body); |
