diff options
| author | sigoden <sigoden@gmail.com> | 2024-03-02 21:11:28 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-03-02 21:11:28 +0800 |
| commit | 3c16aff59145134885f32056fa3be9bd6bed2cce (patch) | |
| tree | f285ee126374b8815840815c7334e1d48fae2791 /src | |
| parent | d275c33632271bf9cdbc01ef005deb9ff851370b (diff) | |
| download | aichat-3c16aff59145134885f32056fa3be9bd6bed2cce.tar.gz | |
feat: add `-c/--code` to generate only code (#327)
Diffstat (limited to 'src')
| -rw-r--r-- | src/cli.rs | 3 | ||||
| -rw-r--r-- | src/config/mod.rs | 5 | ||||
| -rw-r--r-- | src/config/role.rs | 18 | ||||
| -rw-r--r-- | src/main.rs | 25 | ||||
| -rw-r--r-- | src/utils/mod.rs | 18 |
5 files changed, 60 insertions, 9 deletions
@@ -15,6 +15,9 @@ pub struct Cli { /// Execute commands using natural language #[clap(short = 'e', long)] pub execute: bool, + /// Generate only code + #[clap(short = 'c', long)] + pub code: bool, /// Attach files to the message to be sent. #[clap(short = 'f', long, num_args = 1.., value_name = "FILE")] pub file: Option<Vec<String>>, diff --git a/src/config/mod.rs b/src/config/mod.rs index f8d97c2..6b99483 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -288,6 +288,11 @@ impl Config { self.set_role_obj(role) } + pub fn set_code_role(&mut self) -> Result<()> { + let role = Role::for_code(); + 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 a63877b..39f0f17 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -29,7 +29,7 @@ impl Role { _ => "&&", }; Self { - name: "__for_execute__".into(), + name: "__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. @@ -44,7 +44,7 @@ Do not provide markdown formatting such as ```"# pub fn for_describe() -> Self { Self { - name: "__for_describe__".into(), + name: "__describe__".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. @@ -54,6 +54,20 @@ APPLY MARKDOWN formatting when possible."# } } + pub fn for_code() -> Self { + Self { + name: "__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. +If there is a lack of details, provide most logical solution. +You are not allowed to ask for more details. +For example if the prompt is "Hello world Python", you should return "print('Hello world')"."# + .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 bd319e1..f9be816 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::run_command; +use crate::utils::{extract_block, run_command}; use anyhow::{bail, Result}; use clap::Parser; @@ -65,6 +65,8 @@ fn main() -> Result<()> { } 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 @@ -95,7 +97,7 @@ fn main() -> Result<()> { } config.write().prelude()?; if let Err(err) = match text { - Some(text) => start_directive(&config, &text, cli.file, cli.no_stream), + Some(text) => start_directive(&config, &text, cli.file, cli.no_stream, cli.code), None => start_interactive(&config), } { let highlight = stderr().is_terminal() && config.read().highlight; @@ -109,6 +111,7 @@ fn start_directive( text: &str, include: Option<Vec<String>>, no_stream: bool, + code_mode: bool, ) -> Result<()> { if let Some(session) = &config.read().session { session.guard_save()?; @@ -117,14 +120,19 @@ fn start_directive( let mut client = init_client(config)?; ensure_model_capabilities(client.as_mut(), input.required_capabilities())?; config.read().maybe_print_send_tokens(&input); - let output = if no_stream { + let output = if !stdout().is_terminal() || no_stream { let output = client.send_message(input.clone())?; - if stdout().is_terminal() { + let to_print = if code_mode && output.trim_start().starts_with("```") { + extract_block(&output) + } else { + output.clone() + }; + if no_stream { let render_options = config.read().get_render_options()?; let mut markdown_render = MarkdownRender::init(render_options)?; - println!("{}", markdown_render.render(&output).trim()); + println!("{}", markdown_render.render(&to_print).trim()); } else { - println!("{}", output); + println!("{}", to_print); } output } else { @@ -144,7 +152,10 @@ 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 mut eval_str = client.send_message(input.clone())?; + if eval_str.contains("```") { + eval_str = extract_block(&eval_str); + } config.write().save_message(input, &eval_str)?; let render_options = config.read().get_render_options()?; let mut markdown_render = MarkdownRender::init(render_options)?; diff --git a/src/utils/mod.rs b/src/utils/mod.rs index b59b7ad..da98cc6 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -10,10 +10,16 @@ pub use self::prompt_input::*; pub use self::render_prompt::render_prompt; pub use self::tiktoken::cl100k_base_singleton; +use fancy_regex::Regex; +use lazy_static::lazy_static; use sha2::{Digest, Sha256}; use std::env; use std::process::Command; +lazy_static! { + static ref CODE_BLOCK_RE: Regex = Regex::new(r"(?ms)```\w*(.*?)```").unwrap(); +} + pub fn now() -> String { let now = chrono::Local::now(); now.to_rfc3339_opts(chrono::SecondsFormat::Secs, false) @@ -150,6 +156,18 @@ pub fn run_command(eval_str: &str) -> anyhow::Result<i32> { Ok(status.code().unwrap_or_default()) } +pub fn extract_block(input: &str) -> String { + let output: String = CODE_BLOCK_RE + .captures_iter(input) + .filter_map(|m| { + m.ok() + .and_then(|cap| cap.get(1)) + .map(|m| String::from(m.as_str())) + }) + .collect(); + output.trim().to_string() +} + #[cfg(test)] mod tests { use super::*; |
