From b4a40e3fedb438570770a224b890ea24f6e660a9 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 18 May 2024 19:06:21 +0800 Subject: 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 --- src/utils/command.rs | 82 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 82 insertions(+) create mode 100644 src/utils/command.rs (limited to 'src/utils/command.rs') 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>( + cmd: &str, + args: &[T], + envs: Option>, +) -> Result { + 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>( + cmd: &str, + args: &[T], + envs: Option>, +) -> 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())) +} -- cgit v1.2.3