diff options
| author | sigoden <sigoden@gmail.com> | 2024-09-02 17:56:04 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-09-02 17:56:04 +0800 |
| commit | b34b542d31e8d42720676015497c619021aed00d (patch) | |
| tree | cd5b2c40bd8bae3735df8520e687392d71c2fb96 /src/config | |
| parent | dc78636129427111e8949dcb0160aee15c0f3e91 (diff) | |
| download | aichat-b34b542d31e8d42720676015497c619021aed00d.tar.gz | |
refactor: keep the role arguments (#823)
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/mod.rs | 14 | ||||
| -rw-r--r-- | src/config/role.rs | 69 |
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 |
