summaryrefslogtreecommitdiffstats
path: root/src/function.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-07-10 18:53:58 +0800
committerGitHub <noreply@github.com>2024-07-10 18:53:58 +0800
commit4ea6f3933fff4974334e3f9ecb2a423aab43a787 (patch)
tree37e27557072165ebd7a61b27e21b7e3e55d713dd /src/function.rs
parent48219d2808bafc88fa4fa4d6997aa9dc92e6f686 (diff)
downloadaichat-4ea6f3933fff4974334e3f9ecb2a423aab43a787.tar.gz
feat: adjust the way of obtaining function call results (#695)
Diffstat (limited to 'src/function.rs')
-rw-r--r--src/function.rs78
1 files changed, 14 insertions, 64 deletions
diff --git a/src/function.rs b/src/function.rs
index 0016405..6a7e81e 100644
--- a/src/function.rs
+++ b/src/function.rs
@@ -5,7 +5,6 @@ use crate::{
use anyhow::{anyhow, bail, Context, Result};
use indexmap::IndexMap;
-use inquire::{validator::Validation, Text};
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use std::{
@@ -157,7 +156,6 @@ 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().agent {
Some(agent) => {
if let Some(true) = agent.functions().find(&function_name).map(|v| v.agent) {
@@ -205,76 +203,28 @@ impl ToolCall {
if bin_dir.exists() {
envs.insert("PATH".into(), prepend_env_path(&bin_dir)?);
}
+ let temp_file = temp_file("-eval-", "");
+ envs.insert("LLM_OUTPUT".into(), temp_file.display().to_string());
#[cfg(windows)]
let cmd_name = polyfill_cmd_name(&cmd_name, &bin_dir);
-
- let output = if is_dangerously {
- if *IS_STDOUT_TERMINAL {
- println!("{}", dimmed_text(&prompt));
- let answer = Text::new("[1] Run, [2] Run & Retrieve, [3] Skip:")
- .with_default("1")
- .with_validator(|input: &str| match matches!(input, "1" | "2" | "3") {
- true => Ok(Validation::Valid),
- false => Ok(Validation::Invalid(
- "Invalid input, please select 1, 2 or 3".into(),
- )),
- })
- .prompt()?;
- match answer.as_str() {
- "1" => {
- let exit_code = run_command(&cmd_name, &cmd_args, Some(envs))?;
- if exit_code != 0 {
- bail!("Exit {exit_code}");
- }
- Value::Null
- }
- "2" => run_and_retrieve(&cmd_name, &cmd_args, envs)?,
- _ => Value::Null,
- }
- } else {
- println!("Skipped {prompt}");
- Value::Null
- }
- } else {
- println!("{}", dimmed_text(&prompt));
- run_and_retrieve(&cmd_name, &cmd_args, envs)?
- };
-
- Ok(output)
- }
-}
-
-fn run_and_retrieve(
- cmd_name: &str,
- cmd_args: &[String],
- envs: HashMap<String, String>,
-) -> 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(&stderr));
+ println!("{}", dimmed_text(&prompt));
+ let exit_code = run_command(&cmd_name, &cmd_args, Some(envs))?;
+ if exit_code != 0 {
+ bail!("Tool call exit with {exit_code}");
}
- let value = if !stdout.is_empty() {
- serde_json::from_str(&stdout)
+ let output = if temp_file.exists() {
+ let contents =
+ fs::read_to_string(temp_file).context("Failed to retrieve tool call output")?;
+
+ serde_json::from_str(&contents)
.ok()
- .unwrap_or_else(|| json!({"output": stdout}))
+ .unwrap_or_else(|| json!({"result": contents}))
} else {
Value::Null
};
- Ok(value)
- } else {
- let err = if stderr.is_empty() {
- if stdout.is_empty() {
- "Something wrong"
- } else {
- &stdout
- }
- } else {
- &stderr
- };
- bail!("{err}");
+
+ Ok(output)
}
}