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/main.rs | 79 ++++++++++++++++++++++++++++++++++++++++++------------------- 1 file changed, 54 insertions(+), 25 deletions(-) (limited to 'src/main.rs') diff --git a/src/main.rs b/src/main.rs index 25895d7..895617d 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,6 +1,7 @@ mod cli; mod client; mod config; +mod function; mod logger; mod render; mod repl; @@ -12,17 +13,20 @@ mod utils; extern crate log; use crate::cli::Cli; -use crate::client::{ensure_model_capabilities, list_models, send_stream}; +use crate::client::{list_models, send_stream, CompletionOutput}; use crate::config::{ Config, GlobalConfig, Input, InputContext, WorkingMode, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE, }; +use crate::function::eval_tool_calls; use crate::render::{render_error, MarkdownRender}; use crate::repl::Repl; use crate::utils::{create_abort_signal, extract_block, run_command, run_spinner, CODE_BLOCK_RE}; use anyhow::{bail, Result}; +use async_recursion::async_recursion; use clap::Parser; +use function::need_send_call_results; use inquire::{Select, Text}; use is_terminal::IsTerminal; use parking_lot::RwLock; @@ -30,6 +34,7 @@ use std::io::{stderr, stdin, stdout, Read}; use std::process; use std::sync::Arc; use tokio::sync::oneshot; +use utils::detect_shell; #[tokio::main] async fn main() -> Result<()> { @@ -113,7 +118,8 @@ async fn main() -> Result<()> { bail!("No input"); } let input = create_input(&config, text, file)?; - execute(&config, input).await?; + let (_, shell, shell_arg) = detect_shell(); + shell_execute(&config, &shell, shell_arg, input).await?; return Ok(()); } config.write().apply_prelude()?; @@ -130,39 +136,56 @@ async fn main() -> Result<()> { Ok(()) } +#[async_recursion] async fn start_directive( config: &GlobalConfig, input: Input, no_stream: bool, code_mode: bool, ) -> Result<()> { - let mut client = input.create_client()?; - ensure_model_capabilities(client.as_mut(), input.required_capabilities())?; + let client = input.create_client()?; let is_terminal_stdout = stdout().is_terminal(); let extract_code = !is_terminal_stdout && code_mode; - let output = if no_stream || extract_code { - let (output, _) = client.send_message(input.clone()).await?; - let output = if extract_code && output.trim_start().starts_with("```") { - extract_block(&output) - } else { - output.clone() - }; - if is_terminal_stdout { - let render_options = config.read().get_render_options()?; - let mut markdown_render = MarkdownRender::init(render_options)?; - println!("{}", markdown_render.render(&output).trim()); + let (output, tool_call_results) = if no_stream || extract_code { + let CompletionOutput { + text, tool_calls, .. + } = client.send_message(input.clone()).await?; + if !tool_calls.is_empty() { + (String::new(), eval_tool_calls(config, tool_calls)?) } else { - println!("{}", output); + let text = if extract_code && text.trim_start().starts_with("```") { + extract_block(&text) + } else { + text.clone() + }; + if is_terminal_stdout { + let render_options = config.read().get_render_options()?; + let mut markdown_render = MarkdownRender::init(render_options)?; + println!("{}", markdown_render.render(&text).trim()); + } else { + println!("{}", text); + } + (text, vec![]) } - output } else { let abort = create_abort_signal(); send_stream(&input, client.as_ref(), config, abort).await? }; - // Save the message/session - config.write().save_message(input, &output)?; + config + .write() + .save_message(&input, &output, &tool_call_results)?; config.write().end_session()?; - Ok(()) + if need_send_call_results(&tool_call_results) { + start_directive( + config, + input.merge_tool_call(output, tool_call_results), + no_stream, + code_mode, + ) + .await + } else { + Ok(()) + } } async fn start_interactive(config: &GlobalConfig) -> Result<()> { @@ -171,7 +194,12 @@ async fn start_interactive(config: &GlobalConfig) -> Result<()> { } #[async_recursion::async_recursion] -async fn execute(config: &GlobalConfig, mut input: Input) -> Result<()> { +async fn shell_execute( + config: &GlobalConfig, + shell: &str, + shell_arg: &str, + mut input: Input, +) -> Result<()> { let client = input.create_client()?; let is_terminal_stdout = stdout().is_terminal(); let ret = if is_terminal_stdout { @@ -183,11 +211,11 @@ async fn execute(config: &GlobalConfig, mut input: Input) -> Result<()> { } else { client.send_message(input.clone()).await }; - let (mut eval_str, _) = ret?; + let mut eval_str = ret?.text; if let Ok(true) = CODE_BLOCK_RE.is_match(&eval_str) { eval_str = extract_block(&eval_str); } - config.write().save_message(input.clone(), &eval_str)?; + config.write().save_message(&input, &eval_str, &[])?; config.read().maybe_copy(&eval_str); let render_options = config.read().get_render_options()?; let mut markdown_render = MarkdownRender::init(render_options)?; @@ -205,7 +233,8 @@ async fn execute(config: &GlobalConfig, mut input: Input) -> Result<()> { match answer { "✅ Execute" => { - let code = run_command(&eval_str)?; + debug!("{} {:?}", shell, &[shell_arg, &eval_str]); + let code = run_command(shell, &[shell_arg, &eval_str], None)?; if code != 0 { process::exit(code); } @@ -214,7 +243,7 @@ async fn execute(config: &GlobalConfig, mut input: Input) -> Result<()> { let revision = Text::new("Enter your revision:").prompt()?; let text = format!("{}\n{revision}", input.text()); input.set_text(text); - return execute(config, input).await; + return shell_execute(config, shell, shell_arg, input).await; } "📙 Explain" => { let role = config.read().retrieve_role(EXPLAIN_SHELL_ROLE)?; -- cgit v1.2.3