summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-03-03 09:57:50 +0800
committerGitHub <noreply@github.com>2024-03-03 09:57:50 +0800
commit8421f23b450643ca3c66cb3f6fd21ef862a2369d (patch)
treed64addd1104a01f3d60193f9912b4d0a3508edf3 /src
parentb2f86f2899b79b291eb56435f99402509105b7ec (diff)
downloadaichat-8421f23b450643ca3c66cb3f6fd21ef862a2369d.tar.gz
feat: allow overriding execute/code role (#331)
Diffstat (limited to 'src')
-rw-r--r--src/config/mod.rs14
-rw-r--r--src/config/role.rs12
-rw-r--r--src/main.rs28
-rw-r--r--src/utils/mod.rs8
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)]