summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-09-04 06:19:32 +0800
committerGitHub <noreply@github.com>2024-09-04 06:19:32 +0800
commit9654445c32e72865c6e3af32dc1e5a3dc411a593 (patch)
tree410abbfcc248a050f3324a8863c4250f0228a1b2
parentd57f11445da31eb1183c34d78a93f8b601c2333e (diff)
downloadaichat-9654445c32e72865c6e3af32dc1e5a3dc411a593.tar.gz
fix: `:` can be used as seperator for role arguments (#830)
-rw-r--r--src/config/mod.rs3
-rw-r--r--src/config/role.rs34
2 files changed, 20 insertions, 17 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index f7f2b2e..3ef68c6 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -796,6 +796,9 @@ impl Config {
},
None => bail!("No role"),
};
+ if role_name.contains('#') {
+ bail!("Unable to save role with arguments")
+ }
if role_name == TEMP_ROLE_NAME {
role_name = Text::new("Role name:")
.with_validator(|input: &str| {
diff --git a/src/config/role.rs b/src/config/role.rs
index 04328f9..cd00157 100644
--- a/src/config/role.rs
+++ b/src/config/role.rs
@@ -124,15 +124,15 @@ impl Role {
if names.contains(&name.to_string()) {
Some(name.to_string())
} else {
- let parts: Vec<&str> = name.split(':').collect();
+ let parts: Vec<&str> = name.split('#').collect();
let parts_len = parts.len();
if parts_len < 2 {
return None;
}
- let prefix = format!("{}:", parts[0]);
+ let prefix = format!("{}#", parts[0]);
names
.iter()
- .find(|v| v.starts_with(&prefix) && v.split(':').count() == parts_len)
+ .find(|v| v.starts_with(&prefix) && v.split('#').count() == parts_len)
.cloned()
}
}
@@ -329,7 +329,7 @@ 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() {
+ for (i, arg) in name.split('#').skip(1).enumerate() {
prompt = prompt.replace(&format!("__ARG{}__", i + 1), arg);
}
prompt
@@ -407,11 +407,11 @@ mod tests {
#[test]
fn test_merge_prompt_name() {
assert_eq!(
- complete_prompt_args("convert __ARG1__", "convert:foo"),
+ complete_prompt_args("convert __ARG1__", "convert#foo"),
"convert foo"
);
assert_eq!(
- complete_prompt_args("convert __ARG1__ to __ARG2__", "convert:foo:bar"),
+ complete_prompt_args("convert __ARG1__ to __ARG2__", "convert#foo#bar"),
"convert foo to bar"
);
}
@@ -419,8 +419,8 @@ mod tests {
#[test]
fn test_match_name() {
let names = vec![
- "convert:yaml:json".into(),
- "convert:yaml".into(),
+ "convert#yaml#json".into(),
+ "convert#yaml".into(),
"convert".into(),
];
assert_eq!(
@@ -428,22 +428,22 @@ mod tests {
Some("convert".to_string())
);
assert_eq!(
- Role::match_name(&names, "convert:yaml"),
- Some("convert:yaml".to_string())
+ Role::match_name(&names, "convert#yaml"),
+ Some("convert#yaml".to_string())
);
assert_eq!(
- Role::match_name(&names, "convert:json"),
- Some("convert:yaml".to_string())
+ 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())
+ 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())
+ Role::match_name(&names, "convert#json#yaml"),
+ Some("convert#yaml#json".to_string())
);
- assert_eq!(Role::match_name(&names, "convert:yaml:json:simple"), None,);
+ assert_eq!(Role::match_name(&names, "convert#yaml#json#simple"), None,);
}
#[test]