summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--Cargo.lock14
-rw-r--r--Cargo.toml16
-rw-r--r--src/main.rs73
3 files changed, 79 insertions, 24 deletions
diff --git a/Cargo.lock b/Cargo.lock
index 4171e55..1498741 100644
--- a/Cargo.lock
+++ b/Cargo.lock
@@ -615,6 +615,7 @@ dependencies = [
"dyn-clone",
"lazy_static",
"newline-converter",
+ "tempfile",
"thiserror",
"unicode-segmentation",
"unicode-width",
@@ -1054,6 +1055,7 @@ dependencies = [
"serde_urlencoded",
"tokio",
"tokio-native-tls",
+ "tokio-socks",
"tokio-util",
"tower-service",
"url",
@@ -1396,6 +1398,18 @@ dependencies = [
]
[[package]]
+name = "tokio-socks"
+version = "0.5.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "51165dfa029d2a65969413a6cc96f354b86b464498702f174a4efa13608fd8c0"
+dependencies = [
+ "either",
+ "futures-util",
+ "thiserror",
+ "tokio",
+]
+
+[[package]]
name = "tokio-util"
version = "0.7.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
diff --git a/Cargo.toml b/Cargo.toml
index 360aa82..6441939 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -2,6 +2,13 @@
name = "aichat"
version = "0.1.0"
edition = "2021"
+authors = ["sigoden <sigoden@gmail.com>"]
+description = "Chat with OpenAI GPT-3.5 in the terminal."
+license = "MIT OR Apache-2.0"
+homepage = "https://github.com/sigoden/aichat"
+repository = "https://github.com/sigoden/aichat"
+categories = ["command-line-utilities"]
+keywords = ["chatgpt", "ai"]
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
@@ -12,10 +19,15 @@ clap = { version = "4.1.8", features = ["derive", "string"] }
dirs = "4.0.0"
eventsource-stream = "0.2.3"
futures-util = "0.3.26"
-inquire = "0.5.3"
+inquire = { version = "0.5.3", features = ["editor"] }
reedline = "0.16.0"
-reqwest = { version = "0.11.14", features = ["json", "stream"] }
+reqwest = { version = "0.11.14", features = ["json", "stream", "socks"] }
serde = { version = "1.0.152", features = ["derive"] }
serde_json = "1.0.93"
tokio = { version = "1.26.0", features = ["full"] }
toml = "0.7.2"
+
+[profile.release]
+lto = true
+strip = true
+opt-level = "z"
diff --git a/src/main.rs b/src/main.rs
index a1add67..3a76049 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -13,7 +13,7 @@ use clap::{Arg, ArgAction, Command};
use eventsource_stream::{EventStream, Eventsource};
use futures_util::Stream;
use futures_util::StreamExt;
-use inquire::{Confirm, Text};
+use inquire::{Confirm, Editor, Text};
use reedline::{
default_emacs_keybindings, ColumnarMenu, DefaultCompleter, DefaultPrompt, DefaultPromptSegment,
Emacs, FileBackedHistory, KeyCode, KeyModifiers, Reedline, ReedlineEvent, ReedlineMenu, Signal,
@@ -25,13 +25,14 @@ use tokio::runtime::Runtime;
const CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
const API_URL: &str = "https://api.openai.com/v1/chat/completions";
const MODEL: &str = "gpt-3.5-turbo";
-const REPL_COMMANDS: [(&str, &str); 6] = [
+const REPL_COMMANDS: [(&str, &str); 7] = [
(".clear", "Clear the screen"),
(".clear-history", "Clear the history"),
(".exit", "Exit the REPL"),
(".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"),
];
fn main() {
@@ -146,11 +147,13 @@ fn run_repl(
.with_edit_mode(edit_mode);
let prompt = DefaultPrompt::new(DefaultPromptSegment::Empty, DefaultPromptSegment::Empty);
let mut trigged_ctrlc = false;
+ let mut output = String::new();
let mut role: Option<Role> = None;
let handle_line = |line: String,
line_editor: &mut Reedline,
trigged_ctrlc: &mut bool,
- role: &mut Option<Role>|
+ role: &mut Option<Role>,
+ output: &mut String|
-> Result<bool> {
if line.starts_with('.') {
let (name, args) = match line.split_once(' ') {
@@ -158,11 +161,21 @@ fn run_repl(
None => (line.as_str(), None),
};
match name {
+ ".view" => {
+ if output.is_empty() {
+ return Ok(false);
+ }
+ let _ = Editor::new("view ai reply with an external editor")
+ .with_file_extension(".md")
+ .with_predefined_text(output)
+ .prompt()?;
+ dump("", 1);
+ }
".exit" => {
return Ok(true);
}
".help" => {
- dump(get_repl_help());
+ dump(get_repl_help(), 2);
}
".clear" => {
line_editor.clear_scrollback()?;
@@ -175,40 +188,42 @@ fn run_repl(
}
".history" => {
line_editor.print_history()?;
+ dump("", 1);
}
".role" => match args {
Some(name) => match config.roles.iter().find(|v| v.name == name) {
Some(role_) => {
*role = Some(role_.clone());
}
- None => dump("Unknown role."),
+ None => dump("Unknown role.", 2),
},
- None => dump("Usage: .role <name>."),
+ None => dump("Usage: .role <name>.", 2),
},
_ => {
- dump("Unknown command. Type \".help\" for more information.");
+ dump("Unknown command. Type \".help\" for more information.", 2);
}
}
} else {
- let line = if let Some(role) = role.take() {
+ let input = if let Some(role) = role.take() {
role.generate(&line)
} else {
line
};
+ output.clear();
*trigged_ctrlc = false;
- if line.is_empty() {
+ if input.is_empty() {
return Ok(false);
}
runtime.block_on(async {
tokio::select! {
- ret = handle_input(&client, &config, &line) => {
+ ret = handle_input(&client, &config, &input, output) => {
if let Err(err) = ret {
- dump(format!("error: {err}"));
+ dump(format!("error: {err}"), 2);
}
}
_ = tokio::signal::ctrl_c() => {
*trigged_ctrlc = true;
- dump(" Abort current session.")
+ dump(" Abort current session.", 2)
}
}
});
@@ -221,13 +236,20 @@ fn run_repl(
&mut line_editor,
&mut trigged_ctrlc,
&mut role,
+ &mut output,
)?;
}
loop {
let sig = line_editor.read_line(&prompt);
match sig {
Ok(Signal::Success(line)) => {
- let quit = handle_line(line, &mut line_editor, &mut trigged_ctrlc, &mut role)?;
+ let quit = handle_line(
+ line,
+ &mut line_editor,
+ &mut trigged_ctrlc,
+ &mut role,
+ &mut output,
+ )?;
if quit {
break;
}
@@ -235,7 +257,7 @@ fn run_repl(
Ok(Signal::CtrlC) => {
if !trigged_ctrlc {
trigged_ctrlc = true;
- dump("(To exit, press Ctrl+C again or Ctrl+D or type .exit)");
+ dump("(To exit, press Ctrl+C again or Ctrl+D or type .exit)", 2);
} else {
break;
}
@@ -252,17 +274,24 @@ fn run_repl(
Ok(())
}
-async fn handle_input(client: &Client, config: &Config, text: &str) -> Result<()> {
+async fn handle_input(
+ client: &Client,
+ config: &Config,
+ input: &str,
+ output: &mut String,
+) -> Result<()> {
if config.dry_run {
- dump(text);
+ output.push_str(input);
+ dump(input, 2);
return Ok(());
}
- let mut stream = acquire_stream(client, config, text).await?;
+ let mut stream = acquire_stream(client, config, input).await?;
let mut virgin = true;
while let Some(part) = stream.next().await {
let chunk = part?.data;
if chunk == "[DONE]" {
- dump("\n");
+ output.push('\n');
+ dump("", 2);
break;
} else {
let data: Value = serde_json::from_str(&chunk)?;
@@ -278,8 +307,8 @@ async fn handle_input(client: &Client, config: &Config, text: &str) -> Result<()
continue;
}
}
- print!("{text}");
- stdout().flush().unwrap();
+ output.push_str(text);
+ dump(text, 0);
}
}
Ok(())
@@ -379,8 +408,8 @@ async fn acquire_stream(
Ok(stream)
}
-fn dump<T: ToString>(text: T) {
- println!("{}", text.to_string());
+fn dump<T: ToString>(text: T, newlines: usize) {
+ print!("{}{}", text.to_string(), "\n".repeat(newlines));
stdout().flush().unwrap();
}