summaryrefslogtreecommitdiffstats
path: root/src/render
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-03-07 22:33:34 +0800
committerGitHub <noreply@github.com>2023-03-07 22:33:34 +0800
commit360264121ca2931cb7f2fa40594ba590b245ace5 (patch)
tree71e81208cf50b02116c89974aca17a3cc762b1fc /src/render
parentc7fcdb174440817706b818154a6bbcaf58948caa (diff)
downloadaichat-360264121ca2931cb7f2fa40594ba590b245ace5.tar.gz
feat: command mode supports stream out (#31)
* feat: command mode supports stream out * update cli
Diffstat (limited to 'src/render')
-rw-r--r--src/render/cmd.rs37
-rw-r--r--src/render/mod.rs153
-rw-r--r--src/render/repl.rs123
3 files changed, 196 insertions, 117 deletions
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<ReplyStreamEvent>, 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<ReplyStreamEvent>, 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<ReplyStreamEvent>,
+use crossbeam::channel::unbounded;
+use crossbeam::sync::WaitGroup;
+use std::thread::spawn;
+
+pub fn render_stream(
+ input: &str,
+ prompt: Option<String>,
+ 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<String> {
+ 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<ReplyStreamEvent>, 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<ReplyStreamEvent>,
+ 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(())
+}