From 1cc89eff514669d6459f62cc4eed8e0b34d8ef0c Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 23 Apr 2024 14:32:06 +0800 Subject: refactor: more async code (#427) --- src/render/mod.rs | 119 ++++++--------------------------------------------- src/render/stream.rs | 72 +++++++++++++++---------------- 2 files changed, 47 insertions(+), 144 deletions(-) (limited to 'src/render') diff --git a/src/render/mod.rs b/src/render/mod.rs index dc4c081..146577b 100644 --- a/src/render/mod.rs +++ b/src/render/mod.rs @@ -4,61 +4,26 @@ mod stream; pub use self::markdown::{MarkdownRender, RenderOptions}; use self::stream::{markdown_stream, raw_stream}; -use crate::client::Client; -use crate::config::{GlobalConfig, Input}; use crate::utils::AbortSignal; +use crate::{client::ReplyEvent, config::GlobalConfig}; -use anyhow::{Context, Result}; -use crossbeam::channel::{unbounded, Sender}; -use crossbeam::sync::WaitGroup; +use anyhow::Result; use is_terminal::IsTerminal; use nu_ansi_term::{Color, Style}; use std::io::stdout; -use std::thread::spawn; +use tokio::sync::mpsc::UnboundedReceiver; -pub fn render_stream( - input: &Input, - client: &dyn Client, +pub async fn render_stream( + rx: UnboundedReceiver, config: &GlobalConfig, abort: AbortSignal, -) -> Result { - let wg = WaitGroup::new(); - let wg_cloned = wg.clone(); - let render_options = config.read().get_render_options()?; - let mut stream_handler = { - let (tx, rx) = unbounded(); - let abort_clone = abort.clone(); - let highlight = config.read().highlight; - spawn(move || { - let run = move || { - if stdout().is_terminal() { - let mut render = MarkdownRender::init(render_options)?; - markdown_stream(&rx, &mut render, &abort) - } else { - raw_stream(&rx, &abort) - } - }; - if let Err(err) = run() { - render_error(err, highlight); - } - drop(wg_cloned); - }); - ReplyHandler::new(tx, abort_clone) - }; - let ret = client.send_message_streaming(input, &mut stream_handler); - wg.wait(); - let output = stream_handler.get_buffer().to_string(); - match ret { - Ok(_) => { - println!(); - Ok(output) - } - Err(err) => { - if !output.is_empty() { - println!(); - } - Err(err) - } +) -> Result<()> { + if stdout().is_terminal() { + let render_options = config.read().get_render_options()?; + let mut render = MarkdownRender::init(render_options)?; + markdown_stream(rx, &mut render, &abort).await + } else { + raw_stream(rx, &abort).await } } @@ -71,63 +36,3 @@ pub fn render_error(err: anyhow::Error, highlight: bool) { eprintln!("{err}"); } } - -pub struct ReplyHandler { - sender: Sender, - buffer: String, - abort: AbortSignal, -} - -impl ReplyHandler { - pub fn new(sender: Sender, abort: AbortSignal) -> Self { - Self { - sender, - abort, - buffer: String::new(), - } - } - - pub fn text(&mut self, text: &str) -> Result<()> { - debug!("ReplyText: {}", text); - if text.is_empty() { - return Ok(()); - } - self.buffer.push_str(text); - let ret = self - .sender - .send(ReplyEvent::Text(text.to_string())) - .with_context(|| "Failed to send ReplyEvent:Text"); - self.safe_ret(ret)?; - Ok(()) - } - - pub fn done(&mut self) -> Result<()> { - debug!("ReplyDone"); - let ret = self - .sender - .send(ReplyEvent::Done) - .with_context(|| "Failed to send ReplyEvent::Done"); - self.safe_ret(ret)?; - Ok(()) - } - - pub fn get_buffer(&self) -> &str { - &self.buffer - } - - pub fn get_abort(&self) -> AbortSignal { - self.abort.clone() - } - - fn safe_ret(&self, ret: Result<()>) -> Result<()> { - if ret.is_err() && self.abort.aborted() { - return Ok(()); - } - ret - } -} - -pub enum ReplyEvent { - Text(String), - Done, -} diff --git a/src/render/stream.rs b/src/render/stream.rs index 9fdad5a..6007690 100644 --- a/src/render/stream.rs +++ b/src/render/stream.rs @@ -1,9 +1,8 @@ use super::{MarkdownRender, ReplyEvent}; -use crate::utils::{AbortSignal, Spinner}; +use crate::utils::{run_spinner, AbortSignal}; use anyhow::Result; -use crossbeam::channel::Receiver; use crossterm::{ cursor, event::{self, Event, KeyCode, KeyModifiers}, @@ -12,32 +11,32 @@ use crossterm::{ }; use std::{ io::{self, stdout, Stdout, Write}, - ops::Div, - time::{Duration, Instant}, + time::Duration, }; use textwrap::core::display_width; +use tokio::sync::{mpsc::UnboundedReceiver, oneshot}; -pub fn markdown_stream( - rx: &Receiver, +pub async fn markdown_stream( + rx: UnboundedReceiver, render: &mut MarkdownRender, abort: &AbortSignal, ) -> Result<()> { enable_raw_mode()?; let mut stdout = io::stdout(); - let ret = markdown_stream_inner(rx, render, abort, &mut stdout); + let ret = markdown_stream_inner(rx, render, abort, &mut stdout).await; disable_raw_mode()?; ret } -pub fn raw_stream(rx: &Receiver, abort: &AbortSignal) -> Result<()> { +pub async fn raw_stream(mut rx: UnboundedReceiver, abort: &AbortSignal) -> Result<()> { loop { if abort.aborted() { return Ok(()); } - if let Ok(evt) = rx.try_recv() { + if let Some(evt) = rx.recv().await { match evt { ReplyEvent::Text(text) => { print!("{}", text); @@ -52,30 +51,29 @@ pub fn raw_stream(rx: &Receiver, abort: &AbortSignal) -> Result<()> Ok(()) } -fn markdown_stream_inner( - rx: &Receiver, +async fn markdown_stream_inner( + mut rx: UnboundedReceiver, render: &mut MarkdownRender, abort: &AbortSignal, writer: &mut Stdout, ) -> Result<()> { - let mut last_tick = Instant::now(); - let tick_rate = Duration::from_millis(50); - let mut buffer = String::new(); let mut buffer_rows = 1; let columns = terminal::size()?.0; - let mut spinner = Spinner::new(" Generating"); + let (spinner_tx, spinner_rx) = oneshot::channel(); + let mut spinner_tx = Some(spinner_tx); + tokio::spawn(run_spinner(" Generating", spinner_rx)); 'outer: loop { if abort.aborted() { return Ok(()); } - spinner.step(writer)?; - - for reply_event in gather_events(rx) { - spinner.stop(writer)?; + for reply_event in gather_events(&mut rx).await { + if let Some(spinner_tx) = spinner_tx.take() { + let _ = spinner_tx.send(()); + } match reply_event { ReplyEvent::Text(mut text) => { @@ -135,10 +133,7 @@ fn markdown_stream_inner( } } - let timeout = tick_rate - .checked_sub(last_tick.elapsed()) - .unwrap_or_else(|| tick_rate.div(2)); - if crossterm::event::poll(timeout)? { + if crossterm::event::poll(Duration::from_millis(25))? { if let Event::Key(key) = event::read()? { match key.code { KeyCode::Char('c') if key.modifiers == KeyModifiers::CONTROL => { @@ -153,28 +148,31 @@ fn markdown_stream_inner( } } } - - if last_tick.elapsed() >= tick_rate { - last_tick = Instant::now(); - } } - spinner.stop(writer)?; - + if let Some(spinner_tx) = spinner_tx.take() { + let _ = spinner_tx.send(()); + } Ok(()) } -fn gather_events(rx: &Receiver) -> Vec { +async fn gather_events(rx: &mut UnboundedReceiver) -> Vec { let mut texts = vec![]; let mut done = false; - for reply_event in rx.try_iter() { - match reply_event { - ReplyEvent::Text(v) => texts.push(v), - ReplyEvent::Done => { - done = true; + tokio::select! { + _ = async { + while let Some(reply_event) = rx.recv().await { + match reply_event { + ReplyEvent::Text(v) => texts.push(v), + ReplyEvent::Done => { + done = true; + break; + } + } } - } - } + } => {} + _ = tokio::time::sleep(Duration::from_millis(50)) => {} + }; let mut events = vec![]; if !texts.is_empty() { events.push(ReplyEvent::Text(texts.join(""))) -- cgit v1.2.3