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/role.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/role.rs')
| -rw-r--r-- | src/config/role.rs | 89 |
1 files changed, 89 insertions, 0 deletions
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<f64>, + /// Number of tokens + /// + /// System prompt consume extra 6 tokens + #[serde(skip_deserializing)] + pub tokens: usize, +} + +impl Role { + pub fn new(prompt: &str, temperature: Option<f64>) -> 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<Message> { + 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) +} |
