summaryrefslogtreecommitdiffstats
path: root/src/config/role.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/config/role.rs')
-rw-r--r--src/config/role.rs47
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"
+ );
+ }
+}