summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--config.example.yaml1
-rw-r--r--src/client/common.rs35
-rw-r--r--src/config/mod.rs56
-rw-r--r--src/main.rs60
-rw-r--r--src/repl/mod.rs13
-rw-r--r--src/utils/mod.rs2
-rw-r--r--src/utils/spinner.rs24
7 files changed, 130 insertions, 61 deletions
diff --git a/config.example.yaml b/config.example.yaml
index 29966a0..3f72dc6 100644
--- a/config.example.yaml
+++ b/config.example.yaml
@@ -4,6 +4,7 @@ temperature: null # Set default temperature parameter
top_p: null # Set default top-p parameter, range (0, 1)
# ---- behavior ----
+stream: true # Use stream style by default
save: true # Indicates whether to persist the message
keybindings: emacs # Choose keybinding style (emacs, vi)
buffer_editor: null # Command used to edit the current input with ctrl+o, env: EDITOR
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<Option<(St
}
}
-pub async fn chat_completion_streaming(
+pub async fn call_chat_completions(
+ input: &Input,
+ client: &dyn Client,
+ config: &GlobalConfig,
+) -> Result<(String, Vec<ToolResult>)> {
+ 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<f64>,
pub dry_run: bool,
+ pub stream: bool,
pub save: bool,
pub keybindings: String,
pub buffer_editor: Option<String>,
@@ -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<String> {
- 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<SpinnerEve
}
Ok(())
}
+
+pub async fn run_with_spinner<F, T>(task: F, spinner_message: &str) -> Result<T>
+where
+ F: Future<Output = Result<T>>,
+{
+ if *IS_STDOUT_TERMINAL {
+ let spinner = create_spinner(spinner_message).await;
+ let ret = task.await;
+ spinner.stop();
+ ret
+ } else {
+ task.await
+ }
+}