summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-03-04 09:02:57 +0800
committerGitHub <noreply@github.com>2023-03-04 09:02:57 +0800
commit0fa1ae215ac1434591968818f8f75ad5b35d52ec (patch)
tree2417bc0e95fad6886dd414cf407bb9fc0dbf3108 /src
parent3ffebce8bbc530429086302ec4f9a95268c25d07 (diff)
downloadaichat-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.rs41
-rw-r--r--src/config.rs11
-rw-r--r--src/main.rs16
-rw-r--r--src/render.rs189
-rw-r--r--src/repl.rs126
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() {