diff options
| author | sigoden <sigoden@gmail.com> | 2023-03-16 17:02:09 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-03-16 17:02:09 +0800 |
| commit | 1ef97b2f32caecdb2eef9fe3755390f36d96eae5 (patch) | |
| tree | 269d9504e5959d2ca161862e0935103adc84b393 /src/config/mod.rs | |
| parent | 4a74f5cd72160585721dbce1e92c110125d046dd (diff) | |
| download | aichat-1ef97b2f32caecdb2eef9fe3755390f36d96eae5.tar.gz | |
feat: support multiple models (#71)
Diffstat (limited to 'src/config/mod.rs')
| -rw-r--r-- | src/config/mod.rs | 55 |
1 files changed, 46 insertions, 9 deletions
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<String>, + /// Openai model + #[serde(rename(serialize = "model", deserialize = "model"))] + pub model_name: Option<String>, /// What sampling temperature to use, between 0 and 2 pub temperature: Option<f64>, /// Whether to persistently save chat messages @@ -65,12 +74,15 @@ pub struct Config { /// Current conversation #[serde(skip)] pub conversation: Option<Conversation>, + #[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<Vec<Message>> { 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<String> { 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())); |
