summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-03-02 21:11:28 +0800
committerGitHub <noreply@github.com>2024-03-02 21:11:28 +0800
commit3c16aff59145134885f32056fa3be9bd6bed2cce (patch)
treef285ee126374b8815840815c7334e1d48fae2791 /src
parentd275c33632271bf9cdbc01ef005deb9ff851370b (diff)
downloadaichat-3c16aff59145134885f32056fa3be9bd6bed2cce.tar.gz
feat: add `-c/--code` to generate only code (#327)
Diffstat (limited to 'src')
-rw-r--r--src/cli.rs3
-rw-r--r--src/config/mod.rs5
-rw-r--r--src/config/role.rs18
-rw-r--r--src/main.rs25
-rw-r--r--src/utils/mod.rs18
5 files changed, 60 insertions, 9 deletions
diff --git a/src/cli.rs b/src/cli.rs
index 808602e..5f35df4 100644
--- a/src/cli.rs
+++ b/src/cli.rs
@@ -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::*;