summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-11-02 10:45:11 +0800
committerGitHub <noreply@github.com>2023-11-02 10:45:11 +0800
commit7c6841782d36faacc2aa3616dc3ed1b9403fa26e (patch)
tree12d51728e056dbd4e57fc5977b7c51b2b60c06ae /src/config
parent444f4ebe9de8aa68afdf8d39b734a8543b9b5c4b (diff)
downloadaichat-7c6841782d36faacc2aa3616dc3ed1b9403fa26e.tar.gz
refactor: improve code quanity (#197)
- move model_info.rs/message.rs to clients/ - rename SharedConfig to GlobalConfig
Diffstat (limited to 'src/config')
-rw-r--r--src/config/message.rs52
-rw-r--r--src/config/mod.rs10
-rw-r--r--src/config/model_info.rs80
-rw-r--r--src/config/role.rs2
-rw-r--r--src/config/session.rs2
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};