diff options
| author | sigoden <sigoden@gmail.com> | 2024-02-23 13:15:18 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-02-23 13:15:18 +0800 |
| commit | 763841212826142ce248ef3d94524424c6316b0b (patch) | |
| tree | 6805fb523b3fa149744d8534928941a023e0d149 /src | |
| parent | 6c0204e6965bf13c3d883ea5fa65d52415acd530 (diff) | |
| download | aichat-763841212826142ce248ef3d94524424c6316b0b.tar.gz | |
feat: support `-e/--execute` to execute shell command (#318)
Diffstat (limited to 'src')
| -rw-r--r-- | src/cli.rs | 4 | ||||
| -rw-r--r-- | src/config/mod.rs | 16 | ||||
| -rw-r--r-- | src/config/role.rs | 38 | ||||
| -rw-r--r-- | src/main.rs | 132 | ||||
| -rw-r--r-- | src/utils/mod.rs | 46 |
5 files changed, 203 insertions, 33 deletions
@@ -12,6 +12,9 @@ pub struct Cli { /// Create or reuse a session #[clap(short = 's', long)] pub session: Option<Option<String>>, + /// Execute commands using natural language + #[clap(short = 'e', long)] + pub execute: bool, /// Attach files to the message to be sent. #[clap(short = 'f', long, num_args = 1.., value_name = "FILE")] pub file: Option<Vec<String>>, @@ -43,6 +46,7 @@ pub struct Cli { #[clap(long)] pub list_sessions: bool, /// Input text + #[clap(trailing_var_arg = true)] text: Vec<String>, } diff --git a/src/config/mod.rs b/src/config/mod.rs index df74716..d38cbf3 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -158,7 +158,7 @@ impl Config { Ok(config) } - pub fn onstart(&mut self) -> Result<()> { + pub fn prelude(&mut self) -> Result<()> { let prelude = self.prelude.clone(); let err_msg = || format!("Invalid prelude '{}", prelude); match prelude.split_once(':') { @@ -275,6 +275,20 @@ impl Config { pub fn set_role(&mut self, name: &str) -> Result<()> { let role = self.retrieve_role(name)?; + self.set_role_obj(role) + } + + pub fn set_execute_role(&mut self) -> Result<()> { + let role = Role::for_execute(); + self.set_role_obj(role) + } + + pub fn set_describe_role(&mut self) -> Result<()> { + let role = Role::for_describe(); + self.set_role_obj(role) + } + + pub fn set_role_obj(&mut self, role: Role) -> Result<()> { if let Some(session) = self.session.as_mut() { session.update_role(Some(role.clone()))?; } diff --git a/src/config/role.rs b/src/config/role.rs index bd7216c..ff60b55 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -1,4 +1,7 @@ -use crate::client::{Message, MessageContent, MessageRole}; +use crate::{ + client::{Message, MessageContent, MessageRole}, + utils::{detect_os, detect_shell}, +}; use anyhow::{Context, Result}; use serde::{Deserialize, Serialize}; @@ -18,6 +21,39 @@ pub struct Role { } impl Role { + pub fn for_execute() -> Self { + let os = detect_os(); + let shell = detect_shell(); + let shell = match shell.rsplit_once('/') { + Some((_, v)) => v, + None => &shell, + }; + Self { + name: "__builtin__".into(), + prompt: format!( + r#"Provide only {shell} commands for {os} without any description. +If there is a lack of details, provide most logical solution. +Ensure the output is a valid shell command. +If multiple steps required try to combine them together using &&. +Provide only plain text without Markdown formatting. +Do not provide markdown formatting such as ```"# + ), + temperature: None, + } + } + + pub fn for_describe() -> Self { + Self { + name: "__builtin__".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. +APPLY MARKDOWN formatting when possible."# + .into(), + temperature: None, + } + } + pub fn info(&self) -> Result<String> { let output = serde_yaml::to_string(&self) .with_context(|| format!("Unable to show info about role {}", &self.name))?; diff --git a/src/main.rs b/src/main.rs index 5d8315d..8946985 100644 --- a/src/main.rs +++ b/src/main.rs @@ -11,16 +11,20 @@ mod utils; use crate::cli::Cli; use crate::config::{Config, GlobalConfig}; +use crate::utils::{prompt_op_err, run_command}; -use anyhow::Result; +use anyhow::{bail, Result}; use clap::Parser; use client::{ensure_model_capabilities, init_client, list_models}; use config::Input; +use inquire::validator::Validation; +use inquire::Text; use is_terminal::IsTerminal; use parking_lot::RwLock; use render::{render_error, render_stream, MarkdownRender}; use repl::Repl; use std::io::{stderr, stdin, stdout, Read}; +use std::process; use std::sync::Arc; use utils::{cl100k_base_singleton, create_abort_signal}; @@ -56,13 +60,17 @@ fn main() -> Result<()> { if cli.dry_run { config.write().dry_run = true; } - if let Some(name) = &cli.role { - config.write().set_role(name)?; - } - if let Some(session) = &cli.session { - config - .write() - .start_session(session.as_ref().map(|v| v.as_str()))?; + if cli.execute { + config.write().set_execute_role()?; + } else { + if let Some(name) = &cli.role { + config.write().set_role(name)?; + } + 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)?; @@ -75,35 +83,27 @@ fn main() -> Result<()> { println!("{}", info); return Ok(()); } - config.write().onstart()?; - if let Err(err) = start(&config, text, cli.file, cli.no_stream) { + let text = aggregate_text(text)?; + if cli.execute { + match text { + Some(text) => { + execute(&config, &text)?; + return Ok(()); + } + None => bail!("No input text"), + } + } + config.write().prelude()?; + if let Err(err) = match text { + Some(text) => start_directive(&config, &text, cli.file, cli.no_stream), + None => start_interactive(&config), + } { let highlight = stderr().is_terminal() && config.read().highlight; render_error(err, highlight) } Ok(()) } -fn start( - config: &GlobalConfig, - text: Option<String>, - include: Option<Vec<String>>, - no_stream: bool, -) -> Result<()> { - if stdin().is_terminal() { - match text { - Some(text) => start_directive(config, &text, include, no_stream), - None => start_interactive(config), - } - } else { - let mut input = String::new(); - stdin().read_to_string(&mut input)?; - if let Some(text) = text { - input = format!("{text}\n{input}"); - } - start_directive(config, &input, include, no_stream) - } -} - fn start_directive( config: &GlobalConfig, text: &str, @@ -139,3 +139,73 @@ fn start_interactive(config: &GlobalConfig) -> Result<()> { let mut repl: Repl = Repl::init(config)?; repl.run() } + +fn execute(config: &GlobalConfig, text: &str) -> Result<()> { + let input = Input::from_str(text); + let client = init_client(config)?; + config.read().maybe_print_send_tokens(&input); + let eval_str = client.send_message(input.clone())?; + let render_options = config.read().get_render_options()?; + let mut markdown_render = MarkdownRender::init(render_options)?; + if config.read().dry_run { + println!("{}", markdown_render.render(&eval_str).trim()); + return Ok(()); + } + if stdout().is_terminal() { + println!("{}", markdown_render.render(&eval_str).trim()); + let mut describe = false; + loop { + let anwser = Text::new("[e]xecute, [d]escribe, [a]bort: ") + .with_default("e") + .with_validator(|input: &str| { + match matches!(input, "E" | "e" | "D" | "d" | "A" | "a") { + true => Ok(Validation::Valid), + false => Ok(Validation::Invalid( + "Invalid input, choice one of e, d or a".into(), + )), + } + }) + .prompt() + .map_err(prompt_op_err)?; + + match anwser.as_str() { + "E" | "e" => { + let code = run_command(&eval_str)?; + if code != 0 { + process::exit(code); + } + } + "D" | "d" => { + if !describe { + config.write().set_describe_role()?; + } + let input = Input::from_str(&eval_str); + let abort = create_abort_signal(); + render_stream(&input, client.as_ref(), config, abort)?; + describe = true; + continue; + } + _ => {} + } + break; + } + } else { + println!("{}", eval_str); + } + Ok(()) +} + +fn aggregate_text(text: Option<String>) -> Result<Option<String>> { + let text = if stdin().is_terminal() { + text + } else { + let mut stdin_text = String::new(); + stdin().read_to_string(&mut stdin_text)?; + if let Some(text) = text { + Some(format!("{text}\n{stdin_text}")) + } else { + Some(stdin_text) + } + }; + Ok(text) +} diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 04c5668..7dacd5c 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -11,6 +11,8 @@ pub use self::render_prompt::render_prompt; pub use self::tiktoken::cl100k_base_singleton; use sha2::{Digest, Sha256}; +use std::env; +use std::process::Command; pub fn now() -> String { let now = chrono::Local::now(); @@ -87,6 +89,50 @@ pub fn sha256sum(input: &str) -> String { format!("{:x}", result) } +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 fn detect_shell() -> String { + let os = env::consts::OS; + if os == "windows" { + if let Some(true) = env::var("PSModulePath") + .ok() + .map(|v| v.split(';').count() >= 3) + { + "powershell.exe".into() + } else { + "cmd.exe".into() + } + } else { + env::var("SHELL").unwrap_or_else(|_| "/bin/sh".to_string()) + } +} + +pub fn run_command(eval_str: &str) -> anyhow::Result<i32> { + let shell = detect_shell(); + let mut command = Command::new(&shell); + if shell == "powershell.exe" { + command.arg("-Command").arg(eval_str); + } else if shell == "cmd.exe" { + command.arg("/c").arg(eval_str); + } else { + command.arg("-c").arg(eval_str); + }; + let status = command.status()?; + Ok(status.code().unwrap_or_default()) +} + #[cfg(test)] mod tests { use super::*; |
