From f9c40e52dabda7b037805c0635b84ccb6d75f5a8 Mon Sep 17 00:00:00 2001 From: sigoden Date: Fri, 3 Nov 2023 06:52:57 +0800 Subject: refactor: improve code quanity (#203) - update field name of ModelInfo - rename ModelInfo to Model --- src/client/azure_openai.rs | 12 +++---- src/client/common.rs | 18 +++++------ src/client/localai.rs | 10 +++--- src/client/mod.rs | 4 +-- src/client/model.rs | 80 ++++++++++++++++++++++++++++++++++++++++++++++ src/client/model_info.rs | 80 ---------------------------------------------- src/client/openai.rs | 10 +++--- 7 files changed, 107 insertions(+), 107 deletions(-) create mode 100644 src/client/model.rs delete mode 100644 src/client/model_info.rs (limited to 'src/client') diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs index f8a9dae..d1dc43b 100644 --- a/src/client/azure_openai.rs +++ b/src/client/azure_openai.rs @@ -1,5 +1,5 @@ use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS}; -use super::{AzureOpenAIClient, ExtraConfig, PromptType, SendData, ModelInfo}; +use super::{AzureOpenAIClient, ExtraConfig, PromptType, SendData, Model}; use crate::utils::PromptKind; @@ -42,14 +42,14 @@ impl AzureOpenAIClient { ), ]; - pub fn list_models(local_config: &AzureOpenAIConfig, index: usize) -> Vec { - let client = Self::name(local_config); + pub fn list_models(local_config: &AzureOpenAIConfig, client_index: usize) -> Vec { + let client_name = Self::name(local_config); local_config .models .iter() .map(|v| { - ModelInfo::new(index, client, &v.name) + Model::new(client_index, client_name, &v.name) .set_max_tokens(v.max_tokens) .set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS) }) @@ -70,11 +70,11 @@ impl AzureOpenAIClient { let api_base = self.get_api_base()?; - let body = openai_build_body(data, self.model_info.name.clone()); + 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_info.name + &api_base, self.model.llm_name ); let builder = client.post(url).header("api-key", api_key).json(&body); diff --git a/src/client/common.rs b/src/client/common.rs index 0dc637b..464e46a 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -46,16 +46,16 @@ macro_rules! register_client { pub struct $client { global_config: $crate::config::GlobalConfig, config: $config, - model_info: $crate::client::ModelInfo, + model: $crate::client::Model, } impl $client { pub const NAME: &str = $name; pub fn init(global_config: $crate::config::GlobalConfig) -> Option> { - let model_info = global_config.read().model_info.clone(); + let model = global_config.read().model.clone(); let config = { - if let ClientConfig::$config_key(c) = &global_config.read().clients[model_info.index] { + if let ClientConfig::$config_key(c) = &global_config.read().clients[model.client_index] { c.clone() } else { return None; @@ -64,7 +64,7 @@ macro_rules! register_client { Some(Box::new(Self { global_config, config, - model_info, + model, })) } @@ -79,11 +79,11 @@ macro_rules! register_client { None $(.or_else(|| $client::init(config.clone())))+ .ok_or_else(|| { - let model_info = config.read().model_info.clone(); + let model = config.read().model.clone(); anyhow::anyhow!( - "Unknown client {} at config.clients[{}]", - &model_info.client, - &model_info.index + "Unknown client '{}' at config.clients[{}]", + &model.client_name, + &model.client_index ) }) } @@ -101,7 +101,7 @@ macro_rules! register_client { anyhow::bail!("Unknown client {}", client) } - pub fn list_models(config: &$crate::config::Config) -> Vec<$crate::client::ModelInfo> { + pub fn list_models(config: &$crate::config::Config) -> Vec<$crate::client::Model> { config .clients .iter() diff --git a/src/client/localai.rs b/src/client/localai.rs index 5cc12cc..eb4de65 100644 --- a/src/client/localai.rs +++ b/src/client/localai.rs @@ -1,5 +1,5 @@ use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS}; -use super::{ExtraConfig, LocalAIClient, PromptType, SendData, ModelInfo}; +use super::{ExtraConfig, LocalAIClient, PromptType, SendData, Model}; use crate::utils::PromptKind; @@ -41,14 +41,14 @@ impl LocalAIClient { ), ]; - pub fn list_models(local_config: &LocalAIConfig, index: usize) -> Vec { - let client = Self::name(local_config); + pub fn list_models(local_config: &LocalAIConfig, client_index: usize) -> Vec { + let client_name = Self::name(local_config); local_config .models .iter() .map(|v| { - ModelInfo::new(index, client, &v.name) + Model::new(client_index, client_name, &v.name) .set_max_tokens(v.max_tokens) .set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS) }) @@ -58,7 +58,7 @@ impl LocalAIClient { fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { let api_key = self.get_api_key().ok(); - let body = openai_build_body(data, self.model_info.name.clone()); + let body = openai_build_body(data, self.model.llm_name.clone()); let chat_endpoint = self .config diff --git a/src/client/mod.rs b/src/client/mod.rs index 19a0875..7ac9aa0 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -1,11 +1,11 @@ #[macro_use] mod common; mod message; -mod model_info; +mod model; pub use common::*; pub use message::*; -pub use model_info::*; +pub use model::*; register_client!( (openai, "openai", OpenAI, OpenAIConfig, OpenAIClient), diff --git a/src/client/model.rs b/src/client/model.rs new file mode 100644 index 0000000..d00ad46 --- /dev/null +++ b/src/client/model.rs @@ -0,0 +1,80 @@ +use super::message::Message; + +use crate::utils::count_tokens; + +use anyhow::{bail, Result}; + +pub type TokensCountFactors = (usize, usize); // (per-messages, bias) + +#[derive(Debug, Clone)] +pub struct Model { + pub client_index: usize, + pub client_name: String, + pub llm_name: String, + pub max_tokens: Option, + pub tokens_count_factors: TokensCountFactors, +} + +impl Default for Model { + fn default() -> Self { + Model::new(0, "", "") + } +} + +impl Model { + pub fn new(client_index: usize, client_name: &str, name: &str) -> Self { + Self { + client_index, + client_name: client_name.into(), + llm_name: name.into(), + max_tokens: None, + tokens_count_factors: Default::default(), + } + } + + pub fn id(&self) -> String { + format!("{}:{}", self.client_name, self.llm_name) + } + + pub fn set_max_tokens(mut self, max_tokens: Option) -> Self { + match max_tokens { + None | Some(0) => self.max_tokens = None, + _ => self.max_tokens = max_tokens, + } + self + } + + pub fn set_tokens_count_factors(mut self, tokens_count_factors: TokensCountFactors) -> Self { + self.tokens_count_factors = tokens_count_factors; + self + } + + pub fn messages_tokens(&self, messages: &[Message]) -> usize { + messages.iter().map(|v| count_tokens(&v.content)).sum() + } + + pub fn total_tokens(&self, messages: &[Message]) -> usize { + if messages.is_empty() { + return 0; + } + let num_messages = messages.len(); + let message_tokens = self.messages_tokens(messages); + let (per_messages, _) = self.tokens_count_factors; + if messages[num_messages - 1].role.is_user() { + num_messages * per_messages + message_tokens + } else { + (num_messages - 1) * per_messages + message_tokens + } + } + + pub fn max_tokens_limit(&self, messages: &[Message]) -> Result<()> { + let (_, bias) = self.tokens_count_factors; + let total_tokens = self.total_tokens(messages) + bias; + if let Some(max_tokens) = self.max_tokens { + if total_tokens >= max_tokens { + bail!("Exceed max tokens limit") + } + } + Ok(()) + } +} diff --git a/src/client/model_info.rs b/src/client/model_info.rs deleted file mode 100644 index 9e74951..0000000 --- a/src/client/model_info.rs +++ /dev/null @@ -1,80 +0,0 @@ -use super::message::Message; - -use crate::utils::count_tokens; - -use anyhow::{bail, Result}; - -pub type TokensCountFactors = (usize, usize); // (per-messages, bias) - -#[derive(Debug, Clone)] -pub struct ModelInfo { - pub client: String, - pub name: String, - pub index: usize, - pub max_tokens: Option, - pub tokens_count_factors: TokensCountFactors, -} - -impl Default for ModelInfo { - fn default() -> Self { - ModelInfo::new(0, "", "") - } -} - -impl ModelInfo { - pub fn new(index: usize, client: &str, name: &str) -> Self { - Self { - index, - client: client.into(), - name: name.into(), - max_tokens: None, - tokens_count_factors: Default::default(), - } - } - - pub fn set_max_tokens(mut self, max_tokens: Option) -> Self { - match max_tokens { - None | Some(0) => self.max_tokens = None, - _ => self.max_tokens = max_tokens, - } - self - } - - pub fn set_tokens_count_factors(mut self, tokens_count_factors: TokensCountFactors) -> Self { - self.tokens_count_factors = tokens_count_factors; - self - } - - pub fn id(&self) -> String { - format!("{}:{}", self.client, self.name) - } - - pub fn messages_tokens(&self, messages: &[Message]) -> usize { - messages.iter().map(|v| count_tokens(&v.content)).sum() - } - - pub fn total_tokens(&self, messages: &[Message]) -> usize { - if messages.is_empty() { - return 0; - } - let num_messages = messages.len(); - let message_tokens = self.messages_tokens(messages); - let (per_messages, _) = self.tokens_count_factors; - if messages[num_messages - 1].role.is_user() { - num_messages * per_messages + message_tokens - } else { - (num_messages - 1) * per_messages + message_tokens - } - } - - pub fn max_tokens_limit(&self, messages: &[Message]) -> Result<()> { - let (_, bias) = self.tokens_count_factors; - let total_tokens = self.total_tokens(messages) + bias; - if let Some(max_tokens) = self.max_tokens { - if total_tokens >= max_tokens { - bail!("Exceed max tokens limit") - } - } - Ok(()) - } -} diff --git a/src/client/openai.rs b/src/client/openai.rs index 5589d2d..f1243f0 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -1,6 +1,6 @@ use super::{ ExtraConfig, OpenAIClient, PromptType, SendData, - ModelInfo, TokensCountFactors, + Model, TokensCountFactors, }; use crate::{ @@ -44,12 +44,12 @@ impl OpenAIClient { pub const PROMPTS: [PromptType<'static>; 1] = [("api_key", "API Key:", true, PromptKind::String)]; - pub fn list_models(local_config: &OpenAIConfig, index: usize) -> Vec { - let client = Self::name(local_config); + pub fn list_models(local_config: &OpenAIConfig, client_index: usize) -> Vec { + let client_name = Self::name(local_config); MODELS .into_iter() .map(|(name, max_tokens)| { - ModelInfo::new(index, client, name) + Model::new(client_index, client_name, name) .set_max_tokens(Some(max_tokens)) .set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS) }) @@ -59,7 +59,7 @@ impl OpenAIClient { fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { let api_key = self.get_api_key()?; - let body = openai_build_body(data, self.model_info.name.clone()); + let body = openai_build_body(data, self.model.llm_name.clone()); let env_prefix = Self::name(&self.config).to_uppercase(); let api_base = env::var(format!("{env_prefix}_API_BASE")) -- cgit v1.2.3