summaryrefslogtreecommitdiffstats
path: root/src/main.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/main.rs')
-rw-r--r--src/main.rs79
1 files changed, 54 insertions, 25 deletions
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)?;