summaryrefslogtreecommitdiffstats
path: root/src/main.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-02-23 13:15:18 +0800
committerGitHub <noreply@github.com>2024-02-23 13:15:18 +0800
commit763841212826142ce248ef3d94524424c6316b0b (patch)
tree6805fb523b3fa149744d8534928941a023e0d149 /src/main.rs
parent6c0204e6965bf13c3d883ea5fa65d52415acd530 (diff)
downloadaichat-763841212826142ce248ef3d94524424c6316b0b.tar.gz
feat: support `-e/--execute` to execute shell command (#318)
Diffstat (limited to 'src/main.rs')
-rw-r--r--src/main.rs132
1 files changed, 101 insertions, 31 deletions
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)
+}