summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/cli.rs6
-rw-r--r--src/client.rs4
-rw-r--r--src/config/conversation.rs6
-rw-r--r--src/config/message.rs6
-rw-r--r--src/config/mod.rs55
-rw-r--r--src/main.rs9
-rw-r--r--src/repl/handler.rs5
-rw-r--r--src/repl/mod.rs7
-rw-r--r--src/repl/prompt.rs6
9 files changed, 80 insertions, 24 deletions
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<String>,
/// Add a GPT prompt
#[clap(short, long)]
pub prompt: Option<String>,
@@ -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<String>,
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<RequestBuilder> {
+ 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<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()));
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 <name>\n\n"),
+ },
".role" => match args {
Some(name) => handler.handle(ReplCmd::SetRole(name.to_string()))?,
None => print_now!("Usage: .role <name>\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<str> {
- 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()
}
}