diff options
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/message.rs | 52 | ||||
| -rw-r--r-- | src/config/mod.rs | 10 | ||||
| -rw-r--r-- | src/config/model_info.rs | 80 | ||||
| -rw-r--r-- | src/config/role.rs | 2 | ||||
| -rw-r--r-- | src/config/session.rs | 2 |
5 files changed, 5 insertions, 141 deletions
diff --git a/src/config/message.rs b/src/config/message.rs deleted file mode 100644 index 55b2663..0000000 --- a/src/config/message.rs +++ /dev/null @@ -1,52 +0,0 @@ -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/config/mod.rs b/src/config/mod.rs index 731cc10..da8e9b7 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1,16 +1,12 @@ -mod message; -mod model_info; mod role; mod session; -pub use self::message::Message; -pub use self::model_info::{ModelInfo, TokensCountFactors}; use self::role::Role; use self::session::{Session, TEMP_SESSION_NAME}; use crate::client::{ - all_models, create_client_config, list_client_types, ClientConfig, ExtraConfig, OpenAIClient, - SendData, + all_models, create_client_config, list_client_types, ClientConfig, ExtraConfig, Message, + ModelInfo, OpenAIClient, SendData, }; use crate::render::{MarkdownRender, RenderOptions}; use crate::utils::{get_env_name, light_theme_from_colorfgbg, now, prompt_op_err}; @@ -118,7 +114,7 @@ impl Default for Config { } } -pub type SharedConfig = Arc<RwLock<Config>>; +pub type GlobalConfig = Arc<RwLock<Config>>; impl Config { pub fn init(is_interactive: bool) -> Result<Self> { diff --git a/src/config/model_info.rs b/src/config/model_info.rs deleted file mode 100644 index 7a52e63..0000000 --- a/src/config/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<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/config/role.rs b/src/config/role.rs index 819cc12..eaf805f 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -1,4 +1,4 @@ -use super::message::{Message, MessageRole}; +use crate::client::{Message, MessageRole}; use anyhow::{Context, Result}; use serde::{Deserialize, Serialize}; diff --git a/src/config/session.rs b/src/config/session.rs index 7ed5714..446b682 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -1,7 +1,7 @@ -use super::message::{Message, MessageRole}; use super::role::Role; use super::ModelInfo; +use crate::client::{Message, MessageRole}; use crate::render::MarkdownRender; use anyhow::{bail, Context, Result}; |
