From 49b61129c95a3528eaf25dabcb55825b5ed7be72 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sun, 28 Jul 2024 08:56:00 +0800 Subject: feat: add `config.stream` and `.set stream` repl command (#759) --- src/client/common.rs | 35 ++++++++++++++++++++++++------ src/config/mod.rs | 56 ++++++++++++++++++++++++++++++++---------------- src/main.rs | 60 ++++++++++++++++++++++++++++++---------------------- src/repl/mod.rs | 13 +++++++----- src/utils/mod.rs | 2 +- src/utils/spinner.rs | 24 ++++++++++++++++----- 6 files changed, 129 insertions(+), 61 deletions(-) (limited to 'src') diff --git a/src/client/common.rs b/src/client/common.rs index 2108941..ec1f37d 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -397,7 +397,28 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result Result<(String, Vec)> { + let task = client.chat_completions(input.clone()); + let ret = run_with_spinner(task, "Generating").await; + match ret { + Ok(ret) => { + let ChatCompletionsOutput { + text, tool_calls, .. + } = ret; + if !text.is_empty() { + config.read().print_markdown(&text)?; + } + Ok((text, eval_tool_calls(config, tool_calls)?)) + } + Err(err) => Err(err), + } +} + +pub async fn call_chat_completions_streaming( input: &Input, client: &dyn Client, config: &GlobalConfig, @@ -406,23 +427,23 @@ pub async fn chat_completion_streaming( let (tx, rx) = unbounded_channel(); let mut handler = SseHandler::new(tx, abort.clone()); - let (send_ret, rend_ret) = tokio::join!( + let (send_ret, render_ret) = tokio::join!( client.chat_completions_streaming(input, &mut handler), render_stream(rx, config, abort.clone()), ); - if let Err(err) = rend_ret { + if let Err(err) = render_ret { render_error(err, config.read().highlight); } - let (output, calls) = handler.take(); + let (text, tool_calls) = handler.take(); match send_ret { Ok(_) => { - if !output.is_empty() && !output.ends_with('\n') { + if !text.is_empty() && !text.ends_with('\n') { println!(); } - Ok((output, eval_tool_calls(config, calls)?)) + Ok((text, eval_tool_calls(config, tool_calls)?)) } Err(err) => { - if !output.is_empty() { + if !text.is_empty() { println!(); } Err(err) diff --git a/src/config/mod.rs b/src/config/mod.rs index 6d196d6..9ebf34b 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -89,6 +89,7 @@ pub struct Config { pub top_p: Option, pub dry_run: bool, + pub stream: bool, pub save: bool, pub keybindings: String, pub buffer_editor: Option, @@ -156,6 +157,7 @@ impl Default for Config { top_p: None, dry_run: false, + stream: true, save: false, keybindings: "emacs".into(), buffer_editor: None, @@ -516,6 +518,7 @@ impl Config { ("temperature", format_option_value(&role.temperature())), ("top_p", format_option_value(&role.top_p())), ("dry_run", self.dry_run.to_string()), + ("stream", self.stream.to_string()), ("save", self.save.to_string()), ("keybindings", self.keybindings.clone()), ("wrap", wrap), @@ -570,6 +573,18 @@ impl Config { let value = parse_value(value)?; self.set_top_p(value); } + "dry_run" => { + let value = value.parse().with_context(|| "Invalid value")?; + self.dry_run = value; + } + "stream" => { + let value = value.parse().with_context(|| "Invalid value")?; + self.stream = value; + } + "save" => { + let value = value.parse().with_context(|| "Invalid value")?; + self.save = value; + } "rag_reranker_model" => { self.rag_reranker_model = if value == "null" { None @@ -593,26 +608,18 @@ impl Config { let value = parse_value(value)?; self.set_use_tools(value); } - "compress_threshold" => { - let value = parse_value(value)?; - self.set_compress_threshold(value); - } - "save" => { - let value = value.parse().with_context(|| "Invalid value")?; - self.save = value; - } "save_session" => { let value = parse_value(value)?; self.set_save_session(value); } + "compress_threshold" => { + let value = parse_value(value)?; + self.set_compress_threshold(value); + } "highlight" => { let value = value.parse().with_context(|| "Invalid value")?; self.highlight = value; } - "dry_run" => { - let value = value.parse().with_context(|| "Invalid value")?; - self.dry_run = value; - } _ => bail!("Unknown key `{key}`"), } Ok(()) @@ -1229,6 +1236,7 @@ impl Config { "temperature", "top_p", "dry_run", + "stream", "save", "save_session", "compress_threshold", @@ -1251,6 +1259,7 @@ impl Config { None => vec![], }, "dry_run" => complete_bool(self.dry_run), + "stream" => complete_bool(self.stream), "save" => complete_bool(self.save), "save_session" => { let save_session = if let Some(session) = &self.session { @@ -1338,12 +1347,6 @@ impl Config { Ok(RenderOptions::new(theme, wrap, self.wrap_code, truecolor)) } - pub fn markdown_render(&self, text: &str) -> Result { - let render_options = self.render_options()?; - let mut markdown_render = MarkdownRender::init(render_options)?; - Ok(markdown_render.render(text)) - } - pub fn render_prompt_left(&self) -> String { let variables = self.generate_prompt_context(); let left_prompt = self.left_prompt.as_deref().unwrap_or(LEFT_PROMPT); @@ -1356,6 +1359,17 @@ impl Config { render_prompt(right_prompt, &variables) } + pub fn print_markdown(&self, text: &str) -> Result<()> { + if *IS_STDOUT_TERMINAL { + let render_options = self.render_options()?; + let mut markdown_render = MarkdownRender::init(render_options)?; + println!("{}", markdown_render.render(text)); + } else { + println!("{text}"); + } + Ok(()) + } + fn generate_prompt_context(&self) -> HashMap<&str, String> { let mut output = HashMap::new(); let role = self.extract_role(); @@ -1382,6 +1396,9 @@ impl Config { if self.dry_run { output.insert("dry_run", "true".to_string()); } + if self.stream { + output.insert("stream", "true".to_string()); + } if self.save { output.insert("save", "true".to_string()); } @@ -1557,6 +1574,9 @@ impl Config { if let Some(Some(v)) = read_env_bool("dry_run") { self.dry_run = v; } + if let Some(Some(v)) = read_env_bool("stream") { + self.stream = v; + } if let Some(Some(v)) = read_env_bool("save") { self.save = v; } diff --git a/src/main.rs b/src/main.rs index 172c8e3..04915ed 100644 --- a/src/main.rs +++ b/src/main.rs @@ -13,7 +13,9 @@ mod utils; extern crate log; use crate::cli::Cli; -use crate::client::{chat_completion_streaming, list_chat_models, ChatCompletionsOutput}; +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, @@ -23,7 +25,7 @@ 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, - AbortSignal, Shell, CODE_BLOCK_RE, IS_STDOUT_TERMINAL, + run_with_spinner, AbortSignal, Shell, CODE_BLOCK_RE, IS_STDOUT_TERMINAL, }; use anyhow::{bail, Result}; @@ -145,6 +147,9 @@ async fn run( 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)); } @@ -165,7 +170,7 @@ async fn run( false => { let mut input = create_input(&config, text, &cli.file).await?; input.use_embeddings(abort_signal.clone()).await?; - start_directive(&config, input, cli.no_stream, cli.code, abort_signal).await + start_directive(&config, input, cli.code, abort_signal).await } true => start_interactive(&config).await, } @@ -175,34 +180,35 @@ async fn run( async fn start_directive( config: &GlobalConfig, input: Input, - no_stream: bool, 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 no_stream || extract_code { - let ChatCompletionsOutput { - text, tool_calls, .. - } = client.chat_completions(input.clone()).await?; - if !tool_calls.is_empty() { - (String::new(), eval_tool_calls(config, tool_calls)?) - } else { - let text = if extract_code && text.trim_start().starts_with("```") { - extract_block(&text) - } else { - text.clone() - }; - if *IS_STDOUT_TERMINAL { - println!("{}", config.read().markdown_render(&text)?); - } else { - println!("{}", text); + 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)?) } - (text, vec![]) + Err(err) => return Err(err), } } else { - chat_completion_streaming(&input, client.as_ref(), config, abort_signal.clone()).await? + call_chat_completions_streaming(&input, client.as_ref(), config, abort_signal.clone()) + .await? }; config .write() @@ -214,7 +220,6 @@ async fn start_directive( start_directive( config, input.merge_tool_call(output, tool_results), - no_stream, code_mode, abort_signal, ) @@ -249,7 +254,7 @@ async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) - .write() .after_chat_completion(&input, &eval_str, &[])?; if config.read().dry_run { - println!("{}", config.read().markdown_render(&eval_str)?); + config.read().print_markdown(&eval_str)?; return Ok(()); } if *IS_STDOUT_TERMINAL { @@ -278,7 +283,12 @@ async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) - let role = config.read().retrieve_role(EXPLAIN_SHELL_ROLE)?; let input = Input::from_str(config, &eval_str, Some(role)); let abort = create_abort_signal(); - chat_completion_streaming(&input, client.as_ref(), config, abort).await?; + 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; } _ => {} diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 3cfd7de..0711a6d 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -6,7 +6,7 @@ use self::completer::ReplCompleter; use self::highlighter::ReplHighlighter; use self::prompt::ReplPrompt; -use crate::client::chat_completion_streaming; +use crate::client::{call_chat_completions, call_chat_completions_streaming}; use crate::config::{AssertState, Config, GlobalConfig, Input, StateFlags}; use crate::function::need_send_tool_results; use crate::render::render_error; @@ -286,8 +286,7 @@ impl Repl { } None => { let banner = self.config.read().agent_banner()?; - let output = self.config.read().markdown_render(&banner)?; - println!("{output}"); + self.config.read().print_markdown(&banner)?; } }, ".variable" => match args { @@ -569,8 +568,12 @@ async fn ask( let client = input.create_client()?; config.write().before_chat_completion(&input)?; - let (output, tool_results) = - chat_completion_streaming(&input, client.as_ref(), config, abort_signal.clone()).await?; + let (output, tool_results) = if config.read().stream { + call_chat_completions_streaming(&input, client.as_ref(), config, abort_signal.clone()) + .await? + } else { + call_chat_completions(&input, client.as_ref(), config).await? + }; config .write() .after_chat_completion(&input, &output, &tool_results)?; diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 340359c..4e5428a 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -16,7 +16,7 @@ pub use self::path::*; pub use self::prompt_input::*; pub use self::render_prompt::render_prompt; pub use self::request::*; -pub use self::spinner::{create_spinner, Spinner}; +pub use self::spinner::*; use anyhow::{Context, Result}; use fancy_regex::Regex; diff --git a/src/utils/spinner.rs b/src/utils/spinner.rs index 8f386db..53969f4 100644 --- a/src/utils/spinner.rs +++ b/src/utils/spinner.rs @@ -1,7 +1,9 @@ +use super::IS_STDOUT_TERMINAL; + use anyhow::Result; use crossterm::{cursor, queue, style, terminal}; -use is_terminal::IsTerminal; use std::{ + future::Future, io::{stdout, Write}, time::Duration, }; @@ -10,7 +12,6 @@ use tokio::{sync::mpsc, time::interval}; pub struct SpinnerInner { index: usize, message: String, - is_not_terminal: bool, } impl SpinnerInner { @@ -20,12 +21,11 @@ impl SpinnerInner { SpinnerInner { index: 0, message: message.to_string(), - is_not_terminal: !stdout().is_terminal(), } } fn step(&mut self) -> Result<()> { - if self.is_not_terminal || self.message.is_empty() { + if !*IS_STDOUT_TERMINAL || self.message.is_empty() { return Ok(()); } let mut writer = stdout(); @@ -50,7 +50,7 @@ impl SpinnerInner { } fn clear_message(&mut self) -> Result<()> { - if self.is_not_terminal || self.message.is_empty() { + if !*IS_STDOUT_TERMINAL || self.message.is_empty() { return Ok(()); } self.message.clear(); @@ -126,3 +126,17 @@ async fn run_spinner(message: String, mut rx: mpsc::UnboundedReceiver(task: F, spinner_message: &str) -> Result +where + F: Future>, +{ + if *IS_STDOUT_TERMINAL { + let spinner = create_spinner(spinner_message).await; + let ret = task.await; + spinner.stop(); + ret + } else { + task.await + } +} -- cgit v1.2.3