diff options
| author | sigoden <sigoden@gmail.com> | 2023-11-02 10:45:11 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-11-02 10:45:11 +0800 |
| commit | 7c6841782d36faacc2aa3616dc3ed1b9403fa26e (patch) | |
| tree | 12d51728e056dbd4e57fc5977b7c51b2b60c06ae /src/client | |
| parent | 444f4ebe9de8aa68afdf8d39b734a8543b9b5c4b (diff) | |
| download | aichat-7c6841782d36faacc2aa3616dc3ed1b9403fa26e.tar.gz | |
refactor: improve code quanity (#197)
- move model_info.rs/message.rs to clients/
- rename SharedConfig to GlobalConfig
Diffstat (limited to 'src/client')
| -rw-r--r-- | src/client/azure_openai.rs | 4 | ||||
| -rw-r--r-- | src/client/common.rs | 26 | ||||
| -rw-r--r-- | src/client/localai.rs | 4 | ||||
| -rw-r--r-- | src/client/message.rs | 52 | ||||
| -rw-r--r-- | src/client/mod.rs | 4 | ||||
| -rw-r--r-- | src/client/model_info.rs | 80 | ||||
| -rw-r--r-- | src/client/openai.rs | 2 |
7 files changed, 155 insertions, 17 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs index fe3ec0f..f8a9dae 100644 --- a/src/client/azure_openai.rs +++ b/src/client/azure_openai.rs @@ -1,7 +1,7 @@ use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS}; -use super::{AzureOpenAIClient, ExtraConfig, PromptType, SendData}; +use super::{AzureOpenAIClient, ExtraConfig, PromptType, SendData, ModelInfo}; -use crate::{config::ModelInfo, utils::PromptKind}; +use crate::utils::PromptKind; use anyhow::{anyhow, Result}; use async_trait::async_trait; diff --git a/src/client/common.rs b/src/client/common.rs index 0d0c0e2..a7844f3 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -1,8 +1,12 @@ +use super::{openai::OpenAIConfig, ClientConfig, Message}; + use crate::{ - config::{Message, SharedConfig}, + config::GlobalConfig, render::ReplyHandler, - repl::AbortSignal, - utils::{init_tokio_runtime, prompt_input_integer, prompt_input_string, tokenize, PromptKind}, + utils::{ + init_tokio_runtime, prompt_input_integer, prompt_input_string, tokenize, AbortSignal, + PromptKind, + }, }; use anyhow::{Context, Result}; @@ -13,8 +17,6 @@ use serde_json::{json, Value}; use std::{env, time::Duration}; use tokio::time::sleep; -use super::{openai::OpenAIConfig, ClientConfig}; - #[macro_export] macro_rules! register_client { ( @@ -42,15 +44,15 @@ macro_rules! register_client { $( #[derive(Debug)] pub struct $client { - global_config: $crate::config::SharedConfig, + global_config: $crate::config::GlobalConfig, config: $config, - model_info: $crate::config::ModelInfo, + model_info: $crate::client::ModelInfo, } impl $client { pub const NAME: &str = $name; - pub fn init(global_config: $crate::config::SharedConfig) -> Option<Box<dyn Client>> { + pub fn init(global_config: $crate::config::GlobalConfig) -> Option<Box<dyn Client>> { let model_info = global_config.read().model_info.clone(); let config = { if let ClientConfig::$config_key(c) = &global_config.read().clients[model_info.index] { @@ -73,7 +75,7 @@ macro_rules! register_client { )+ - pub fn init_client(config: $crate::config::SharedConfig) -> anyhow::Result<Box<dyn Client>> { + pub fn init_client(config: $crate::config::GlobalConfig) -> anyhow::Result<Box<dyn Client>> { None $(.or_else(|| $client::init(config.clone())))+ .ok_or_else(|| { @@ -99,7 +101,7 @@ macro_rules! register_client { anyhow::bail!("Unknown client {}", client) } - pub fn all_models(config: &$crate::config::Config) -> Vec<$crate::config::ModelInfo> { + pub fn all_models(config: &$crate::config::Config) -> Vec<$crate::client::ModelInfo> { config .clients .iter() @@ -122,7 +124,7 @@ macro_rules! openai_compatible_client { fn config( &self, ) -> ( - &$crate::config::SharedConfig, + &$crate::config::GlobalConfig, &Option<$crate::client::ExtraConfig>, ) { (&self.global_config, &self.config.extra) @@ -169,7 +171,7 @@ macro_rules! config_get_fn { #[async_trait] pub trait Client { - fn config(&self) -> (&SharedConfig, &Option<ExtraConfig>); + fn config(&self) -> (&GlobalConfig, &Option<ExtraConfig>); fn build_client(&self) -> Result<ReqwestClient> { let mut builder = ReqwestClient::builder(); diff --git a/src/client/localai.rs b/src/client/localai.rs index 796b574..5cc12cc 100644 --- a/src/client/localai.rs +++ b/src/client/localai.rs @@ -1,7 +1,7 @@ use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS}; -use super::{ExtraConfig, LocalAIClient, PromptType, SendData}; +use super::{ExtraConfig, LocalAIClient, PromptType, SendData, ModelInfo}; -use crate::{config::ModelInfo, utils::PromptKind}; +use crate::utils::PromptKind; use anyhow::Result; use async_trait::async_trait; diff --git a/src/client/message.rs b/src/client/message.rs new file mode 100644 index 0000000..55b2663 --- /dev/null +++ b/src/client/message.rs @@ -0,0 +1,52 @@ +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct Message { + pub role: MessageRole, + pub content: String, +} + +impl Message { + pub fn new(content: &str) -> Self { + Self { + role: MessageRole::User, + content: content.to_string(), + } + } +} + +#[derive(Debug, Clone, Copy, Deserialize, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum MessageRole { + System, + Assistant, + User, +} + +#[allow(dead_code)] +impl MessageRole { + pub fn is_system(&self) -> bool { + matches!(self, MessageRole::System) + } + + pub fn is_user(&self) -> bool { + matches!(self, MessageRole::User) + } + + pub fn is_assistant(&self) -> bool { + matches!(self, MessageRole::Assistant) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_serde() { + assert_eq!( + serde_json::to_string(&Message::new("Hello World")).unwrap(), + "{\"role\":\"user\",\"content\":\"Hello World\"}" + ); + } +} diff --git a/src/client/mod.rs b/src/client/mod.rs index e55055d..19a0875 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -1,7 +1,11 @@ #[macro_use] mod common; +mod message; +mod model_info; pub use common::*; +pub use message::*; +pub use model_info::*; register_client!( (openai, "openai", OpenAI, OpenAIConfig, OpenAIClient), diff --git a/src/client/model_info.rs b/src/client/model_info.rs new file mode 100644 index 0000000..7a52e63 --- /dev/null +++ b/src/client/model_info.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 ModelInfo { + pub client: String, + pub name: String, + pub index: usize, + pub max_tokens: Option<usize>, + 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<usize>) -> 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 full_name(&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 98de5a8..5589d2d 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -1,9 +1,9 @@ use super::{ ExtraConfig, OpenAIClient, PromptType, SendData, + ModelInfo, TokensCountFactors, }; use crate::{ - config::{ModelInfo, TokensCountFactors}, render::ReplyHandler, utils::PromptKind, }; |
