diff options
| author | sigoden <sigoden@gmail.com> | 2024-03-03 09:57:50 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-03-03 09:57:50 +0800 |
| commit | 8421f23b450643ca3c66cb3f6fd21ef862a2369d (patch) | |
| tree | d64addd1104a01f3d60193f9912b4d0a3508edf3 /src | |
| parent | b2f86f2899b79b291eb56435f99402509105b7ec (diff) | |
| download | aichat-8421f23b450643ca3c66cb3f6fd21ef862a2369d.tar.gz | |
feat: allow overriding execute/code role (#331)
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/mod.rs | 14 | ||||
| -rw-r--r-- | src/config/role.rs | 12 | ||||
| -rw-r--r-- | src/main.rs | 28 | ||||
| -rw-r--r-- | src/utils/mod.rs | 8 |
4 files changed, 37 insertions, 25 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs index cbf1020..5220c5d 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -285,17 +285,23 @@ impl Config { } pub fn set_execute_role(&mut self) -> Result<()> { - let role = Role::for_execute(); + let role = self + .retrieve_role(Role::EXECUTE) + .unwrap_or_else(|_| Role::for_execute()); self.set_role_obj(role) } - pub fn set_describe_role(&mut self) -> Result<()> { - let role = Role::for_describe(); + pub fn set_describe_command_role(&mut self) -> Result<()> { + let role = self + .retrieve_role(Role::DESCRIBE_COMMAND) + .unwrap_or_else(|_| Role::for_describe_command()); self.set_role_obj(role) } pub fn set_code_role(&mut self) -> Result<()> { - let role = Role::for_code(); + let role = self + .retrieve_role(Role::CODE) + .unwrap_or_else(|_| Role::for_code()); self.set_role_obj(role) } diff --git a/src/config/role.rs b/src/config/role.rs index 39f0f17..1acf027 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -21,6 +21,10 @@ pub struct Role { } impl Role { + pub const EXECUTE: &'static str = "__execute__"; + pub const DESCRIBE_COMMAND: &'static str = "__describe_command__"; + pub const CODE: &'static str = "__code__"; + pub fn for_execute() -> Self { let os = detect_os(); let (shell, _, _) = detect_shell(); @@ -29,7 +33,7 @@ impl Role { _ => "&&", }; Self { - name: "__execute__".into(), + name: Self::EXECUTE.into(), prompt: format!( r#"Provide only {shell} commands for {os} without any description. If there is a lack of details, provide most logical solution. @@ -42,9 +46,9 @@ Do not provide markdown formatting such as ```"# } } - pub fn for_describe() -> Self { + pub fn for_describe_command() -> Self { Self { - name: "__describe__".into(), + name: Self::DESCRIBE_COMMAND.into(), prompt: 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. @@ -56,7 +60,7 @@ APPLY MARKDOWN formatting when possible."# pub fn for_code() -> Self { Self { - name: "__code__".into(), + name: Self::CODE.into(), prompt: r#"Provide only code as output without any description. Provide only code in plain text format without Markdown formatting. Do not include symbols such as ``` or ```python. diff --git a/src/main.rs b/src/main.rs index 30e7f8f..f6cafa3 100644 --- a/src/main.rs +++ b/src/main.rs @@ -11,7 +11,7 @@ mod utils; use crate::cli::Cli; use crate::config::{Config, GlobalConfig}; -use crate::utils::{extract_block, run_command}; +use crate::utils::{extract_block, run_command, CODE_BLOCK_RE}; use anyhow::{bail, Result}; use clap::Parser; @@ -60,19 +60,17 @@ fn main() -> Result<()> { if cli.dry_run { config.write().dry_run = true; } - if cli.execute { + if let Some(name) = &cli.role { + config.write().set_role(name)?; + } else if cli.execute { config.write().set_execute_role()?; - } else { - if let Some(name) = &cli.role { - config.write().set_role(name)?; - } else if cli.code { - config.write().set_code_role()?; - } - if let Some(session) = &cli.session { - config - .write() - .start_session(session.as_ref().map(|v| v.as_str()))?; - } + } else if cli.code { + config.write().set_code_role()?; + } + if let Some(session) = &cli.session { + config + .write() + .start_session(session.as_ref().map(|v| v.as_str()))?; } if let Some(model) = &cli.model { config.write().set_model(model)?; @@ -154,7 +152,7 @@ fn execute(config: &GlobalConfig, text: &str) -> Result<()> { let client = init_client(config)?; config.read().maybe_print_send_tokens(&input); let mut eval_str = client.send_message(input.clone())?; - if eval_str.contains("```") { + if let Ok(true) = CODE_BLOCK_RE.is_match(&eval_str) { eval_str = extract_block(&eval_str); } config.write().save_message(input, &eval_str)?; @@ -192,7 +190,7 @@ fn execute(config: &GlobalConfig, text: &str) -> Result<()> { } "D" | "d" => { if !describe { - config.write().set_describe_role()?; + config.write().set_describe_command_role()?; } let input = Input::from_str(&eval_str); let abort = create_abort_signal(); diff --git a/src/utils/mod.rs b/src/utils/mod.rs index da98cc6..15af4ca 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -17,7 +17,7 @@ use std::env; use std::process::Command; lazy_static! { - static ref CODE_BLOCK_RE: Regex = Regex::new(r"(?ms)```\w*(.*?)```").unwrap(); + pub static ref CODE_BLOCK_RE: Regex = Regex::new(r"(?ms)```\w*(.*)```").unwrap(); } pub fn now() -> String { @@ -165,7 +165,11 @@ pub fn extract_block(input: &str) -> String { .map(|m| String::from(m.as_str())) }) .collect(); - output.trim().to_string() + if output.is_empty() { + input.trim().to_string() + } else { + output.trim().to_string() + } } #[cfg(test)] |
