diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-11 09:23:59 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-11 09:23:59 +0800 |
| commit | 1e8fc5d269985048d8d3023a615b94a8908571cf (patch) | |
| tree | e1e43a67bf1e7b5d878d2a7cac5d25a07cda8bdc /src/config | |
| parent | 058299e500af86adbc065eced542e108ae325458 (diff) | |
| download | aichat-1e8fc5d269985048d8d3023a615b94a8908571cf.tar.gz | |
refactor: list roles includeing builtin roles (#499)
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/input.rs | 17 | ||||
| -rw-r--r-- | src/config/mod.rs | 17 | ||||
| -rw-r--r-- | src/config/role.rs | 100 |
3 files changed, 67 insertions, 67 deletions
diff --git a/src/config/input.rs b/src/config/input.rs index b7b896c..0527b49 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -106,7 +106,7 @@ impl Input { } pub fn session<'a>(&self, session: &'a Option<Session>) -> Option<&'a Session> { - if self.context.in_session { + if self.context.session { session.as_ref() } else { None @@ -114,7 +114,7 @@ impl Input { } pub fn session_mut<'a>(&self, session: &'a mut Option<Session>) -> Option<&'a mut Session> { - if self.context.in_session { + if self.context.session { session.as_mut() } else { None @@ -199,12 +199,19 @@ impl Input { #[derive(Debug, Clone, Default)] pub struct InputContext { role: Option<Role>, - in_session: bool, + session: bool, } impl InputContext { - pub fn new(role: Option<Role>, in_session: bool) -> Self { - Self { role, in_session } + pub fn new(role: Option<Role>, session: bool) -> Self { + Self { role, session } + } + + pub fn role(role: Role) -> Self { + Self { + role: Some(role), + session: false, + } } } diff --git a/src/config/mod.rs b/src/config/mod.rs index e166969..4a1867f 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -3,7 +3,7 @@ mod role; mod session; pub use self::input::{Input, InputContext}; -pub use self::role::{Role, CODE_ROLE, EXPLAIN_ROLE, SHELL_ROLE}; +pub use self::role::{Role, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; use self::session::{Session, TEMP_SESSION_NAME}; use crate::client::{ @@ -190,7 +190,6 @@ impl Config { role.complete_prompt_args(name); role }) - .or_else(|| Role::find_system_role(name)) .ok_or_else(|| anyhow!("Unknown role `{name}`")) } @@ -682,10 +681,6 @@ impl Config { Ok(()) } - pub fn has_session(&self) -> bool { - self.session.is_some() - } - pub fn clear_session_messages(&mut self) -> Result<()> { if let Some(session) = self.session.as_mut() { session.clear_messages(); @@ -818,7 +813,7 @@ impl Config { } pub fn input_context(&self) -> InputContext { - InputContext::new(self.role.clone(), self.has_session()) + InputContext::new(self.role.clone(), self.session.is_some()) } pub fn maybe_print_send_tokens(&self, input: &Input) { @@ -978,7 +973,15 @@ impl Config { .with_context(|| format!("Failed to load roles at {}", path.display()))?; let roles: Vec<Role> = serde_yaml::from_str(&content).with_context(|| "Invalid roles config")?; + + let exist_roles: HashSet<_> = roles.iter().map(|v| v.name.clone()).collect(); self.roles = roles; + let builtin_roles = Role::builtin(); + for role in builtin_roles { + if !exist_roles.contains(&role.name) { + self.roles.push(role); + } + } Ok(()) } diff --git a/src/config/role.rs b/src/config/role.rs index 39b777c..4fc34c9 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -9,7 +9,7 @@ use serde::{Deserialize, Serialize}; pub const TEMP_ROLE: &str = "%%"; pub const SHELL_ROLE: &str = "%shell%"; -pub const EXPLAIN_ROLE: &str = "%explain%"; +pub const EXPLAIN_SHELL_ROLE: &str = "%explain-shell%"; pub const CODE_ROLE: &str = "%code%"; pub const INPUT_PLACEHOLDER: &str = "__INPUT__"; @@ -32,61 +32,20 @@ impl Role { } } - pub fn find_system_role(name: &str) -> Option<Self> { - match name { - SHELL_ROLE => Some(Self::shell()), - EXPLAIN_ROLE => Some(Self::explain()), - CODE_ROLE => Some(Self::code()), - _ => None, - } - } - - pub fn shell() -> Self { - let os = detect_os(); - let (detected_shell, _, _) = detect_shell(); - let (shell, use_semicolon) = match (detected_shell.as_str(), os.as_str()) { - // GPT doesn’t know much about nushell - ("nushell", "windows") => ("cmd", true), - ("nushell", _) => ("bash", true), - ("powershell", _) => ("powershell", true), - ("pwsh", _) => ("powershell", false), - _ => (detected_shell.as_str(), false), - }; - let combine = if use_semicolon { - "\nIf multiple steps required try to combine them together using ';'.\nIf it already combined with '&&' try to replace it with ';'.".to_string() - } else { - "\nIf multiple steps required try to combine them together using &&.".to_string() - }; - Self { - name: SHELL_ROLE.into(), - prompt: format!( - r#"Provide only {shell} commands for {os} without any description. -Ensure the output is a valid {shell} command. {combine} -If there is a lack of details, provide most logical solution. -Output plain text only, without any markdown formatting."# - ), - temperature: None, - top_p: None, - } - } - - pub fn explain() -> Self { - Self { - name: EXPLAIN_ROLE.into(), - prompt: r#"Provide a terse, single sentence description of the given shell command. + pub fn builtin() -> 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(), - temperature: None, - top_p: None, - } - } - - pub fn code() -> Self { - Self { - name: CODE_ROLE.into(), - prompt: r#"Provide only code without comments or explanations. + .into(), + ), + ( + CODE_ROLE, + r#"Provide only code without comments or explanations. ### INPUT: async sleep in js ### OUTPUT: @@ -96,10 +55,17 @@ async function timeout(ms) { } ``` "# - .into(), + .into(), + ), + ] + .into_iter() + .map(|(name, prompt)| Self { + name: name.into(), + prompt, temperature: None, top_p: None, - } + }) + .collect() } pub fn export(&self) -> Result<String> { @@ -241,6 +207,30 @@ fn parse_structure_prompt(prompt: &str) -> (&str, Vec<(&str, &str)>) { (prompt, vec![]) } +fn shell_prompt() -> String { + let os = detect_os(); + let (detected_shell, _, _) = detect_shell(); + let (shell, use_semicolon) = match (detected_shell.as_str(), os.as_str()) { + // GPT doesn’t know much about nushell + ("nushell", "windows") => ("cmd", true), + ("nushell", _) => ("bash", true), + ("powershell", _) => ("powershell", true), + ("pwsh", _) => ("powershell", false), + _ => (detected_shell.as_str(), false), + }; + let combine = if use_semicolon { + "\nIf multiple steps required try to combine them together using ';'.\nIf it already combined with '&&' try to replace it with ';'.".to_string() + } else { + "\nIf 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. {combine} +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::*; |
