mod cli; mod client; mod config; mod function; mod rag; mod render; mod repl; mod serve; #[macro_use] mod utils; #[macro_use] extern crate log; use crate::cli::Cli; use crate::client::{ call_chat_completions, call_chat_completions_streaming, list_chat_models, ChatCompletionsOutput, }; use crate::config::{ ensure_parent_exists, list_agents, load_env_file, Config, GlobalConfig, Input, WorkingMode, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE, TEMP_SESSION_NAME, }; use crate::function::{eval_tool_calls, need_send_tool_results}; use crate::render::render_error; use crate::repl::Repl; use crate::utils::{ create_abort_signal, create_spinner, detect_shell, extract_block, get_env_name, run_command, run_with_spinner, AbortSignal, Shell, CODE_BLOCK_RE, IS_STDOUT_TERMINAL, }; use anyhow::{bail, Result}; use async_recursion::async_recursion; use clap::Parser; use inquire::{Select, Text}; use is_terminal::IsTerminal; use parking_lot::RwLock; use simplelog::{format_description, ConfigBuilder, LevelFilter, SimpleLogger, WriteLogger}; use std::{ env, io::{stderr, stdin, Read}, process, sync::Arc, }; #[tokio::main] async fn main() -> Result<()> { load_env_file()?; let cli = Cli::parse(); let text = cli.text(); let text = aggregate_text(text)?; let working_mode = if cli.serve.is_some() { WorkingMode::Serve } else if text.is_none() && cli.file.is_empty() { WorkingMode::Repl } else { WorkingMode::Command }; let mut model_id = cli.model.clone(); if model_id.is_none() { if let Ok(v) = env::var(get_env_name("model")) { model_id = Some(v); } } setup_logger(working_mode.is_serve())?; let config = Arc::new(RwLock::new(Config::init( working_mode, model_id.as_deref(), )?)); let highlight = config.read().highlight; if let Err(err) = run(config, cli, text, model_id).await { let highlight = stderr().is_terminal() && highlight; render_error(err, highlight); std::process::exit(1); } Ok(()) } async fn run( config: GlobalConfig, cli: Cli, text: Option, model_id: Option, ) -> Result<()> { let abort_signal = create_abort_signal(); if let Some(addr) = cli.serve { return serve::run(config, addr).await; } if cli.list_models { for model in list_chat_models(&config.read()) { println!("{}", model.id()); } return Ok(()); } if cli.list_roles { config .read() .roles .iter() .for_each(|v| println!("{}", v.name())); return Ok(()); } if cli.list_agents { let agents = list_agents().join("\n"); println!("{agents}"); return Ok(()); } if cli.list_rags { let rags = config.read().list_rags().join("\n"); println!("{rags}"); return Ok(()); } if cli.dry_run { config.write().dry_run = true; } if let Some(agent) = &cli.agent { let session = cli.session.as_ref().map(|v| match v { Some(v) => v.as_str(), None => TEMP_SESSION_NAME, }); Config::use_agent(&config, agent, session, abort_signal.clone()).await? } else { if let Some(prompt) = &cli.prompt { config.write().use_prompt(prompt)?; } else if let Some(name) = &cli.role { config.write().use_role(name)?; } else if cli.execute { config.write().use_role(SHELL_ROLE)?; } else if cli.code { config.write().use_role(CODE_ROLE)?; } if let Some(session) = &cli.session { config .write() .use_session(session.as_ref().map(|v| v.as_str()))?; } if let Some(rag) = &cli.rag { Config::use_rag(&config, Some(rag), abort_signal.clone()).await?; } } if cli.list_sessions { let sessions = config.read().list_sessions().join("\n"); println!("{sessions}"); return Ok(()); } if let Some(model_id) = &model_id { config.write().set_model(model_id)?; } if cli.no_stream { config.write().stream = false; } if cli.save_session { config.write().set_save_session(Some(true)); } if cli.info { let info = config.read().info()?; println!("{}", info); return Ok(()); } let is_repl = config.read().working_mode.is_repl(); if cli.execute && !is_repl { let input = create_input(&config, text, &cli.file).await?; let shell = detect_shell(); shell_execute(&config, &shell, input).await?; return Ok(()); } config.write().apply_prelude()?; match is_repl { false => { let mut input = create_input(&config, text, &cli.file).await?; input.use_embeddings(abort_signal.clone()).await?; start_directive(&config, input, cli.code, abort_signal).await } true => start_interactive(&config).await, } } #[async_recursion] async fn start_directive( config: &GlobalConfig, input: Input, code_mode: bool, abort_signal: AbortSignal, ) -> Result<()> { let client = input.create_client()?; let extract_code = !*IS_STDOUT_TERMINAL && code_mode; config.write().before_chat_completion(&input)?; let (output, tool_results) = if !config.read().stream || extract_code { let task = client.chat_completions(input.clone()); let ret = run_with_spinner(task, "Generating").await; match ret { Ok(ret) => { let ChatCompletionsOutput { mut text, tool_calls, .. } = ret; if !text.is_empty() { if extract_code && text.trim_start().starts_with("```") { text = extract_block(&text); } config.read().print_markdown(&text)?; } (text, eval_tool_calls(config, tool_calls)?) } Err(err) => return Err(err), } } else { call_chat_completions_streaming(&input, client.as_ref(), config, abort_signal.clone()) .await? }; config .write() .after_chat_completion(&input, &output, &tool_results)?; config.write().exit_session()?; if need_send_tool_results(&tool_results) { start_directive( config, input.merge_tool_call(output, tool_results), code_mode, abort_signal, ) .await } else { Ok(()) } } async fn start_interactive(config: &GlobalConfig) -> Result<()> { let mut repl: Repl = Repl::init(config)?; repl.run().await } #[async_recursion::async_recursion] async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) -> Result<()> { let client = input.create_client()?; config.write().before_chat_completion(&input)?; let ret = if *IS_STDOUT_TERMINAL { let spinner = create_spinner("Generating").await; let ret = client.chat_completions(input.clone()).await; spinner.stop(); ret } else { client.chat_completions(input.clone()).await }; 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() .after_chat_completion(&input, &eval_str, &[])?; if config.read().dry_run { config.read().print_markdown(&eval_str)?; return Ok(()); } if *IS_STDOUT_TERMINAL { loop { let answer = Select::new( eval_str.trim(), vec!["✅ Execute", "🔄️ Revise", "📖 Explain", "❌ Cancel"], ) .prompt()?; match answer { "✅ Execute" => { debug!("{} {:?}", shell.cmd, &[&shell.arg, &eval_str]); let code = run_command(&shell.cmd, &[&shell.arg, &eval_str], None)?; if code != 0 { process::exit(code); } } "🔄️ Revise" => { let revision = Text::new("Enter your revision:").prompt()?; let text = format!("{}\n{revision}", input.text()); input.set_text(text); return shell_execute(config, shell, input).await; } "📖 Explain" => { let role = config.read().retrieve_role(EXPLAIN_SHELL_ROLE)?; let input = Input::from_str(config, &eval_str, Some(role)); let abort = create_abort_signal(); if config.read().stream { call_chat_completions_streaming(&input, client.as_ref(), config, abort) .await?; } else { call_chat_completions(&input, client.as_ref(), config).await?; } continue; } _ => {} } break; } } else { println!("{}", eval_str); } Ok(()) } fn aggregate_text(text: Option) -> Result> { let text = if stdin().is_terminal() { text } else { let mut stdin_text = String::new(); stdin().read_to_string(&mut stdin_text)?; if let Some(text) = text { Some(format!("{text}\n{stdin_text}")) } else { Some(stdin_text) } }; Ok(text) } async fn create_input( config: &GlobalConfig, text: Option, file: &[String], ) -> Result { let input = if file.is_empty() { Input::from_str(config, &text.unwrap_or_default(), None) } else { Input::from_files(config, &text.unwrap_or_default(), file.to_vec(), None).await? }; if input.is_empty() { bail!("No input"); } Ok(input) } fn setup_logger(is_serve: bool) -> Result<()> { let (log_level, log_path) = Config::log(is_serve)?; if log_level == LevelFilter::Off { return Ok(()); } let crate_name = env!("CARGO_CRATE_NAME"); let log_filter = match std::env::var(get_env_name("log_filter")) { Ok(v) => v, Err(_) => match is_serve { true => format!("{crate_name}::serve"), false => crate_name.into(), }, }; let config = ConfigBuilder::new() .add_filter_allow(log_filter) .set_time_format_custom(format_description!( "[year]-[month]-[day]T[hour]:[minute]:[second].[subsecond digits:3]Z" )) .set_thread_level(LevelFilter::Off) .build(); match log_path { None => { SimpleLogger::init(log_level, config)?; } Some(log_path) => { ensure_parent_exists(&log_path)?; let log_file = std::fs::File::create(log_path)?; WriteLogger::init(log_level, config, log_file)?; } } Ok(()) }