diff options
| author | sigoden <sigoden@gmail.com> | 2024-09-10 18:35:34 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-09-10 18:35:34 +0800 |
| commit | e181ae9b0d4814a7c005c94837426126dde63c6c (patch) | |
| tree | ed727036d586f0673198a4b30217bbe2e34a094b /src | |
| parent | 84e9515509c559ed01e4b0a67539f10cd2c065e6 (diff) | |
| download | aichat-e181ae9b0d4814a7c005c94837426126dde63c6c.tar.gz | |
refactor: extract built-in roles to embedded files (#853)
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/agent.rs | 21 | ||||
| -rw-r--r-- | src/config/mod.rs | 19 | ||||
| -rw-r--r-- | src/config/role.rs | 106 | ||||
| -rw-r--r-- | src/config/session.rs | 9 | ||||
| -rw-r--r-- | src/repl/mod.rs | 2 | ||||
| -rw-r--r-- | src/utils/command.rs | 15 | ||||
| -rw-r--r-- | src/utils/mod.rs | 2 | ||||
| -rw-r--r-- | src/utils/variables.rs | 33 |
8 files changed, 79 insertions, 128 deletions
diff --git a/src/config/agent.rs b/src/config/agent.rs index bf4ef54..36af15e 100644 --- a/src/config/agent.rs +++ b/src/config/agent.rs @@ -292,9 +292,7 @@ impl AgentDefinition { for variable in &self.variables { output = output.replace(&format!("{{{{{}}}}}", variable.name), &variable.value) } - for (key, value) in builtin_variables() { - output = output.replace(&format!("{{{{{}}}}}", key), &value); - } + interpolate_variables(&mut output); output } @@ -412,20 +410,3 @@ fn save_variables(variables_path: &Path, variables: &[AgentVariable]) -> Result< .with_context(|| format!("Failed to save variables to '{}'", variables_path.display()))?; Ok(()) } - -fn builtin_variables() -> Vec<(&'static str, String)> { - vec![ - ("__os__", env::consts::OS.to_string()), - ("__os_family__", env::consts::FAMILY.to_string()), - ("__arch__", env::consts::ARCH.to_string()), - ("__shell__", SHELL.name.clone()), - ("__locale__", sys_locale::get_locale().unwrap_or_default()), - ("__now__", now()), - ( - "__cwd__", - env::current_dir() - .map(|v| v.display().to_string()) - .unwrap_or_default(), - ), - ] -} diff --git a/src/config/mod.rs b/src/config/mod.rs index 5cd266f..dae2018 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -5,7 +5,7 @@ mod session; pub use self::agent::{list_agents, Agent}; pub use self::input::Input; -pub use self::role::{Role, RoleLike, BUILTIN_ROLES, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; +pub use self::role::{Role, RoleLike, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; use self::session::Session; use crate::client::{ @@ -774,11 +774,7 @@ impl Config { let content = read_to_string(&path)?; Role::new(name, &content) } else { - BUILTIN_ROLES - .iter() - .find(|v| v.name() == name) - .cloned() - .ok_or_else(|| anyhow!("Unknown role `{name}`"))? + Role::builtin(name)? }; match role.model_id() { Some(model_id) => { @@ -839,8 +835,11 @@ impl Config { if role_name == TEMP_ROLE_NAME { role_name = Text::new("Role name:") .with_validator(|input: &str| { - if input.trim().is_empty() { - Ok(Validation::Invalid("This field is required".into())) + let input = input.trim(); + if input.is_empty() { + Ok(Validation::Invalid("This name is required".into())) + } else if input == TEMP_ROLE_NAME { + Ok(Validation::Invalid("This name is reserved".into())) } else { Ok(Validation::Valid) } @@ -862,7 +861,7 @@ impl Config { } pub fn all_roles() -> Vec<Role> { - let mut roles: HashMap<String, Role> = BUILTIN_ROLES + let mut roles: HashMap<String, Role> = Role::list_builtin_roles() .iter() .map(|v| (v.name().to_string(), v.clone())) .collect(); @@ -894,7 +893,7 @@ impl Config { } } if with_builtin { - names.extend(BUILTIN_ROLES.iter().map(|v| v.name().to_string())); + names.extend(Role::list_builtin_role_names()); } let mut names: Vec<_> = names.into_iter().collect(); names.sort_unstable(); diff --git a/src/config/role.rs b/src/config/role.rs index 74d9953..68366d0 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -4,6 +4,7 @@ use crate::client::{Message, MessageContent, MessageRole, Model}; use anyhow::Result; use fancy_regex::Regex; +use rust_embed::Embed; use serde::{Deserialize, Serialize}; use serde_json::Value; @@ -13,68 +14,11 @@ pub const CODE_ROLE: &str = "%code%"; pub const INPUT_PLACEHOLDER: &str = "__INPUT__"; -lazy_static::lazy_static! { - pub static ref BUILTIN_ROLES: Vec<Role> = { - [ - (SHELL_ROLE, shell_prompt()), - ( - EXPLAIN_SHELL_ROLE, - r#"Provide a terse, single sentence description of the given shell command. -Describe each argument and option of the command. -Provide short responses in about 80 words. -APPLY MARKDOWN formatting when possible."# - .into(), - ), - ( - CODE_ROLE, - r#"Provide only code without comments or explanations. -### INPUT: -async sleep in js -### OUTPUT: -```javascript -async function timeout(ms) { - return new Promise(resolve => setTimeout(resolve, ms)); -} -``` -"# - .into(), - ), - ( - "%create-prompt%", - r#"As a professional Prompt Engineer, your role is to create effective and innovative prompts for interacting with AI models. - -Your core skills include: -1. **CO-STAR Framework Application**: Utilize the CO-STAR framework to build efficient prompts, ensuring effective communication with large language models. -2. **Contextual Awareness**: Construct prompts that adapt to complex conversation contexts, ensuring relevant and coherent responses. -3. **Chain-of-Thought Prompting**: Create prompts that elicit AI models to demonstrate their reasoning process, enhancing the transparency and accuracy of answers. -4. **Zero-shot Learning**: Design prompts that enable AI models to perform specific tasks without requiring examples, reducing dependence on training data. -5. **Few-shot Learning**: Guide AI models to quickly learn and execute new tasks through a few examples. - -Your output format should include: -- **Context**: Provide comprehensive background information for the task to ensure the AI understands the specific scenario and offers relevant feedback. -- **Objective**: Clearly define the task objective, guiding the AI to focus on achieving specific goals. -- **Style**: Specify writing styles according to requirements, such as imitating a particular person or industry expert. -- **Tone**: Set an appropriate emotional tone to ensure the AI's response aligns with the expected emotional context. -- **Audience**: Tailor AI responses for a specific audience, ensuring content appropriateness and ease of understanding. -- **Response**: Specify output formats for easy execution of downstream tasks, such as lists, JSON, or professional reports. -- **Workflow**: Instruct the AI on how to step-by-step complete tasks, clarifying inputs, outputs, and specific actions for each step. -- **Examples**: Show a case of input and output that fits the scenario. - -Your workflow should be: -1. **Analyze User Input**: Extract key information from user requests to determine design objectives. -2. **Conceive New Prompts**: Based on user needs, create prompts that meet requirements, with each part being professional and detailed. -3. **Generate Output**: Must only output the newly generated and optimized prompts, without explanation, and without wrapping it in markdown code block."#.into(), - ), - ("%functions%", r#"--- -use_tools: all ---- - "#.into()), - ] - .into_iter() - .map(|(name, content)| Role::new(name, &content)) - .collect() - }; +#[derive(Embed)] +#[folder = "assets/roles/"] +struct RolesAsset; +lazy_static::lazy_static! { static ref RE_METADATA: Regex = Regex::new(r"(?s)-{3,}\s*(.*?)\s*-{3,}\s*(.*)").unwrap(); } @@ -122,9 +66,11 @@ impl Role { prompt = prompt_value.as_str().trim(); } } + let mut prompt = complete_prompt_args(prompt, name); + interpolate_variables(&mut prompt); let mut role = Self { name: name.to_string(), - prompt: complete_prompt_args(prompt, name), + prompt, ..Default::default() }; if !metadata.is_empty() { @@ -145,6 +91,25 @@ impl Role { role } + pub fn builtin(name: &str) -> Result<Self> { + let content = RolesAsset::get(&format!("{name}.md")) + .ok_or_else(|| anyhow!("Unknown role `{name}`"))?; + let content = unsafe { std::str::from_utf8_unchecked(&content.data) }; + Ok(Role::new(name, content)) + } + + pub fn list_builtin_role_names() -> Vec<String> { + RolesAsset::iter() + .filter_map(|v| v.strip_suffix(".md").map(|v| v.to_string())) + .collect() + } + + pub fn list_builtin_roles() -> Vec<Self> { + RolesAsset::iter() + .filter_map(|v| Role::builtin(&v).ok()) + .collect() + } + pub fn match_name(names: &[String], name: &str) -> Option<String> { if names.contains(&name.to_string()) { Some(name.to_string()) @@ -412,23 +377,6 @@ fn parse_structure_prompt(prompt: &str) -> (&str, Vec<(&str, &str)>) { (prompt, vec![]) } -fn shell_prompt() -> String { - let os = OS.as_str(); - let shell = SHELL.name.as_str(); - let combinator = if shell == "powershell" { - "If multiple steps required try to combine them together using ';'.\nIf it already combined with '&&' try to replace it with ';'.".to_string() - } else { - "If multiple steps required try to combine them together using '&&'.".to_string() - }; - format!( - r#"Provide only {shell} commands for {os} without any description. -Ensure the output is a valid {shell} command. -{combinator} -If there is a lack of details, provide most logical solution. -Output plain text only, without any markdown formatting."# - ) -} - #[cfg(test)] mod tests { use super::*; diff --git a/src/config/session.rs b/src/config/session.rs index 476ad51..0bb55e8 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -298,11 +298,14 @@ impl Session { if !ans { return Ok(()); } - while session_name == TEMP_SESSION_NAME { + if session_name == TEMP_SESSION_NAME { session_name = Text::new("Session name:") .with_validator(|input: &str| { - if input.trim().is_empty() { - Ok(Validation::Invalid("This field is required".into())) + let input = input.trim(); + if input.is_empty() { + Ok(Validation::Invalid("This name is required".into())) + } else if input == TEMP_SESSION_NAME { + Ok(Validation::Invalid("This name is reserved".into())) } else { Ok(Validation::Valid) } diff --git a/src/repl/mod.rs b/src/repl/mod.rs index b6c71a9..8c68f78 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -326,7 +326,7 @@ impl Repl { self.config.write().save_session(name)?; } _ => { - println!(r#"Usage: .save session [name]"#) + println!(r#"Usage: .save <role|session> [name]"#) } } } diff --git a/src/utils/command.rs b/src/utils/command.rs index caf70c7..8afef14 100644 --- a/src/utils/command.rs +++ b/src/utils/command.rs @@ -5,24 +5,9 @@ use std::{collections::HashMap, env, ffi::OsStr, path::Path, process::Command}; use anyhow::{anyhow, bail, Context, Result}; lazy_static::lazy_static! { - pub static ref OS: String = detect_os(); pub static ref SHELL: Shell = detect_shell(); } -pub fn detect_os() -> String { - let os = env::consts::OS; - if os == "linux" { - if let Ok(contents) = std::fs::read_to_string("/etc/os-release") { - for line in contents.lines() { - if let Some(id) = line.strip_prefix("ID=") { - return format!("{os} ({id})"); - } - } - } - } - os.to_string() -} - pub struct Shell { pub name: String, pub cmd: String, diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 1090c4e..e937e6b 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -8,6 +8,7 @@ mod prompt_input; mod render_prompt; mod request; mod spinner; +mod variables; pub use self::abort_signal::*; pub use self::clipboard::set_text; @@ -19,6 +20,7 @@ pub use self::prompt_input::*; pub use self::render_prompt::render_prompt; pub use self::request::*; pub use self::spinner::*; +pub use self::variables::*; use anyhow::{Context, Result}; use fancy_regex::Regex; diff --git a/src/utils/variables.rs b/src/utils/variables.rs new file mode 100644 index 0000000..91c9912 --- /dev/null +++ b/src/utils/variables.rs @@ -0,0 +1,33 @@ +use super::*; +use fancy_regex::{Captures, Regex}; + +lazy_static::lazy_static! { + pub static ref RE_VARIABLE: Regex = Regex::new(r"\{\{(\w+)\}\}").unwrap(); +} +pub fn interpolate_variables(text: &mut String) { + *text = RE_VARIABLE + .replace_all(text, |caps: &Captures<'_>| { + let key = &caps[1]; + match key { + "__os__" => env::consts::OS.to_string(), + "__os_distro__" => { + let info = os_info::get(); + if env::consts::OS == "linux" { + format!("{info} (linux)") + } else { + info.to_string() + } + } + "__os_family__" => env::consts::FAMILY.to_string(), + "__arch__" => env::consts::ARCH.to_string(), + "__shell__" => SHELL.name.clone(), + "__locale__" => sys_locale::get_locale().unwrap_or_default(), + "__now__" => now(), + "__cwd__" => env::current_dir() + .map(|v| v.display().to_string()) + .unwrap_or_default(), + _ => format!("{{{{{}}}}}", key), + } + }) + .to_string(); +} |
