From 1ef97b2f32caecdb2eef9fe3755390f36d96eae5 Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 16 Mar 2023 17:02:09 +0800 Subject: feat: support multiple models (#71) --- src/config/mod.rs | 55 ++++++++++++++++++++++++++++++++++++++++++++++--------- 1 file changed, 46 insertions(+), 9 deletions(-) (limited to 'src/config/mod.rs') diff --git a/src/config/mod.rs b/src/config/mod.rs index b7833e4..f543ee9 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -21,6 +21,12 @@ use std::{ sync::Arc, }; +pub const MODELS: [(&str, usize); 3] = [ + ("gpt-4", 8192), + ("gpt-4-32k", 32768), + ("gpt-3.5-turbo", 4096), +]; + const CONFIG_FILE_NAME: &str = "config.yaml"; const ROLES_FILE_NAME: &str = "roles.yaml"; const HISTORY_FILE_NAME: &str = "history.txt"; @@ -42,6 +48,9 @@ const SET_COMPLETIONS: [&str; 9] = [ pub struct Config { /// Openai api key pub api_key: Option, + /// Openai model + #[serde(rename(serialize = "model", deserialize = "model"))] + pub model_name: Option, /// What sampling temperature to use, between 0 and 2 pub temperature: Option, /// Whether to persistently save chat messages @@ -65,12 +74,15 @@ pub struct Config { /// Current conversation #[serde(skip)] pub conversation: Option, + #[serde(skip)] + pub model: (String, usize), } impl Default for Config { fn default() -> Self { Self { api_key: None, + model_name: None, temperature: None, save: false, highlight: true, @@ -81,6 +93,7 @@ impl Default for Config { roles: vec![], role: None, conversation: None, + model: ("gpt-3.5-turbo".into(), 4096), } } } @@ -105,6 +118,9 @@ impl Config { if config.api_key.is_none() { bail!("api_key not set"); } + if let Some(name) = config.model_name.clone() { + config.set_model(&name)?; + } config.merge_env_vars(); config.maybe_proxy(); config.load_roles()?; @@ -251,6 +267,10 @@ impl Config { } } + pub fn get_model(&self) -> (String, usize) { + self.model.clone() + } + pub fn build_messages(&self, content: &str) -> Result> { let messages = if let Some(conversation) = self.conversation.as_ref() { conversation.build_emssages(content) @@ -260,11 +280,28 @@ impl Config { let message = Message::new(content); vec![message] }; - within_max_tokens_limit(&messages)?; + within_max_tokens_limit(&messages, self.model.1)?; Ok(messages) } + pub fn set_model(&mut self, name: &str) -> Result<()> { + if let Some(token) = MODELS.iter().find(|(v, _)| *v == name).map(|(_, v)| *v) { + self.model = (name.to_string(), token); + } else { + bail!("Invalid model") + } + Ok(()) + } + + pub fn get_reamind_tokens(&self) -> usize { + let mut tokens = self.model.1; + if let Some(conversation) = self.conversation.as_ref() { + tokens = tokens.saturating_sub(conversation.tokens); + } + tokens + } + pub fn info(&self) -> Result { let file_info = |path: &Path| { let state = if path.exists() { "" } else { " ⚠️" }; @@ -284,6 +321,7 @@ impl Config { ("roles_file", file_info(&Config::roles_file()?)), ("messages_file", file_info(&Config::messages_file()?)), ("api_key", self.get_api_key().to_string()), + ("model", self.model.0.to_string()), ("temperature", temperature), ("save", self.save.to_string()), ("highlight", self.highlight.to_string()), @@ -307,6 +345,7 @@ impl Config { .collect(); completion.extend(SET_COMPLETIONS.map(|v| v.to_string())); + completion.extend(MODELS.map(|(v, _)| format!(".model {}", v))); completion } @@ -359,14 +398,12 @@ impl Config { } pub fn start_conversation(&mut self) -> Result<()> { - if let Some(conversation) = self.conversation.as_ref() { - if conversation.reamind_tokens() > 0 { - let ans = Confirm::new("Already in a conversation, start a new one?") - .with_default(true) - .prompt()?; - if !ans { - return Ok(()); - } + if self.conversation.is_some() && self.get_reamind_tokens() > 0 { + let ans = Confirm::new("Already in a conversation, start a new one?") + .with_default(true) + .prompt()?; + if !ans { + return Ok(()); } } self.conversation = Some(Conversation::new(self.role.clone())); -- cgit v1.2.3