summaryrefslogtreecommitdiffstats
path: root/src/config/mod.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/config/mod.rs')
-rw-r--r--src/config/mod.rs64
1 files changed, 21 insertions, 43 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 64be3bd..93168fe 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -1,14 +1,17 @@
mod conversation;
+mod message;
+mod role;
use self::conversation::Conversation;
+use self::message::{Message, MESSAGE_EXTRA_TOKENS};
+use self::role::Role;
use crate::utils::{count_tokens, now};
use anyhow::{anyhow, bail, Context, Result};
use inquire::{Confirm, Text};
use parking_lot::Mutex;
-use serde::{Deserialize, Serialize};
-use serde_json::{json, Value};
+use serde::Deserialize;
use std::{
env,
fs::{create_dir_all, read_to_string, File, OpenOptions},
@@ -19,12 +22,10 @@ use std::{
};
const MAX_TOKENS: usize = 4096;
-const MESSAGE_EXTRA_TOKENS: usize = 6;
const CONFIG_FILE_NAME: &str = "config.yaml";
const ROLES_FILE_NAME: &str = "roles.yaml";
const HISTORY_FILE_NAME: &str = "history.txt";
const MESSAGE_FILE_NAME: &str = "messages.md";
-const TEMP_ROLE_NAME: &str = "P";
const SET_COMPLETIONS: [&str; 9] = [
".set api_key",
".set temperature",
@@ -126,7 +127,7 @@ impl Config {
format!("# CHAT:[{timestamp}]\n{input}\n--------\n{output}\n--------\n\n",)
}
Some(v) => {
- if v.name == TEMP_ROLE_NAME {
+ if v.is_temp() {
format!(
"# CHAT:[{timestamp}]\n{}\n{input}\n--------\n{output}\n--------\n\n",
v.prompt
@@ -166,7 +167,7 @@ impl Config {
}
match self.find_role(name) {
Some(mut role) => {
- role.tokens = count_tokens(&role.prompt);
+ role.tokens = role.consume_tokens();
let output =
serde_yaml::to_string(&role).unwrap_or("Unable to echo role details".into());
self.role = Some(role);
@@ -178,12 +179,7 @@ impl Config {
pub fn create_temp_role(&mut self, prompt: &str) -> Result<()> {
self.ensure_no_conversation()?;
- self.role = Some(Role {
- name: TEMP_ROLE_NAME.into(),
- prompt: prompt.into(),
- temperature: self.temperature,
- tokens: count_tokens(prompt),
- });
+ self.role = Some(Role::new(prompt, self.temperature));
Ok(())
}
@@ -198,33 +194,32 @@ impl Config {
if let Some(conversation) = self.conversation.as_ref() {
conversation.echo_messages(content)
} else if let Some(role) = self.role.as_ref() {
- format!("{}\n{content}", role.prompt)
+ role.echo_messages(content)
} else {
content.to_string()
}
}
- pub fn build_messages(&self, content: &str) -> Result<Value> {
- let tokens = count_tokens(content) + MESSAGE_EXTRA_TOKENS;
+ pub fn build_messages(&self, content: &str) -> Result<Vec<Message>> {
+ let content_tokens = count_tokens(content);
let check_tokens = |tokens| {
if tokens >= MAX_TOKENS {
bail!("Exceed max tokens limit")
}
Ok(())
};
- check_tokens(tokens)?;
- let user_message = json!({ "role": "user", "content": content });
- let value = if let Some(conversation) = self.conversation.as_ref() {
- check_tokens(tokens + conversation.tokens)?;
+ let messages = if let Some(conversation) = self.conversation.as_ref() {
+ check_tokens(content_tokens + conversation.tokens)?;
conversation.build_emssages(content)
} else if let Some(role) = self.role.as_ref() {
- check_tokens(tokens + role.tokens + MESSAGE_EXTRA_TOKENS)?;
- let system_message = json!({ "role": "system", "content": role.prompt });
- json!([system_message, user_message])
+ check_tokens(content_tokens + role.tokens)?;
+ role.build_emssages(content)
} else {
- json!([user_message])
+ let message = Message::new(content);
+ check_tokens(content_tokens + MESSAGE_EXTRA_TOKENS)?;
+ vec![message]
};
- Ok(value)
+ Ok(messages)
}
pub fn info(&self) -> Result<String> {
@@ -327,11 +322,7 @@ impl Config {
return Ok(());
}
}
- let mut conversation = Conversation::new();
- if let Some(role) = self.role.as_ref() {
- conversation.add_prompt(&role.prompt);
- }
- self.conversation = Some(conversation);
+ self.conversation = Some(Conversation::new(self.role.clone()));
Ok(())
}
@@ -341,7 +332,7 @@ impl Config {
pub fn save_conversation(&mut self, input: &str, output: &str) -> Result<()> {
if let Some(conversation) = self.conversation.as_mut() {
- conversation.add_chat(input, output)?;
+ conversation.add_message(input, output)?;
}
Ok(())
}
@@ -376,19 +367,6 @@ impl Config {
}
}
-#[derive(Debug, Clone, Deserialize, Serialize)]
-pub struct Role {
- /// Role name
- pub name: String,
- /// Prompt text send to ai for setting up a role
- pub prompt: String,
- /// What sampling temperature to use, between 0 and 2
- pub temperature: Option<f64>,
- /// Number of tokens
- #[serde(skip_deserializing)]
- pub tokens: usize,
-}
-
fn create_config_file(config_path: &Path) -> Result<()> {
let confirm_map_err = |_| anyhow!("Not finish questionnaire, try again later.");
let text_map_err = |_| anyhow!("An error happened when asking for your key, try again later.");