summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/main.rs56
-rw-r--r--src/render/cmd.rs64
-rw-r--r--src/render/markdown.rs40
-rw-r--r--src/render/mod.rs43
-rw-r--r--src/render/stream.rs (renamed from src/render/repl.rs)42
-rw-r--r--src/repl/mod.rs17
-rw-r--r--src/utils/mod.rs2
-rw-r--r--src/utils/split_line.rs212
8 files changed, 88 insertions, 388 deletions
diff --git a/src/main.rs b/src/main.rs
index 368779e..dd9f9cf 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -10,20 +10,17 @@ extern crate log;
mod utils;
use crate::cli::Cli;
-use crate::client::Client;
use crate::config::{Config, GlobalConfig};
use anyhow::Result;
use clap::Parser;
use client::{init_client, list_models};
-use crossbeam::sync::WaitGroup;
use is_terminal::IsTerminal;
use parking_lot::RwLock;
-use render::{render_stream, MarkdownRender};
+use render::{render_error, render_stream, MarkdownRender};
use repl::Repl;
-use std::io::{stdin, Read};
+use std::io::{stderr, stdin, stdout, Read};
use std::sync::Arc;
-use std::{io::stdout, process::exit};
use utils::{cl100k_base_singleton, create_abort_signal};
fn main() -> Result<()> {
@@ -36,18 +33,18 @@ fn main() -> Result<()> {
.roles
.iter()
.for_each(|v| println!("{}", v.name));
- exit(0);
+ return Ok(());
}
if cli.list_models {
for model in list_models(&config.read()) {
println!("{}", model.id());
}
- exit(0);
+ return Ok(());
}
if cli.list_sessions {
let sessions = config.read().list_sessions().join("\n");
println!("{sessions}");
- exit(0);
+ return Ok(());
}
if let Some(wrap) = &cli.wrap {
config.write().set_wrap(wrap)?;
@@ -75,15 +72,22 @@ fn main() -> Result<()> {
if cli.info {
let info = config.read().info()?;
println!("{}", info);
- exit(0);
+ return Ok(());
}
config.write().onstart()?;
let no_stream = cli.no_stream;
- let client = init_client(&config)?;
+ if let Err(err) = start(&config, text, no_stream) {
+ let highlight = stderr().is_terminal() && config.read().highlight;
+ render_error(err, highlight)
+ }
+ Ok(())
+}
+
+fn start(config: &GlobalConfig, text: Option<String>, no_stream: bool) -> Result<()> {
if stdin().is_terminal() {
match text {
- Some(text) => start_directive(client.as_ref(), &config, &text, no_stream),
- None => start_interactive(&config),
+ Some(text) => start_directive(config, &text, no_stream),
+ None => start_interactive(config),
}
} else {
let mut input = String::new();
@@ -91,40 +95,34 @@ fn main() -> Result<()> {
if let Some(text) = text {
input = format!("{text}\n{input}");
}
- start_directive(client.as_ref(), &config, &input, no_stream)
+ start_directive(config, &input, no_stream)
}
}
-fn start_directive(
- client: &dyn Client,
- config: &GlobalConfig,
- input: &str,
- no_stream: bool,
-) -> Result<()> {
+fn start_directive(config: &GlobalConfig, input: &str, no_stream: bool) -> Result<()> {
if let Some(session) = &config.read().session {
session.guard_save()?;
}
- if !stdout().is_terminal() {
- config.write().highlight = false;
- }
+ let client = init_client(config)?;
config.read().maybe_print_send_tokens(input);
let output = if no_stream {
- let render_options = config.read().get_render_options()?;
let output = client.send_message(input)?;
- let mut markdown_render = MarkdownRender::init(render_options)?;
- println!("{}", markdown_render.render(&output).trim());
+ if stdout().is_terminal() {
+ let render_options = config.read().get_render_options()?;
+ let mut markdown_render = MarkdownRender::init(render_options)?;
+ println!("{}", markdown_render.render(&output).trim());
+ } else {
+ println!("{}", output);
+ }
output
} else {
- let wg = WaitGroup::new();
let abort = create_abort_signal();
let abort_clone = abort.clone();
ctrlc::set_handler(move || {
abort_clone.set_ctrlc();
})
.expect("Failed to setting Ctrl-C handler");
- let output = render_stream(input, client, config, false, abort, wg.clone())?;
- wg.wait();
- output
+ render_stream(input, client.as_ref(), config, abort)?
};
config.write().save_message(input, &output)
}
diff --git a/src/render/cmd.rs b/src/render/cmd.rs
deleted file mode 100644
index e942025..0000000
--- a/src/render/cmd.rs
+++ /dev/null
@@ -1,64 +0,0 @@
-use super::{MarkdownRender, ReplyEvent};
-
-use crate::utils::{split_line_sematic, split_line_tail, AbortSignal};
-
-use anyhow::Result;
-use crossbeam::channel::Receiver;
-use textwrap::core::display_width;
-
-pub fn cmd_render_stream(
- rx: &Receiver<ReplyEvent>,
- render: &mut MarkdownRender,
- abort: &AbortSignal,
-) -> Result<()> {
- let mut buffer = String::new();
- let mut indent = 0;
- loop {
- if abort.aborted() {
- return Ok(());
- }
- if let Ok(evt) = rx.try_recv() {
- match evt {
- ReplyEvent::Text(text) => {
- if text.contains('\n') {
- let text = format!("{buffer}{text}");
- let (head, tail) = split_line_tail(&text);
- buffer = tail.to_string();
- let output = render.render_with_indent(head, indent);
- println!("{}", output);
- indent = 0;
- } else {
- buffer = format!("{buffer}{text}");
- if !(render.is_code()
- || buffer.len() < 40
- || buffer.starts_with('#')
- || buffer.starts_with('>')
- || buffer.starts_with('|'))
- {
- if let Some((head, remain)) = split_line_sematic(&buffer) {
- buffer = remain;
- let output = render.render_with_indent(&head, indent);
- let (_, tail) = split_line_tail(&output);
- if let Some(width) = render.wrap_width() {
- if output.contains('\n') {
- indent = display_width(tail);
- } else {
- indent += display_width(&output);
- }
- indent %= width as usize;
- }
- print!("{}", output);
- }
- }
- }
- }
- ReplyEvent::Done => {
- let output = render.render_with_indent(&buffer, indent);
- println!("{}", output);
- break;
- }
- }
- }
- }
- Ok(())
-}
diff --git a/src/render/markdown.rs b/src/render/markdown.rs
index f300854..44e24b3 100644
--- a/src/render/markdown.rs
+++ b/src/render/markdown.rs
@@ -64,17 +64,6 @@ impl MarkdownRender {
})
}
- pub(crate) const fn is_code(&self) -> bool {
- matches!(
- self.prev_line_type,
- LineType::CodeBegin | LineType::CodeInner
- )
- }
-
- pub(crate) const fn wrap_width(&self) -> Option<u16> {
- self.wrap_width
- }
-
pub fn render(&mut self, text: &str) -> String {
text.split('\n')
.map(|line| self.render_line_mut(line))
@@ -82,16 +71,6 @@ impl MarkdownRender {
.join("\n")
}
- pub fn render_with_indent(&mut self, text: &str, indent: usize) -> String {
- let text = format!("{}{}", " ".repeat(indent), text);
- let output = self.render(&text);
- if output.starts_with('\n') {
- output
- } else {
- output.chars().skip(indent).collect()
- }
- }
-
pub fn render_line(&self, line: &str) -> String {
let (_, code_syntax, is_code) = self.check_line(line);
if is_code {
@@ -377,23 +356,4 @@ std::error::Error>> {
let output = render.render(TEXT);
assert_eq!(TEXT_WRAP_ALL, output);
}
-
- #[test]
- fn wrap_with_indent() {
- let options = RenderOptions::default();
- let mut render = MarkdownRender::init(options).unwrap();
- render.wrap_width = Some(80);
-
- let input = "To unzip a file in Rust, you can use the `zip` crate. Here's an example code";
- let output = render.render_with_indent(input, 40);
- let expect =
- "To unzip a file in Rust, you can use the\n`zip` crate. Here's an example code";
- assert_eq!(output, expect);
-
- let input = "Unzip a file";
- let output = render.render_with_indent(input, 76);
- let expect = "\nUnzip a file";
-
- assert_eq!(output, expect);
- }
}
diff --git a/src/render/mod.rs b/src/render/mod.rs
index 2b97557..1e20ccc 100644
--- a/src/render/mod.rs
+++ b/src/render/mod.rs
@@ -1,10 +1,8 @@
-mod cmd;
mod markdown;
-mod repl;
+mod stream;
-use self::cmd::cmd_render_stream;
pub use self::markdown::{MarkdownRender, RenderOptions};
-use self::repl::repl_render_stream;
+use self::stream::{markdown_stream, raw_stream};
use crate::client::Client;
use crate::config::GlobalConfig;
@@ -13,17 +11,19 @@ use crate::utils::AbortSignal;
use anyhow::{Context, Result};
use crossbeam::channel::{unbounded, Sender};
use crossbeam::sync::WaitGroup;
+use is_terminal::IsTerminal;
use nu_ansi_term::{Color, Style};
+use std::io::stdout;
use std::thread::spawn;
pub fn render_stream(
input: &str,
client: &dyn Client,
config: &GlobalConfig,
- repl: bool,
abort: AbortSignal,
- wg: WaitGroup,
) -> Result<String> {
+ 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();
@@ -31,33 +31,44 @@ pub fn render_stream(
let highlight = config.read().highlight;
spawn(move || {
let run = move || {
- if repl {
+ if stdout().is_terminal() {
let mut render = MarkdownRender::init(render_options)?;
- repl_render_stream(&rx, &mut render, &abort)
+ markdown_stream(&rx, &mut render, &abort)
} else {
- let mut render = MarkdownRender::init(render_options)?;
- cmd_render_stream(&rx, &mut render, &abort)
+ raw_stream(&rx, &abort)
}
};
if let Err(err) = run() {
render_error(err, highlight);
}
- drop(wg);
+ drop(wg_cloned);
});
ReplyHandler::new(tx, abort_clone)
};
- client.send_message_streaming(input, &mut stream_handler)?;
- let buffer = stream_handler.get_buffer();
- Ok(buffer.to_string())
+ 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)
+ }
+ }
}
pub fn render_error(err: anyhow::Error, highlight: bool) {
let err = format!("{err:?}");
if highlight {
let style = Style::new().fg(Color::Red);
- println!("{}", style.paint(err.trim()));
+ eprintln!("{}", style.paint(err));
} else {
- println!("{}", err.trim());
+ eprintln!("{err}");
}
}
diff --git a/src/render/repl.rs b/src/render/stream.rs
index eeb5420..189862d 100644
--- a/src/render/repl.rs
+++ b/src/render/stream.rs
@@ -1,6 +1,6 @@
use super::{MarkdownRender, ReplyEvent};
-use crate::utils::{split_line_tail, AbortSignal};
+use crate::utils::AbortSignal;
use anyhow::Result;
use crossbeam::channel::Receiver;
@@ -16,7 +16,7 @@ use std::{
};
use textwrap::core::display_width;
-pub fn repl_render_stream(
+pub fn markdown_stream(
rx: &Receiver<ReplyEvent>,
render: &mut MarkdownRender,
abort: &AbortSignal,
@@ -24,14 +24,33 @@ pub fn repl_render_stream(
enable_raw_mode()?;
let mut stdout = io::stdout();
- let ret = repl_render_stream_inner(rx, render, abort, &mut stdout);
+ let ret = markdown_stream_inner(rx, render, abort, &mut stdout);
disable_raw_mode()?;
ret
}
-fn repl_render_stream_inner(
+pub fn raw_stream(rx: &Receiver<ReplyEvent>, abort: &AbortSignal) -> Result<()> {
+ loop {
+ if abort.aborted() {
+ return Ok(());
+ }
+ if let Ok(evt) = rx.try_recv() {
+ match evt {
+ ReplyEvent::Text(text) => {
+ print!("{}", text);
+ }
+ ReplyEvent::Done => {
+ break;
+ }
+ }
+ }
+ }
+ Ok(())
+}
+
+fn markdown_stream_inner(
rx: &Receiver<ReplyEvent>,
render: &mut MarkdownRender,
abort: &AbortSignal,
@@ -104,13 +123,6 @@ fn repl_render_stream_inner(
writer.flush()?;
}
ReplyEvent::Done => {
- #[cfg(target_os = "windows")]
- let eol = "\n\n";
- #[cfg(not(target_os = "windows"))]
- let eol = "\n";
- queue!(writer, style::Print(eol))?;
- writer.flush()?;
-
break;
}
}
@@ -157,6 +169,14 @@ fn print_block(writer: &mut Stdout, text: &str, columns: u16) -> Result<u16> {
Ok(num)
}
+fn split_line_tail(text: &str) -> (&str, &str) {
+ if let Some((head, tail)) = text.rsplit_once('\n') {
+ (head, tail)
+ } else {
+ ("", text)
+ }
+}
+
fn need_rows(text: &str, columns: u16) -> u16 {
let buffer_width = display_width(text).max(1) as u16;
(buffer_width + columns - 1) / columns
diff --git a/src/repl/mod.rs b/src/repl/mod.rs
index 8ffb36c..511e069 100644
--- a/src/repl/mod.rs
+++ b/src/repl/mod.rs
@@ -12,7 +12,6 @@ use crate::render::{render_error, render_stream};
use crate::utils::{create_abort_signal, set_text, AbortSignal};
use anyhow::{bail, Context, Result};
-use crossbeam::sync::WaitGroup;
use fancy_regex::Regex;
use lazy_static::lazy_static;
use reedline::Signal;
@@ -235,21 +234,11 @@ impl Repl {
return Ok(());
}
self.config.read().maybe_print_send_tokens(input);
- let wg = WaitGroup::new();
let client = init_client(&self.config)?;
- let ret = render_stream(
- input,
- client.as_ref(),
- &self.config,
- true,
- self.abort.clone(),
- wg.clone(),
- );
- wg.wait();
- let buffer = ret?;
- self.config.write().save_message(input, &buffer)?;
+ let output = render_stream(input, client.as_ref(), &self.config, self.abort.clone())?;
+ self.config.write().save_message(input, &output)?;
if self.config.read().auto_copy {
- let _ = self.copy(&buffer);
+ let _ = self.copy(&output);
}
Ok(())
}
diff --git a/src/utils/mod.rs b/src/utils/mod.rs
index b94f83e..0039298 100644
--- a/src/utils/mod.rs
+++ b/src/utils/mod.rs
@@ -1,13 +1,11 @@
mod abort_signal;
mod clipboard;
mod prompt_input;
-mod split_line;
mod tiktoken;
pub use self::abort_signal::{create_abort_signal, AbortSignal};
pub use self::clipboard::set_text;
pub use self::prompt_input::*;
-pub use self::split_line::*;
pub use self::tiktoken::cl100k_base_singleton;
pub fn now() -> String {
diff --git a/src/utils/split_line.rs b/src/utils/split_line.rs
deleted file mode 100644
index da10ad1..0000000
--- a/src/utils/split_line.rs
+++ /dev/null
@@ -1,212 +0,0 @@
-pub fn split_line_sematic(text: &str) -> Option<(String, String)> {
- let mut balance: Vec<Kind> = Vec::new();
- let chars: Vec<char> = text.chars().collect();
- let mut index = 0;
- let len = chars.len();
- while index < len - 1 {
- let ch = chars[index];
- if balance.is_empty()
- && ((matches!(ch, ',' | '.' | ';') && chars[index + 1].is_whitespace())
- || matches!(ch, ',' | '。' | ';'))
- {
- let (output, remain) = chars.split_at(index + 1);
- return Some((output.iter().collect(), remain.iter().collect()));
- }
- if index + 2 < len && do_balance(&mut balance, &chars[index..=index + 2]) {
- index += 3;
- continue;
- }
- if do_balance(&mut balance, &chars[index..=index + 1]) {
- index += 2;
- continue;
- }
- do_balance(&mut balance, &chars[index..=index]);
- index += 1;
- }
-
- None
-}
-
-pub fn split_line_tail(text: &str) -> (&str, &str) {
- if let Some((head, tail)) = text.rsplit_once('\n') {
- (head, tail)
- } else {
- ("", text)
- }
-}
-
-#[derive(Debug, Clone, Copy, Eq, PartialEq)]
-enum Kind {
- ParentheseStart,
- ParentheseEnd,
- BracketStart,
- BracketEnd,
- Asterisk,
- Asterisk2,
- SingleQuota,
- DoubleQuota,
- Tilde,
- Tilde2,
- Backtick,
- Backtick3,
-}
-
-impl Kind {
- fn from_chars(chars: &[char]) -> Option<Self> {
- let kind = match chars.len() {
- 1 => match chars[0] {
- '(' => Self::ParentheseStart,
- ')' => Self::ParentheseEnd,
- '[' => Self::BracketStart,
- ']' => Self::BracketEnd,
- '*' => Self::Asterisk,
- '\'' => Self::SingleQuota,
- '"' => Self::DoubleQuota,
- '~' => Self::Tilde,
- '`' => Self::Backtick,
- _ => return None,
- },
- 2 if chars[0] == chars[1] => match chars[0] {
- '*' => Self::Asterisk2,
- '~' => Self::Tilde2,
- _ => return None,
- },
- 3 => {
- if chars == ['`', '`', '`'] {
- Self::Backtick3
- } else {
- return None;
- }
- }
- _ => return None,
- };
- Some(kind)
- }
-}
-
-fn do_balance(balance: &mut Vec<Kind>, chars: &[char]) -> bool {
- Kind::from_chars(chars).map_or(false, |kind| {
- let last = balance.last();
- match (kind, last) {
- (Kind::ParentheseEnd, Some(&Kind::ParentheseStart))
- | (Kind::BracketEnd, Some(&Kind::BracketStart))
- | (Kind::Asterisk, Some(&Kind::Asterisk))
- | (Kind::Asterisk2, Some(&Kind::Asterisk2))
- | (Kind::SingleQuota, Some(&Kind::SingleQuota))
- | (Kind::DoubleQuota, Some(&Kind::DoubleQuota))
- | (Kind::Tilde, Some(&Kind::Tilde))
- | (Kind::Tilde2, Some(&Kind::Tilde2))
- | (Kind::Backtick, Some(&Kind::Backtick))
- | (Kind::Backtick3, Some(&Kind::Backtick3)) => {
- balance.pop();
- true
- }
- (
- Kind::ParentheseStart
- | Kind::BracketStart
- | Kind::Asterisk
- | Kind::Asterisk2
- | Kind::SingleQuota
- | Kind::DoubleQuota
- | Kind::Tilde
- | Kind::Tilde2
- | Kind::Backtick
- | Kind::Backtick3,
- _,
- ) => {
- balance.push(kind);
- true
- }
- _ => false,
- }
- })
-}
-
-#[cfg(test)]
-mod tests {
- use super::*;
-
- macro_rules! assert_split_line {
- ($a:literal, $b:literal, true) => {
- assert_eq!(
- split_line_sematic(&format!("{}{}", $a, $b)),
- Some(($a.into(), $b.into()))
- );
- };
- ($a:literal, $b:literal, false) => {
- assert_eq!(split_line_sematic(&format!("{}{}", $a, $b)), None);
- };
- }
-
- #[test]
- fn test_split_line() {
- assert_split_line!(
- "Wikipedia is a free online encyclopedia,",
- " that anyone can edit,",
- true
- );
- assert_split_line!(
- "Wikipedia is a free online encyclopedia.",
- " that anyone can edit,",
- true
- );
- assert_split_line!("床前明月光,", "疑是地上霜。", true);
- assert_split_line!("床前明月光。", "疑是地上霜。", true);
- assert_split_line!("床前明月光;", "疑是地上霜。", true);
- assert_split_line!(
- "Wikipedia is (a free online encyclopedia).",
- " that anyone can edit.",
- true
- );
- assert_split_line!(
- "Wikipedia is a free online `encyclopedia,",
- " that` anyone can edit.",
- false
- );
- assert_split_line!(
- "Wikipedia is a free online ```encyclopedia,",
- " that``` anyone can edit.",
- false
- );
- assert_split_line!(
- "Wikipedia is a free online *encyclopedia,",
- " that* anyone can edit.",
- false
- );
- assert_split_line!(
- "Wikipedia is a free online **encyclopedia,",
- " that** anyone can edit.",
- false
- );
- assert_split_line!(
- "Wikipedia is a free online ~encyclopedia,",
- " that~ anyone can edit.",
- false
- );
- assert_split_line!(
- "Wikipedia is a free online ~~encyclopedia,",
- " that~~ anyone can edit.",
- false
- );
- assert_split_line!(
- "Wikipedia is a free online ``encyclopedia,",
- " that`` anyone can edit.",
- true
- );
- assert_split_line!(
- "Wikipedia is a free online \"encyclopedia,",
- " that\" anyone can edit.",
- false
- );
- assert_split_line!(
- "Wikipedia is a free online 'encyclopedia,",
- " that' anyone can edit.",
- false
- );
- assert_split_line!(
- "Wikipedia is a free online encyclopedia.",
- "that anyone can edit.",
- false
- );
- }
-}