diff options
| author | sigoden <sigoden@gmail.com> | 2023-03-04 09:02:57 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-03-04 09:02:57 +0800 |
| commit | 0fa1ae215ac1434591968818f8f75ad5b35d52ec (patch) | |
| tree | 2417bc0e95fad6886dd414cf407bb9fc0dbf3108 /src | |
| parent | 3ffebce8bbc530429086302ec4f9a95268c25d07 (diff) | |
| download | aichat-0fa1ae215ac1434591968818f8f75ad5b35d52ec.tar.gz | |
feat: support highlight reply markdown (#3)
* feat: support highlight reply markdown
* migrate markdown highlighter from termimad to mdcat
* optimize render
No need to clear screen when there is no newline in reply token
* handle ctrl-c when rendering stream
* update readme
* fix ctrlc don't abort acquire_stream when establish connection
* ensure render_stream's dtect_ctrlc is exit before next readline
This will make reedline don't throw err 'The cursor position could not be read within a normal duration'
Diffstat (limited to 'src')
| -rw-r--r-- | src/client.rs | 41 | ||||
| -rw-r--r-- | src/config.rs | 11 | ||||
| -rw-r--r-- | src/main.rs | 16 | ||||
| -rw-r--r-- | src/render.rs | 189 | ||||
| -rw-r--r-- | src/repl.rs | 126 |
5 files changed, 324 insertions, 59 deletions
diff --git a/src/client.rs b/src/client.rs index 3f42eb4..1f42346 100644 --- a/src/client.rs +++ b/src/client.rs @@ -1,4 +1,5 @@ use crate::config::Config; +use crate::repl::ReplyReceiver; use anyhow::{anyhow, Result}; use eventsource_stream::Eventsource; @@ -8,6 +9,7 @@ use serde_json::{json, Value}; use std::sync::atomic::{AtomicBool, Ordering}; use std::{sync::Arc, time::Duration}; use tokio::runtime::Runtime; +use tokio::time::sleep; const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); const API_URL: &str = "https://api.openai.com/v1/chat/completions"; @@ -45,22 +47,31 @@ impl ChatGptClient { .block_on(async { self.acquire_inner(input, prompt).await }) } - pub fn acquire_stream<T>( + pub fn acquire_stream( &self, input: &str, prompt: Option<String>, - output: &mut String, - handler: T, + receiver: &mut ReplyReceiver, ctrlc: Arc<AtomicBool>, - ) -> Result<()> - where - T: FnOnce(&mut String, &str) + Copy, - { + ) -> Result<()> { + async fn watch_ctrlc(ctrlc: Arc<AtomicBool>) { + loop { + if ctrlc.load(Ordering::SeqCst) { + break; + } + sleep(Duration::from_millis(100)).await; + } + } self.runtime.block_on(async { tokio::select! { - ret = self.acquire_stream_inner(input, prompt, handler, output) => { + ret = self.acquire_stream_inner(input, prompt, receiver) => { + receiver.done(); ret } + _ = watch_ctrlc(ctrlc.clone()) => { + receiver.done(); + Ok(()) + }, _ = tokio::signal::ctrl_c() => { ctrlc.store(true, Ordering::SeqCst); Ok(()) @@ -85,19 +96,15 @@ impl ChatGptClient { Ok(output.to_string()) } - async fn acquire_stream_inner<T>( + async fn acquire_stream_inner( &self, content: &str, prompt: Option<String>, - handler: T, - output: &mut String, - ) -> Result<()> - where - T: FnOnce(&mut String, &str) + Copy, - { + receiver: &mut ReplyReceiver, + ) -> Result<()> { let content = combine(content, prompt); if self.config.dry_run { - handler(output, &content); + receiver.text(&content); return Ok(()); } let builder = self.request_builder(&content, true); @@ -121,7 +128,7 @@ impl ChatGptClient { continue; } } - handler(output, text); + receiver.text(text); } } diff --git a/src/config.rs b/src/config.rs index c3ebf3c..cce6b6a 100644 --- a/src/config.rs +++ b/src/config.rs @@ -24,6 +24,9 @@ pub struct Config { /// Whether to persistently save chat messages #[serde(default)] pub save: bool, + /// Whether to highlight reply message + #[serde(default)] + pub highlight: bool, /// Set proxy pub proxy: Option<String>, /// Used only for debugging @@ -172,6 +175,14 @@ fn create_config_file(config_path: &Path) -> Result<()> { raw_config.push_str("save: true\n"); } + let ans = Confirm::new("Whether to highlight reply message?") + .with_default(true) + .prompt() + .map_err(confirm_map_err)?; + if ans { + raw_config.push_str("highlight: true\n"); + } + std::fs::write(config_path, raw_config) .map_err(|err| anyhow!("Failed to write to config file, {err}"))?; Ok(()) diff --git a/src/main.rs b/src/main.rs index ddd856e..893850a 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,17 +1,20 @@ mod cli; mod client; mod config; +mod render; mod repl; -use std::process::exit; use std::sync::Arc; +use std::{io::stdout, process::exit}; use cli::Cli; use client::ChatGptClient; use config::{Config, Role}; +use is_terminal::IsTerminal; use anyhow::{anyhow, Result}; use clap::Parser; +use render::MarkdownRender; use repl::{Repl, ReplCmdHandler}; fn main() { @@ -52,8 +55,15 @@ fn start_directive( ) -> Result<()> { let mut file = config.open_message_file()?; let output = client.acquire(input, role.map(|v| v.prompt))?; - println!("{}", output.trim()); - Config::save_message(file.as_mut(), input, &output); + let output = output.trim(); + if config.highlight && stdout().is_terminal() { + let markdown_render = MarkdownRender::init()?; + markdown_render.print(output)?; + } else { + println!("{output}"); + } + + Config::save_message(file.as_mut(), input, output); Ok(()) } diff --git a/src/render.rs b/src/render.rs new file mode 100644 index 0000000..6a052e2 --- /dev/null +++ b/src/render.rs @@ -0,0 +1,189 @@ +use anyhow::Result; +use crossbeam::sync::WaitGroup; +use crossterm::{ + cursor, + event::{self, Event, KeyCode, KeyEvent, KeyModifiers}, + execute, queue, style, + terminal::{ + self, disable_raw_mode, enable_raw_mode, size, ClearType, EnterAlternateScreen, + LeaveAlternateScreen, + }, +}; +use mdcat::{ + push_tty, + terminal::{TerminalProgram, TerminalSize}, + Environment, ResourceAccess, Settings, +}; +use pulldown_cmark::Parser; +use std::{ + io::{self, Write}, + sync::{ + atomic::{AtomicBool, Ordering}, + mpsc::Receiver, + Arc, + }, + thread, + time::Duration, +}; +use syntect::parsing::SyntaxSet; + +use crate::repl::{dump, ReplyEvent}; + +pub fn render_stream( + rx: Receiver<ReplyEvent>, + ctrlc: Arc<AtomicBool>, + markdown_render: Arc<MarkdownRender>, +) -> Result<()> { + let wg = WaitGroup::new(); + let ctrlc_clone = ctrlc.clone(); + let stream_done = Arc::new(AtomicBool::new(false)); + let stream_done_clone = stream_done.clone(); + let wg_clone = wg.clone(); + thread::spawn(move || { + let _ = detect_ctrlc(ctrlc_clone, stream_done_clone); + drop(wg_clone); + }); + let ret = render_stream_inner(rx, ctrlc, markdown_render); + stream_done.store(true, Ordering::SeqCst); + wg.wait(); + ret +} + +fn detect_ctrlc(ctrlc: Arc<AtomicBool>, stream_done: Arc<AtomicBool>) -> Result<()> { + loop { + if ctrlc.load(Ordering::SeqCst) || stream_done.load(Ordering::SeqCst) { + return Ok(()); + } + if event::poll(Duration::from_millis(100))? { + if let Event::Key(KeyEvent { + code: KeyCode::Char('c'), + modifiers: KeyModifiers::CONTROL, + .. + }) = event::read()? + { + ctrlc.store(true, Ordering::SeqCst); + break; + } + } + } + Ok(()) +} + +fn render_stream_inner( + rx: Receiver<ReplyEvent>, + ctrlc: Arc<AtomicBool>, + markdown_render: Arc<MarkdownRender>, +) -> Result<()> { + // setup terminal + enable_raw_mode()?; + let mut output = String::new(); + let mut stdout = io::stdout(); + execute!(stdout, EnterAlternateScreen)?; + + fn clear(stdout: &mut impl Write) -> io::Result<()> { + queue!( + stdout, + style::ResetColor, + terminal::Clear(ClearType::All), + cursor::Hide, + cursor::MoveTo(0, 0) + ) + } + + clear(&mut stdout)?; + + while let Ok(ev) = rx.recv() { + if ctrlc.load(Ordering::SeqCst) { + break; + } + match ev { + ReplyEvent::Text(text) => { + output.push_str(&text); + let rows = size()?.1 as usize; + let lines: Vec<&str> = output.split('\n').collect(); + let len = lines.len(); + let skip = if len > rows { len - rows } else { 0 }; + let mut selected_lines = vec![]; + let mut count_begin_code = 0; + let mut code = None; + for (index, line) in lines.iter().enumerate() { + if index < skip { + if line.starts_with("```") { + count_begin_code += 1; + code = Some(*line); + } + } else { + selected_lines.push(*line); + } + } + if count_begin_code % 2 == 1 { + if let Some(code) = code { + selected_lines[0] = code + } + }; + let content = selected_lines.join("\n"); + let markdown = markdown_render.render(&content)?; + if text.contains('\n') { + clear(&mut stdout)?; + for line in markdown.split('\n') { + queue!(stdout, style::Print(line), cursor::MoveToNextLine(1))?; + } + } else if let Some(line) = markdown.split('\n').last() { + queue!( + stdout, + style::ResetColor, + terminal::Clear(ClearType::CurrentLine), + cursor::MoveToColumn(0), + style::Print(line) + )?; + } + + stdout.flush()?; + } + ReplyEvent::Done => { + break; + } + } + } + + execute!(stdout, style::ResetColor, cursor::Show)?; + + // restore terminal + disable_raw_mode()?; + execute!(stdout, LeaveAlternateScreen)?; + + Ok(()) +} + +pub struct MarkdownRender { + env: Environment, + settings: Settings, +} + +impl MarkdownRender { + pub fn init() -> Result<Self> { + let terminal = TerminalProgram::detect(); + let env = + Environment::for_local_directory(&std::env::current_dir().expect("Working directory"))?; + let settings = Settings { + resource_access: ResourceAccess::LocalOnly, + syntax_set: SyntaxSet::load_defaults_newlines(), + terminal_capabilities: terminal.capabilities(), + terminal_size: TerminalSize::default(), + }; + Ok(Self { env, settings }) + } + + pub fn print(&self, input: &str) -> Result<()> { + let markdown = self.render(input)?; + dump(markdown, 0); + Ok(()) + } + + pub fn render(&self, input: &str) -> Result<String> { + let source = Parser::new(input); + let mut sink = Vec::new(); + push_tty(&self.settings, &self.env, &mut sink, source)?; + Ok(String::from_utf8_lossy(&sink).into()) + } +} diff --git a/src/repl.rs b/src/repl.rs index 2bb2aef..f8a10a3 100644 --- a/src/repl.rs +++ b/src/repl.rs @@ -1,7 +1,8 @@ use crate::client::ChatGptClient; use crate::config::{Config, Role}; +use crate::render::{self, MarkdownRender}; use anyhow::{anyhow, Result}; -use inquire::Editor; +use crossbeam::sync::WaitGroup; use reedline::{ default_emacs_keybindings, ColumnarMenu, DefaultCompleter, DefaultPrompt, DefaultPromptSegment, Emacs, FileBackedHistory, KeyCode, KeyModifiers, Keybindings, Reedline, ReedlineEvent, @@ -11,9 +12,12 @@ use std::cell::RefCell; use std::fs::File; use std::io::{stdout, Write}; use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::mpsc::channel; +use std::sync::mpsc::Sender; use std::sync::Arc; +use std::thread::spawn; -const REPL_COMMANDS: [(&str, &str); 8] = [ +const REPL_COMMANDS: [(&str, &str); 7] = [ (".clear", "Clear the screen"), (".clear-history", "Clear the history"), (".clear-role", "Clear the role status"), @@ -21,7 +25,6 @@ const REPL_COMMANDS: [(&str, &str); 8] = [ (".help", "Print this help message"), (".history", "Print the history"), (".role", "Specify the role that the AI will play"), - (".view", "Use an external editor to view the AI reply"), ]; const MENU_NAME: &str = "completion_menu"; @@ -86,13 +89,9 @@ impl Repl { Ok(Signal::CtrlD) => { break; } - Err(err) => { - dump(format!("{err:?}"), 1); - break; - } + _ => {} } } - // tx.send(ReplCmd::Quit).unwrap(); Ok(()) } @@ -103,7 +102,6 @@ impl Repl { None => (line.as_str(), None), }; match cmd { - ".view" => handler.handle(ReplCmd::View)?, ".exit" => { return Ok(true); } @@ -187,6 +185,7 @@ pub struct ReplCmdHandler { config: Arc<Config>, state: RefCell<ReplCmdHandlerState>, ctrlc: Arc<AtomicBool>, + render: Option<Arc<MarkdownRender>>, } struct ReplCmdHandlerState { @@ -197,6 +196,11 @@ struct ReplCmdHandlerState { impl ReplCmdHandler { pub fn init(client: ChatGptClient, config: Arc<Config>, role: Option<Role>) -> Result<Self> { + let render = if config.highlight { + Some(Arc::new(MarkdownRender::init()?)) + } else { + None + }; let prompt = role.map(|v| v.prompt).unwrap_or_default(); let save_file = config.open_message_file()?; let ctrlc = Arc::new(AtomicBool::new(false)); @@ -208,14 +212,14 @@ impl ReplCmdHandler { Ok(Self { client, config, - ctrlc, state, + ctrlc, + render, }) } fn handle(&self, cmd: ReplCmd) -> Result<()> { match cmd { ReplCmd::Input(input) => { - let mut output = String::new(); if input.is_empty() { self.state.borrow_mut().output.clear(); return Ok(()); @@ -226,27 +230,37 @@ impl ReplCmdHandler { } else { Some(prompt) }; - self.client.acquire_stream( + let wg = WaitGroup::new(); + let mut receiver = if let Some(markdown_render) = self.render.clone() { + let (tx, rx) = channel(); + let ctrlc = self.ctrlc.clone(); + let wg = wg.clone(); + spawn(move || { + let _ = render::render_stream(rx, ctrlc, markdown_render); + drop(wg); + }); + ReplyReceiver::new(Some(tx)) + } else { + ReplyReceiver::new(None) + }; + self.client + .acquire_stream(&input, prompt, &mut receiver, self.ctrlc.clone())?; + Config::save_message( + self.state.borrow_mut().save_file.as_mut(), &input, - prompt, - &mut output, - dump_and_collect, - self.ctrlc.clone(), - )?; - dump_and_collect(&mut output, "\n\n"); - Config::save_message(self.state.borrow_mut().save_file.as_mut(), &input, &output); - self.state.borrow_mut().output = output; - } - ReplCmd::View => { - let output = self.state.borrow().output.to_string(); - if output.is_empty() { - return Ok(()); + &receiver.output, + ); + wg.wait(); + match self.render.clone() { + Some(markdown_render) => { + markdown_render.print(&receiver.output)?; + dump("", 1); + } + None => { + dump(&receiver.output, 2); + } } - let _ = Editor::new("view ai reply with an external editor") - .with_file_extension(".md") - .with_predefined_text(&output) - .prompt()?; - dump("", 1); + self.state.borrow_mut().output = receiver.output; } ReplCmd::SetRole(name) => match self.config.find_role(&name) { Some(v) => { @@ -265,21 +279,55 @@ impl ReplCmdHandler { } } -pub enum ReplCmd { - View, - UnsetRole, - Input(String), - SetRole(String), +pub struct ReplyReceiver { + output: String, + sender: Option<Sender<ReplyEvent>>, +} + +impl ReplyReceiver { + pub fn new(sender: Option<Sender<ReplyEvent>>) -> Self { + Self { + output: String::new(), + sender, + } + } + pub fn text(&mut self, text: &str) { + match self.sender.as_ref() { + Some(tx) => { + let _ = tx.send(ReplyEvent::Text(text.to_string())); + } + None => { + dump(text, 0); + } + } + self.output.push_str(text); + } + pub fn done(&mut self) { + match self.sender.as_ref() { + Some(tx) => { + let _ = tx.send(ReplyEvent::Done); + } + None => { + dump("", 2); + } + } + } +} + +pub enum ReplyEvent { + Text(String), + Done, } pub fn dump<T: ToString>(text: T, newlines: usize) { print!("{}{}", text.to_string(), "\n".repeat(newlines)); - stdout().flush().unwrap(); + let _ = stdout().flush(); } -fn dump_and_collect(output: &mut String, reply: &str) { - output.push_str(reply); - dump(reply, 0); +enum ReplCmd { + UnsetRole, + Input(String), + SetRole(String), } fn dump_repl_help() { |
