From 360264121ca2931cb7f2fa40594ba590b245ace5 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 7 Mar 2023 22:33:34 +0800 Subject: feat: command mode supports stream out (#31) * feat: command mode supports stream out * update cli --- src/render/cmd.rs | 37 +++++++++++++ src/render/mod.rs | 153 +++++++++++++---------------------------------------- src/render/repl.rs | 123 ++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 196 insertions(+), 117 deletions(-) create mode 100644 src/render/cmd.rs create mode 100644 src/render/repl.rs (limited to 'src/render') diff --git a/src/render/cmd.rs b/src/render/cmd.rs new file mode 100644 index 0000000..f29fe7d --- /dev/null +++ b/src/render/cmd.rs @@ -0,0 +1,37 @@ +use super::MarkdownRender; +use crate::repl::{ReplyStreamEvent, SharedAbortSignal}; +use crate::utils::dump; + +use anyhow::Result; +use crossbeam::channel::Receiver; + +pub fn cmd_render_stream(rx: Receiver, abort: SharedAbortSignal) -> Result<()> { + let mut buffer = String::new(); + let mut markdown_render = MarkdownRender::new(); + loop { + if abort.aborted() { + return Ok(()); + } + if let Ok(evt) = rx.try_recv() { + match evt { + ReplyStreamEvent::Text(text) => { + if text.contains('\n') { + let text = format!("{buffer}{text}"); + let mut lines: Vec<&str> = text.split('\n').collect(); + buffer = lines.pop().unwrap_or_default().to_string(); + let output = lines.join("\n"); + dump(markdown_render.render(&output), 1); + } else { + buffer = format!("{buffer}{text}"); + } + } + ReplyStreamEvent::Done => { + let output = markdown_render.render(&buffer); + dump(output, 2); + break; + } + } + } + } + Ok(()) +} diff --git a/src/render/mod.rs b/src/render/mod.rs index e0f42e7..efbcdd7 100644 --- a/src/render/mod.rs +++ b/src/render/mod.rs @@ -1,125 +1,44 @@ +mod cmd; mod markdown; +mod repl; +use self::cmd::cmd_render_stream; pub use self::markdown::MarkdownRender; -use crate::repl::{ReplyStreamEvent, SharedAbortSignal}; +use self::repl::repl_render_stream; +use crate::client::ChatGptClient; +use crate::repl::{ReplyStreamHandler, SharedAbortSignal}; use anyhow::Result; -use crossbeam::channel::Receiver; -use crossterm::{ - cursor, - event::{self, Event, KeyCode, KeyModifiers}, - queue, style, - terminal::{self, disable_raw_mode, enable_raw_mode}, -}; -use std::{ - io::{self, Stdout, Write}, - time::{Duration, Instant}, -}; -use unicode_width::UnicodeWidthStr; - -pub fn render_stream(rx: Receiver, abort: SharedAbortSignal) -> Result<()> { - enable_raw_mode()?; - let mut stdout = io::stdout(); - queue!(stdout, event::DisableMouseCapture)?; - - let ret = render_stream_inner(rx, abort, &mut stdout); - - queue!(stdout, event::DisableMouseCapture)?; - disable_raw_mode()?; - - ret -} - -pub fn render_stream_inner( - rx: Receiver, +use crossbeam::channel::unbounded; +use crossbeam::sync::WaitGroup; +use std::thread::spawn; + +pub fn render_stream( + input: &str, + prompt: Option, + client: &ChatGptClient, + highlight: bool, + repl: bool, abort: SharedAbortSignal, - writer: &mut Stdout, -) -> Result<()> { - let mut last_tick = Instant::now(); - let tick_rate = Duration::from_millis(100); - let mut buffer = String::new(); - let mut markdown_render = MarkdownRender::new(); - let terminal_columns = terminal::size()?.0; - loop { - if abort.aborted() { - return Ok(()); - } - - if let Ok(evt) = rx.try_recv() { - recover_cursor(writer, terminal_columns, &buffer)?; - - match evt { - ReplyStreamEvent::Text(text) => { - if text.contains('\n') { - let text = format!("{buffer}{text}"); - let mut lines: Vec<&str> = text.split('\n').collect(); - buffer = lines.pop().unwrap_or_default().to_string(); - let output = markdown_render.render(&lines.join("\n")); - for line in output.split('\n') { - queue!( - writer, - style::Print(line), - style::Print("\n"), - cursor::MoveLeft(terminal_columns), - )?; - } - queue!(writer, style::Print(&buffer),)?; - } else { - buffer = format!("{buffer}{text}"); - let output = markdown_render.render_line_stateless(&buffer); - queue!(writer, style::Print(&output))?; - } - writer.flush()?; - } - ReplyStreamEvent::Done => { - let output = markdown_render.render_line_stateless(&buffer); - queue!(writer, style::Print(output), style::Print("\n"))?; - writer.flush()?; - break; - } - } - continue; - } - - let timeout = tick_rate - .checked_sub(last_tick.elapsed()) - .unwrap_or_else(|| Duration::from_secs(0)); - if crossterm::event::poll(timeout)? { - if let Event::Key(key) = event::read()? { - match key.code { - KeyCode::Char('c') if key.modifiers == KeyModifiers::CONTROL => { - abort.set_ctrlc(); - return Ok(()); - } - KeyCode::Char('d') if key.modifiers == KeyModifiers::CONTROL => { - abort.set_ctrld(); - return Ok(()); - } - _ => {} - } - } - } - - if last_tick.elapsed() >= tick_rate { - last_tick = Instant::now(); - } - } - Ok(()) -} - -fn recover_cursor(writer: &mut Stdout, terminal_columns: u16, buffer: &str) -> Result<()> { - let buffer_rows = (buffer.width() as u16 + terminal_columns - 1) / terminal_columns; - let (_, row) = cursor::position()?; - if buffer_rows == 0 { - queue!(writer, cursor::MoveTo(0, row))?; - } else if row + 1 >= buffer_rows { - queue!(writer, cursor::MoveTo(0, row + 1 - buffer_rows))?; + wg: WaitGroup, +) -> Result { + let mut stream_handler = if highlight { + let (tx, rx) = unbounded(); + let abort_clone = abort.clone(); + spawn(move || { + let _ = if repl { + repl_render_stream(rx, abort) + } else { + cmd_render_stream(rx, abort) + }; + drop(wg); + }); + ReplyStreamHandler::new(Some(tx), abort_clone) } else { - queue!( - writer, - terminal::ScrollUp(buffer_rows - 1 - row), - cursor::MoveTo(0, 0) - )?; - } - Ok(()) + drop(wg); + ReplyStreamHandler::new(None, abort) + }; + client.send_message_streaming(input, prompt, &mut stream_handler)?; + let buffer = stream_handler.get_buffer(); + Ok(buffer.to_string()) } diff --git a/src/render/repl.rs b/src/render/repl.rs new file mode 100644 index 0000000..96c6364 --- /dev/null +++ b/src/render/repl.rs @@ -0,0 +1,123 @@ +use super::MarkdownRender; +use crate::repl::{ReplyStreamEvent, SharedAbortSignal}; + +use anyhow::Result; +use crossbeam::channel::Receiver; +use crossterm::{ + cursor, + event::{self, Event, KeyCode, KeyModifiers}, + queue, style, + terminal::{self, disable_raw_mode, enable_raw_mode}, +}; +use std::{ + io::{self, Stdout, Write}, + time::{Duration, Instant}, +}; +use unicode_width::UnicodeWidthStr; + +pub fn repl_render_stream(rx: Receiver, abort: SharedAbortSignal) -> Result<()> { + enable_raw_mode()?; + let mut stdout = io::stdout(); + queue!(stdout, event::DisableMouseCapture)?; + + let ret = repl_render_stream_inner(rx, abort, &mut stdout); + + queue!(stdout, event::DisableMouseCapture)?; + disable_raw_mode()?; + + ret +} + +fn repl_render_stream_inner( + rx: Receiver, + abort: SharedAbortSignal, + writer: &mut Stdout, +) -> Result<()> { + let mut last_tick = Instant::now(); + let tick_rate = Duration::from_millis(100); + let mut buffer = String::new(); + let mut markdown_render = MarkdownRender::new(); + let terminal_columns = terminal::size()?.0; + loop { + if abort.aborted() { + return Ok(()); + } + + if let Ok(evt) = rx.try_recv() { + recover_cursor(writer, terminal_columns, &buffer)?; + + match evt { + ReplyStreamEvent::Text(text) => { + if text.contains('\n') { + let text = format!("{buffer}{text}"); + let mut lines: Vec<&str> = text.split('\n').collect(); + buffer = lines.pop().unwrap_or_default().to_string(); + let output = markdown_render.render(&lines.join("\n")); + for line in output.split('\n') { + queue!( + writer, + style::Print(line), + style::Print("\n"), + cursor::MoveLeft(terminal_columns), + )?; + } + queue!(writer, style::Print(&buffer),)?; + } else { + buffer = format!("{buffer}{text}"); + let output = markdown_render.render_line_stateless(&buffer); + queue!(writer, style::Print(&output))?; + } + writer.flush()?; + } + ReplyStreamEvent::Done => { + let output = markdown_render.render_line_stateless(&buffer); + queue!(writer, style::Print(output), style::Print("\n"))?; + writer.flush()?; + break; + } + } + continue; + } + + let timeout = tick_rate + .checked_sub(last_tick.elapsed()) + .unwrap_or_else(|| Duration::from_secs(0)); + if crossterm::event::poll(timeout)? { + if let Event::Key(key) = event::read()? { + match key.code { + KeyCode::Char('c') if key.modifiers == KeyModifiers::CONTROL => { + abort.set_ctrlc(); + return Ok(()); + } + KeyCode::Char('d') if key.modifiers == KeyModifiers::CONTROL => { + abort.set_ctrld(); + return Ok(()); + } + _ => {} + } + } + } + + if last_tick.elapsed() >= tick_rate { + last_tick = Instant::now(); + } + } + Ok(()) +} + +fn recover_cursor(writer: &mut Stdout, terminal_columns: u16, buffer: &str) -> Result<()> { + let buffer_rows = (buffer.width() as u16 + terminal_columns - 1) / terminal_columns; + let (_, row) = cursor::position()?; + if buffer_rows == 0 { + queue!(writer, cursor::MoveTo(0, row))?; + } else if row + 1 >= buffer_rows { + queue!(writer, cursor::MoveTo(0, row + 1 - buffer_rows))?; + } else { + queue!( + writer, + terminal::ScrollUp(buffer_rows - 1 - row), + cursor::MoveTo(0, 0) + )?; + } + Ok(()) +} -- cgit v1.2.3