diff options
Diffstat (limited to 'src/config/role.rs')
| -rw-r--r-- | src/config/role.rs | 47 |
1 files changed, 42 insertions, 5 deletions
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<f64>, @@ -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" + ); + } +} |
