diff options
| author | sigoden <sigoden@gmail.com> | 2023-03-09 19:10:02 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-03-09 19:10:02 +0800 |
| commit | 9767c07eeeb20cb2f8386faf66cb88452987988b (patch) | |
| tree | fd6eec9a9ccad49798ab9de7329db507987d8a4d /src/config/mod.rs | |
| parent | c45d71cdeabbced7a55381e3b431c3eae3799ac0 (diff) | |
| download | aichat-9767c07eeeb20cb2f8386faf66cb88452987988b.tar.gz | |
feat: support two types of role prompts (#52)
1. embeded prompt
use __INPUT__ placeholder
will generate one user message when send to gpt
```
- name: shell
prompt: >
I want you to act as a linux shell expert.
Q: How to unzip a file
A: unzip file.zip
Q: __INPUT__
A:
```
2. system prompt
no __INPUT__ placeholder
will generate on system message and one user message when send to gpt
```
- name: shell
prompt: |
I want you to act as a linux shell expert.
I want you to answer only with bash code.
Do not write explanations.
```
Diffstat (limited to 'src/config/mod.rs')
| -rw-r--r-- | src/config/mod.rs | 64 |
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."); |
