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/cli.rs | 6 +++++ src/client.rs | 4 ++-- src/config/conversation.rs | 6 +---- src/config/message.rs | 6 ++--- src/config/mod.rs | 55 ++++++++++++++++++++++++++++++++++++++-------- src/main.rs | 9 ++++++++ src/repl/handler.rs | 5 +++++ src/repl/mod.rs | 7 +++++- src/repl/prompt.rs | 6 ++--- 9 files changed, 80 insertions(+), 24 deletions(-) (limited to 'src') diff --git a/src/cli.rs b/src/cli.rs index cfc1476..9211878 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -3,6 +3,9 @@ use clap::Parser; #[derive(Parser, Debug)] #[command(author, version, about, long_about = None)] pub struct Cli { + /// Choose a model + #[clap(short, long)] + pub model: Option, /// Add a GPT prompt #[clap(short, long)] pub prompt: Option, @@ -15,6 +18,9 @@ pub struct Cli { /// List all roles #[clap(long)] pub list_roles: bool, + /// List all models + #[clap(long)] + pub list_models: bool, /// Select a role #[clap(short, long)] pub role: Option, diff --git a/src/client.rs b/src/client.rs index 916a1ef..b438440 100644 --- a/src/client.rs +++ b/src/client.rs @@ -12,7 +12,6 @@ use tokio::time::sleep; const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); const API_URL: &str = "https://api.openai.com/v1/chat/completions"; -const MODEL: &str = "gpt-3.5-turbo"; #[derive(Debug)] pub struct ChatGptClient { @@ -137,9 +136,10 @@ impl ChatGptClient { } fn request_builder(&self, content: &str, stream: bool) -> Result { + let (model, _) = self.config.read().get_model(); let messages = self.config.read().build_messages(content)?; let mut body = json!({ - "model": MODEL, + "model": model, "messages": messages, }); diff --git a/src/config/conversation.rs b/src/config/conversation.rs index c166195..8343400 100644 --- a/src/config/conversation.rs +++ b/src/config/conversation.rs @@ -1,4 +1,4 @@ -use super::message::{num_tokens_from_messages, Message, MessageRole, MAX_TOKENS}; +use super::message::{num_tokens_from_messages, Message, MessageRole}; use super::role::Role; use anyhow::{bail, Result}; @@ -87,8 +87,4 @@ impl Conversation { } messages } - - pub fn reamind_tokens(&self) -> usize { - MAX_TOKENS.saturating_sub(self.tokens) - } } diff --git a/src/config/message.rs b/src/config/message.rs index 3fd5935..5717220 100644 --- a/src/config/message.rs +++ b/src/config/message.rs @@ -3,8 +3,6 @@ use crate::utils::count_tokens; use anyhow::{bail, Result}; use serde::{Deserialize, Serialize}; -pub const MAX_TOKENS: usize = 4096; - #[derive(Debug, Clone, Deserialize, Serialize)] pub struct Message { pub role: MessageRole, @@ -28,9 +26,9 @@ pub enum MessageRole { User, } -pub fn within_max_tokens_limit(messages: &[Message]) -> Result<()> { +pub fn within_max_tokens_limit(messages: &[Message], max_tokens: usize) -> Result<()> { let tokens = num_tokens_from_messages(messages); - if tokens >= MAX_TOKENS { + if tokens >= max_tokens { bail!("Exceed max tokens limit") } Ok(()) 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())); diff --git a/src/main.rs b/src/main.rs index ecbfc60..589e865 100644 --- a/src/main.rs +++ b/src/main.rs @@ -35,6 +35,12 @@ fn main() -> Result<()> { .for_each(|v| println!("{}", v.name)); exit(0); } + if cli.list_models { + config::MODELS + .iter() + .for_each(|(name, _)| println!("{}", name)); + exit(0); + } let role = match &cli.role { Some(name) => Some( config @@ -44,6 +50,9 @@ fn main() -> Result<()> { ), None => None, }; + if let Some(model) = &cli.model { + config.write().set_model(model)?; + } config.write().role = role; if cli.no_highlight { config.write().highlight = false; diff --git a/src/repl/handler.rs b/src/repl/handler.rs index 159127c..014a851 100644 --- a/src/repl/handler.rs +++ b/src/repl/handler.rs @@ -12,6 +12,7 @@ use std::cell::RefCell; pub enum ReplCmd { Submit(String), + SetModel(String), SetRole(String), UpdateConfig(String), Prompt(String), @@ -65,6 +66,10 @@ impl ReplCmdHandler { self.config.write().save_conversation(&input, &buffer)?; *self.reply.borrow_mut() = buffer; } + ReplCmd::SetModel(name) => { + self.config.write().set_model(&name)?; + print_now!("\n"); + } ReplCmd::SetRole(name) => { let output = self.config.write().change_role(&name)?; print_now!("{}\n\n", output.trim_end()); diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 337afb3..663decb 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -19,9 +19,10 @@ use reedline::Signal; use std::borrow::Cow; use std::sync::Arc; -pub const REPL_COMMANDS: [(&str, &str); 11] = [ +pub const REPL_COMMANDS: [(&str, &str); 12] = [ (".info", "Print the information"), (".set", "Modify the configuration temporarily"), + (".model", "Choose a model"), (".prompt", "Add a GPT prompt"), (".role", "Select a role"), (".clear role", "Clear the currently selected role"), @@ -109,6 +110,10 @@ impl Repl { self.editor.print_history()?; print_now!("\n"); } + ".model" => match args { + Some(name) => handler.handle(ReplCmd::SetModel(name.to_string()))?, + None => print_now!("Usage: .model \n\n"), + }, ".role" => match args { Some(name) => handler.handle(ReplCmd::SetRole(name.to_string()))?, None => print_now!("Usage: .role \n\n"), diff --git a/src/repl/prompt.rs b/src/repl/prompt.rs index 1a59ba3..8c670f4 100644 --- a/src/repl/prompt.rs +++ b/src/repl/prompt.rs @@ -76,10 +76,10 @@ impl Prompt for ReplPrompt { } fn render_prompt_right(&self) -> Cow { - if let Some(conversation) = self.config.read().conversation.as_ref() { - conversation.reamind_tokens().to_string().into() - } else { + if self.config.read().conversation.is_none() { Cow::Borrowed("") + } else { + self.config.read().get_reamind_tokens().to_string().into() } } -- cgit v1.2.3