diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-11 14:01:45 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-11 14:01:45 +0800 |
| commit | 861529374726609a56f6b8af1145a0df5c924dcc (patch) | |
| tree | 78c7636f7e65d8bea24a7566010b2663840b9e2a /src | |
| parent | 822688a06a2dd2c45c7a1d5fd32f8f0d415d8620 (diff) | |
| download | aichat-861529374726609a56f6b8af1145a0df5c924dcc.tar.gz | |
feat: add config `dangerously_functions` (#582)
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/bot.rs | 6 | ||||
| -rw-r--r-- | src/config/mod.rs | 30 | ||||
| -rw-r--r-- | src/function.rs | 31 | ||||
| -rw-r--r-- | src/utils/mod.rs | 8 |
4 files changed, 40 insertions, 35 deletions
diff --git a/src/config/bot.rs b/src/config/bot.rs index 19fd5cd..80d3a5d 100644 --- a/src/config/bot.rs +++ b/src/config/bot.rs @@ -105,6 +105,10 @@ impl Bot { &self.name } + pub fn config(&self) -> &BotConfig { + &self.config + } + pub fn functions(&self) -> &Functions { &self.functions } @@ -170,6 +174,8 @@ pub struct BotConfig { pub temperature: Option<f64>, #[serde(skip_serializing_if = "Option::is_none")] pub top_p: Option<f64>, + #[serde(skip_serializing_if = "Option::is_none")] + pub dangerously_functions: Option<FunctionsFilter>, } impl BotConfig { diff --git a/src/config/mod.rs b/src/config/mod.rs index 7c99acb..3af4467 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -12,15 +12,13 @@ use crate::client::{ create_client_config, list_chat_models, list_client_types, ClientConfig, Model, OPENAI_COMPATIBLE_PLATFORMS, }; -use crate::function::{FunctionDeclaration, Functions, ToolCallResult}; +use crate::function::{FunctionDeclaration, Functions, FunctionsFilter, ToolCallResult}; use crate::rag::Rag; use crate::render::{MarkdownRender, RenderOptions}; -use crate::utils::{ - format_option_value, fuzzy_match, get_env_name, light_theme_from_colorfgbg, now, render_prompt, - set_text, warning_text, AbortSignal, IS_STDOUT_TERMINAL, -}; +use crate::utils::*; use anyhow::{anyhow, bail, Context, Result}; +use fancy_regex::Regex; use inquire::{Confirm, Select}; use parking_lot::RwLock; use serde::Deserialize; @@ -96,6 +94,7 @@ pub struct Config { pub rag_top_k: usize, pub rag_template: Option<String>, pub function_calling: bool, + pub dangerously_functions: Option<FunctionsFilter>, pub compress_threshold: usize, pub summarize_prompt: Option<String>, pub summary_prompt: Option<String>, @@ -144,6 +143,7 @@ impl Default for Config { rag_top_k: 4, rag_template: None, function_calling: false, + dangerously_functions: None, compress_threshold: 4000, summarize_prompt: None, summary_prompt: None, @@ -965,6 +965,26 @@ impl Config { functions } + pub fn is_dangerously_function(&self, name: &str) -> bool { + if get_env_bool("no_dangerously_functions") { + return false; + } + let dangerously_functions = match &self.bot { + Some(bot) => bot.config().dangerously_functions.as_ref(), + None => self.dangerously_functions.as_ref(), + }; + match dangerously_functions { + None => false, + Some(regex) => { + let regex = match Regex::new(&format!("^({regex})$")) { + Ok(v) => v, + Err(_) => return false, + }; + regex.is_match(name).unwrap_or(false) + } + } + } + pub fn buffer_editor(&self) -> Option<String> { self.buffer_editor .clone() diff --git a/src/function.rs b/src/function.rs index 0896540..f16fe3f 100644 --- a/src/function.rs +++ b/src/function.rs @@ -1,9 +1,6 @@ use crate::{ config::{Config, GlobalConfig}, - utils::{ - dimmed_text, get_env_bool, indent_text, run_command, run_command_with_output, warning_text, - IS_STDOUT_TERMINAL, - }, + utils::*, }; use anyhow::{anyhow, bail, Context, Result}; @@ -20,6 +17,7 @@ use std::{ pub const FUNCTION_ALL_MATCHER: &str = ".*"; pub type ToolResults = (Vec<ToolCallResult>, String); +pub type FunctionsFilter = String; pub fn eval_tool_calls( config: &GlobalConfig, @@ -171,6 +169,7 @@ impl ToolCall { pub fn eval(&self, config: &GlobalConfig) -> Result<Value> { let function_name = self.name.clone(); + let is_dangerously = config.read().is_dangerously_function(&function_name); let (call_name, cmd_name, mut cmd_args) = match &config.read().bot { Some(bot) => { if !bot.functions().contains(&function_name) { @@ -219,11 +218,11 @@ impl ToolCall { #[cfg(windows)] let cmd_name = polyfill_cmd_name(&cmd_name, &bin_dir); - let output = if self.is_execute() { + let output = if is_dangerously { if *IS_STDOUT_TERMINAL { println!("{prompt}"); let answer = Text::new("[1] Run, [2] Run & Retrieve, [3] Skip:") - .with_default("1") + .with_default("2") .with_validator(|input: &str| match matches!(input, "1" | "2" | "3") { true => Ok(Validation::Valid), false => Ok(Validation::Invalid( @@ -239,7 +238,7 @@ impl ToolCall { } Value::Null } - "2" => run_and_retrieve(&cmd_name, &cmd_args, envs, &prompt)?, + "2" => run_and_retrieve(&cmd_name, &cmd_args, envs)?, _ => Value::Null, } } else { @@ -248,35 +247,23 @@ impl ToolCall { } } else { println!("{}", dimmed_text(&prompt)); - run_and_retrieve(&cmd_name, &cmd_args, envs, &prompt)? + run_and_retrieve(&cmd_name, &cmd_args, envs)? }; Ok(output) } - - pub fn is_execute(&self) -> bool { - if get_env_bool("function_auto_execute") { - false - } else { - self.name.starts_with("may_") || self.name.contains("__may_") - } - } } fn run_and_retrieve( cmd_name: &str, cmd_args: &[String], envs: HashMap<String, String>, - prompt: &str, ) -> Result<Value> { let (success, stdout, stderr) = run_command_with_output(cmd_name, cmd_args, Some(envs))?; if success { if !stderr.is_empty() { - eprintln!( - "{}", - warning_text(&format!("{prompt}:\n{}", indent_text(&stderr, 4))) - ); + eprintln!("{}", warning_text(&stderr)); } let value = if !stdout.is_empty() { serde_json::from_str(&stdout) @@ -296,7 +283,7 @@ fn run_and_retrieve( } else { &stderr }; - bail!("{}", &format!("{prompt}:\n{}", indent_text(err, 4))); + bail!("{err}"); } } diff --git a/src/utils/mod.rs b/src/utils/mod.rs index cd00cc5..5c34a54 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -151,14 +151,6 @@ pub fn dimmed_text(input: &str) -> String { nu_ansi_term::Style::new().dimmed().paint(input).to_string() } -pub fn indent_text(text: &str, spaces: usize) -> String { - let indent_size = " ".repeat(spaces); - text.lines() - .map(|line| format!("{}{}", indent_size, line)) - .collect::<Vec<String>>() - .join("\n") -} - #[cfg(test)] mod tests { use super::*; |
