From 9767c07eeeb20cb2f8386faf66cb88452987988b Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 9 Mar 2023 19:10:02 +0800 Subject: 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. ``` --- src/config/role.rs | 89 ++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 89 insertions(+) create mode 100644 src/config/role.rs (limited to 'src/config/role.rs') diff --git a/src/config/role.rs b/src/config/role.rs new file mode 100644 index 0000000..d0ed623 --- /dev/null +++ b/src/config/role.rs @@ -0,0 +1,89 @@ +use super::message::{Message, MessageRole, MESSAGE_EXTRA_TOKENS}; + +use crate::utils::count_tokens; + +use serde::{Deserialize, Serialize}; + +const TEMP_NAME: &str = "P"; +const INPUT_PLACEHOLDER: &str = "__INPUT__"; +const INPUT_PLACEHOLDER_TOKENS: usize = 3; + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct Role { + /// Role name + pub name: String, + /// Prompt text send to ai for setting up a role. + /// + /// If prmopt contains __INPUT___, it's embeded prompt + /// If prmopt don't contain __INPUT___, it's system prompt + pub prompt: String, + /// What sampling temperature to use, between 0 and 2 + pub temperature: Option, + /// Number of tokens + /// + /// System prompt consume extra 6 tokens + #[serde(skip_deserializing)] + pub tokens: usize, +} + +impl Role { + pub fn new(prompt: &str, temperature: Option) -> Self { + let mut value = Self { + name: TEMP_NAME.into(), + prompt: prompt.into(), + temperature, + tokens: 0, + }; + value.tokens = value.consume_tokens(); + value + } + + pub fn is_temp(&self) -> bool { + self.name == TEMP_NAME + } + + pub fn consume_tokens(&self) -> usize { + if self.embeded() { + count_tokens(&self.prompt) + MESSAGE_EXTRA_TOKENS - INPUT_PLACEHOLDER_TOKENS + } else { + count_tokens(&self.prompt) + 2 * MESSAGE_EXTRA_TOKENS + } + } + + pub fn embeded(&self) -> bool { + self.prompt.contains(INPUT_PLACEHOLDER) + } + + pub fn echo_messages(&self, content: &str) -> String { + if self.embeded() { + merge_prompt_content(&self.prompt, content) + } else { + format!("{}{content}", self.prompt) + } + } + + pub fn build_emssages(&self, content: &str) -> Vec { + if self.embeded() { + let content = merge_prompt_content(&self.prompt, content); + vec![Message { + role: MessageRole::User, + content, + }] + } else { + vec![ + Message { + role: MessageRole::System, + content: self.prompt.clone(), + }, + Message { + role: MessageRole::User, + content: content.to_string(), + }, + ] + } + } +} + +pub fn merge_prompt_content(prompt: &str, content: &str) -> String { + prompt.replace(INPUT_PLACEHOLDER, content) +} -- cgit v1.2.3