diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-18 19:06:21 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-18 19:06:21 +0800 |
| commit | b4a40e3fedb438570770a224b890ea24f6e660a9 (patch) | |
| tree | 344b96102da7cbedf1034d023aa82599940388b1 /src/utils | |
| parent | 1348a62e5f8bc140a7218fbfe1b73f990ab16101 (diff) | |
| download | aichat-b4a40e3fedb438570770a224b890ea24f6e660a9.tar.gz | |
feat: support function calling (#514)
* feat: support function calling
* fix on Windows OS
* implement multi-steps function calling
* fix on Windows OS
* add error for client not support function calling
* refactor message data structure and make claude client supporting function calling
* support reuse previous call results
* improve error handling for function calling
* use prefix `may_` as indicator for `execute` type fucntions
Diffstat (limited to 'src/utils')
| -rw-r--r-- | src/utils/command.rs | 82 | ||||
| -rw-r--r-- | src/utils/mod.rs | 90 |
2 files changed, 110 insertions, 62 deletions
diff --git a/src/utils/command.rs b/src/utils/command.rs new file mode 100644 index 0000000..77ca813 --- /dev/null +++ b/src/utils/command.rs @@ -0,0 +1,82 @@ +use std::{collections::HashMap, env, ffi::OsStr, process::Command}; + +use anyhow::{Context, 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, String, &'static str) { + let os = env::consts::OS; + if os == "windows" { + if env::var("NU_VERSION").is_ok() { + ("nushell".into(), "nu.exe".into(), "-c") + } else if let Some(ret) = env::var("PSModulePath").ok().and_then(|v| { + let v = v.to_lowercase(); + if v.split(';').count() >= 3 { + if v.contains("powershell\\7\\") { + Some(("pwsh".into(), "pwsh.exe".into(), "-c")) + } else { + Some(("powershell".into(), "powershell.exe".into(), "-Command")) + } + } else { + None + } + }) { + ret + } else { + ("cmd".into(), "cmd.exe".into(), "/C") + } + } else if env::var("NU_VERSION").is_ok() { + ("nushell".into(), "nu".into(), "-c") + } else { + let shell_cmd = env::var("SHELL").unwrap_or_else(|_| "/bin/sh".to_string()); + let shell_name = match shell_cmd.rsplit_once('/') { + Some((_, name)) => name.to_string(), + None => shell_cmd.clone(), + }; + let shell_name = if shell_name == "nu" { + "nushell".into() + } else { + shell_name + }; + (shell_name, shell_cmd, "-c") + } +} + +pub fn run_command<T: AsRef<OsStr>>( + cmd: &str, + args: &[T], + envs: Option<HashMap<String, String>>, +) -> Result<i32> { + let status = Command::new(cmd) + .args(args.iter()) + .envs(envs.unwrap_or_default()) + .status()?; + Ok(status.code().unwrap_or_default()) +} + +pub fn run_command_with_output<T: AsRef<OsStr>>( + cmd: &str, + args: &[T], + envs: Option<HashMap<String, String>>, +) -> Result<(bool, String, String)> { + let output = Command::new(cmd) + .args(args.iter()) + .envs(envs.unwrap_or_default()) + .output()?; + let status = output.status; + let stdout = std::str::from_utf8(&output.stdout).context("Invalid UTF-8 in stdout")?; + let stderr = std::str::from_utf8(&output.stderr).context("Invalid UTF-8 in stderr")?; + Ok((status.success(), stdout.to_string(), stderr.to_string())) +} diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 31bb971..3ce2a78 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -1,5 +1,6 @@ mod abort_signal; mod clipboard; +mod command; mod crypto; mod prompt_input; mod render_prompt; @@ -7,6 +8,7 @@ mod spinner; pub use self::abort_signal::{create_abort_signal, AbortSignal}; pub use self::clipboard::set_text; +pub use self::command::*; pub use self::crypto::*; pub use self::prompt_input::*; pub use self::render_prompt::render_prompt; @@ -15,7 +17,6 @@ pub use self::spinner::run_spinner; use fancy_regex::Regex; use lazy_static::lazy_static; use std::env; -use std::process::Command; lazy_static! { pub static ref CODE_BLOCK_RE: Regex = Regex::new(r"(?ms)```\w*(.*)```").unwrap(); @@ -79,67 +80,6 @@ pub fn light_theme_from_colorfgbg(colorfgbg: &str) -> Option<bool> { Some(light) } -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, String, &'static str) { - let os = env::consts::OS; - if os == "windows" { - if env::var("NU_VERSION").is_ok() { - ("nushell".into(), "nu.exe".into(), "-c") - } else if let Some(ret) = env::var("PSModulePath").ok().and_then(|v| { - let v = v.to_lowercase(); - if v.split(';').count() >= 3 { - if v.contains("powershell\\7\\") { - Some(("pwsh".into(), "pwsh.exe".into(), "-c")) - } else { - Some(("powershell".into(), "powershell.exe".into(), "-Command")) - } - } else { - None - } - }) { - ret - } else { - ("cmd".into(), "cmd.exe".into(), "/C") - } - } else if env::var("NU_VERSION").is_ok() { - ("nushell".into(), "nu".into(), "-c") - } else { - let shell_cmd = env::var("SHELL").unwrap_or_else(|_| "/bin/sh".to_string()); - let shell_name = match shell_cmd.rsplit_once('/') { - Some((_, name)) => name.to_string(), - None => shell_cmd.clone(), - }; - let shell_name = if shell_name == "nu" { - "nushell".into() - } else { - shell_name - }; - (shell_name, shell_cmd, "-c") - } -} - -pub fn run_command(eval_str: &str) -> anyhow::Result<i32> { - let (_shell_name, shell_cmd, shell_arg) = detect_shell(); - let status = Command::new(shell_cmd) - .arg(shell_arg) - .arg(eval_str) - .status()?; - Ok(status.code().unwrap_or_default()) -} - pub fn extract_block(input: &str) -> String { let output: String = CODE_BLOCK_RE .captures_iter(input) @@ -183,6 +123,32 @@ pub fn fuzzy_match(text: &str, pattern: &str) -> bool { pattern_index == pattern_chars.len() } +pub fn error_text(input: &str) -> String { + nu_ansi_term::Style::new() + .fg(nu_ansi_term::Color::Red) + .paint(input) + .to_string() +} + +pub fn warning_text(input: &str) -> String { + nu_ansi_term::Style::new() + .fg(nu_ansi_term::Color::Yellow) + .paint(input) + .to_string() +} + +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::*; |
