From 9c04455e363d076cd7582fe9d476aa0fd6d86ea8 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 13 Mar 2023 10:09:01 +0800 Subject: feat: support role args (#69) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat: support role args We can use role args to pass some additional arguments to the prompt. ``` - name: convert:json:yaml prompt: convert __ARG1__ below to __ARG2__ ``` `:json:yaml` is `role args`. It has two args: - arg1 `json`, it will replace __ARG1__ in prompt - arg2 `yaml`, it will replace __ARG2__ in prompt ``` 〉.role convert:json:yaml name: convert:json:yaml prompt: convert json below to yaml temperature: null 〉.role convert:yaml:json name: convert:yaml:json prompt: convert yaml below to json temperature: null ``` different role args, will generate different prompts. * small updates --- src/config/role.rs | 47 ++++++++++++++++++++++++++++++++++++++++++----- 1 file changed, 42 insertions(+), 5 deletions(-) (limited to 'src/config/role.rs') diff --git a/src/config/role.rs b/src/config/role.rs index 16f7bc1..2def03d 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -9,10 +9,7 @@ const INPUT_PLACEHOLDER: &str = "__INPUT__"; 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 + /// Prompt text pub prompt: String, /// What sampling temperature to use, between 0 and 2 pub temperature: Option, @@ -35,6 +32,21 @@ impl Role { self.prompt.contains(INPUT_PLACEHOLDER) } + pub fn complete_prompt_args(&mut self, name: &str) { + self.name = name.to_string(); + self.prompt = complete_prompt_args(&self.prompt, &self.name); + } + + pub fn match_name(&self, name: &str) -> bool { + if self.name.contains(':') { + let role_name_parts: Vec<&str> = self.name.split(':').collect(); + let name_parts: Vec<&str> = name.split(':').collect(); + role_name_parts[0] == name_parts[0] && role_name_parts.len() == name_parts.len() + } else { + self.name == name + } + } + pub fn echo_messages(&self, content: &str) -> String { if self.embeded() { merge_prompt_content(&self.prompt, content) @@ -65,6 +77,31 @@ impl Role { } } -pub fn merge_prompt_content(prompt: &str, content: &str) -> String { +fn merge_prompt_content(prompt: &str, content: &str) -> String { prompt.replace(INPUT_PLACEHOLDER, content) } + +fn complete_prompt_args(prompt: &str, name: &str) -> String { + let mut prompt = prompt.to_string(); + for (i, arg) in name.split(':').skip(1).enumerate() { + prompt = prompt.replace(&format!("__ARG{}__", i + 1), arg); + } + prompt +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_merge_prompt_name() { + assert_eq!( + complete_prompt_args("convert __ARG1__", "convert:foo"), + "convert foo" + ); + assert_eq!( + complete_prompt_args("convert __ARG1__ to __ARG2__", "convert:foo:bar"), + "convert foo to bar" + ); + } +} -- cgit v1.2.3