summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--src/config/mod.rs14
-rw-r--r--src/config/role.rs69
2 files changed, 77 insertions, 6 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 5d1172b..f7f2b2e 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -733,9 +733,10 @@ impl Config {
}
pub fn retrieve_role(&self, name: &str) -> Result<Role> {
- let mut role = if Self::list_roles(false).contains(&name.to_string()) {
- let path = Self::role_file(name)?;
- let content = read_to_string(path)?;
+ let names = Self::list_roles(false);
+ let mut role = if let Some(role_name) = Role::match_name(&names, name) {
+ let path = Self::role_file(&role_name)?;
+ let content = read_to_string(&path)?;
Role::new(name, &content)
} else {
BUILTIN_ROLES
@@ -777,7 +778,9 @@ impl Config {
}
pub fn upsert_role(&mut self, name: &str) -> Result<()> {
- let role_path = Self::role_file(name)?;
+ let names = Self::list_roles(false);
+ let role_name = Role::match_name(&names, name).unwrap_or_else(|| name.to_string());
+ let role_path = Self::role_file(&role_name)?;
ensure_parent_exists(&role_path)?;
let editor = self.editor()?;
edit_file(&editor, &role_path)?;
@@ -859,7 +862,8 @@ impl Config {
}
pub fn has_role(name: &str) -> bool {
- Self::list_roles(true).iter().any(|v| v == name)
+ let names = Self::list_roles(true);
+ Role::match_name(&names, name).is_some()
}
pub fn use_session(&mut self, session_name: Option<&str>) -> Result<()> {
diff --git a/src/config/role.rs b/src/config/role.rs
index c61aa53..04328f9 100644
--- a/src/config/role.rs
+++ b/src/config/role.rs
@@ -99,7 +99,7 @@ impl Role {
}
let mut role = Self {
name: name.to_string(),
- prompt: prompt.to_string(),
+ prompt: complete_prompt_args(prompt, name),
..Default::default()
};
if !metadata.is_empty() {
@@ -120,6 +120,23 @@ impl Role {
role
}
+ pub fn match_name(names: &[String], name: &str) -> Option<String> {
+ if names.contains(&name.to_string()) {
+ Some(name.to_string())
+ } else {
+ let parts: Vec<&str> = name.split(':').collect();
+ let parts_len = parts.len();
+ if parts_len < 2 {
+ return None;
+ }
+ let prefix = format!("{}:", parts[0]);
+ names
+ .iter()
+ .find(|v| v.starts_with(&prefix) && v.split(':').count() == parts_len)
+ .cloned()
+ }
+ }
+
pub fn export(&self) -> String {
let mut metadata = vec![];
if let Some(model) = self.model_id() {
@@ -310,6 +327,14 @@ impl RoleLike for Role {
}
}
+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
+}
+
fn parse_structure_prompt(prompt: &str) -> (&str, Vec<(&str, &str)>) {
let mut text = prompt;
let mut search_input = true;
@@ -380,6 +405,48 @@ 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"
+ );
+ }
+
+ #[test]
+ fn test_match_name() {
+ let names = vec![
+ "convert:yaml:json".into(),
+ "convert:yaml".into(),
+ "convert".into(),
+ ];
+ assert_eq!(
+ Role::match_name(&names, "convert"),
+ Some("convert".to_string())
+ );
+ assert_eq!(
+ Role::match_name(&names, "convert:yaml"),
+ Some("convert:yaml".to_string())
+ );
+ assert_eq!(
+ Role::match_name(&names, "convert:json"),
+ Some("convert:yaml".to_string())
+ );
+ assert_eq!(
+ Role::match_name(&names, "convert:yaml:json"),
+ Some("convert:yaml:json".to_string())
+ );
+ assert_eq!(
+ Role::match_name(&names, "convert:json:yaml"),
+ Some("convert:yaml:json".to_string())
+ );
+ assert_eq!(Role::match_name(&names, "convert:yaml:json:simple"), None,);
+ }
+
+ #[test]
fn test_parse_structure_prompt1() {
let prompt = r#"
System message