summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-03-05 22:51:29 +0800
committerGitHub <noreply@github.com>2023-03-05 22:51:29 +0800
commit4b1d6c16b31cf2605fc928d4f783eb4c49d643ea (patch)
tree074970f34280d8b8590df74ede9ed57c307b55cf
parent957ea431c23dd6cbaa512e6a67b4a25e2adf08e4 (diff)
downloadaichat-4b1d6c16b31cf2605fc928d4f783eb4c49d643ea.tar.gz
feat: add `.set` command (#20)
* feat: add `.set` command * Add config.role
-rw-r--r--Cargo.toml2
-rw-r--r--src/client.rs58
-rw-r--r--src/config.rs152
-rw-r--r--src/main.rs39
-rw-r--r--src/repl.rs125
5 files changed, 232 insertions, 144 deletions
diff --git a/Cargo.toml b/Cargo.toml
index 9098634..6d2403c 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -27,7 +27,7 @@ serde = { version = "1.0.152", features = ["derive"] }
serde_json = "1.0.93"
serde_yaml = "0.9.17"
tokio = { version = "1.26.0", features = ["full"] }
-mdcat = { version = "1.1.0", default_features = false, features =["static"] }
+mdcat = { version = "1.1.0", default-features = false, features =["static"] }
pulldown-cmark = { version = "0.9.2", default-features = false, features = ['simd'] }
crossbeam = "0.8.2"
crossterm = "0.26.1"
diff --git a/src/client.rs b/src/client.rs
index f2aef38..698fffc 100644
--- a/src/client.rs
+++ b/src/client.rs
@@ -1,4 +1,4 @@
-use crate::config::Config;
+use crate::config::SharedConfig;
use crate::repl::ReplyReceiver;
use anyhow::{anyhow, Context, Result};
@@ -17,28 +17,16 @@ const MODEL: &str = "gpt-3.5-turbo";
#[derive(Debug)]
pub struct ChatGptClient {
- client: Client,
- config: Arc<Config>,
+ config: SharedConfig,
runtime: Runtime,
}
impl ChatGptClient {
- pub fn init(config: Arc<Config>) -> Result<Self> {
- let mut builder = Client::builder();
- if let Some(proxy) = config.proxy.as_ref() {
- builder = builder.proxy(Proxy::all(proxy).with_context(|| "Invalid config.proxy")?);
- }
- let client = builder
- .connect_timeout(CONNECT_TIMEOUT)
- .build()
- .with_context(|| "Failed to init http client")?;
-
+ pub fn init(config: SharedConfig) -> Result<Self> {
let runtime = init_runtime()?;
- Ok(Self {
- client,
- config,
- runtime,
- })
+ let s = Self { config, runtime };
+ let _ = s.build_client()?; // check error
+ Ok(s)
}
pub fn acquire(&self, input: &str, prompt: Option<String>) -> Result<String> {
@@ -80,10 +68,10 @@ impl ChatGptClient {
}
async fn acquire_inner(&self, content: &str, prompt: Option<String>) -> Result<String> {
- if self.config.dry_run {
+ if self.config.borrow().dry_run {
return Ok(combine(content, prompt));
}
- let builder = self.request_builder(content, prompt, false);
+ let builder = self.request_builder(content, prompt, false)?;
let data: Value = builder.send().await?.json().await?;
@@ -100,11 +88,11 @@ impl ChatGptClient {
prompt: Option<String>,
receiver: &mut ReplyReceiver,
) -> Result<()> {
- if self.config.dry_run {
+ if self.config.borrow().dry_run {
receiver.text(&combine(content, prompt));
return Ok(());
}
- let builder = self.request_builder(content, prompt, true);
+ let builder = self.request_builder(content, prompt, true)?;
let mut stream = builder.send().await?.bytes_stream().eventsource();
let mut virgin = true;
while let Some(part) = stream.next().await {
@@ -132,12 +120,24 @@ impl ChatGptClient {
Ok(())
}
+ fn build_client(&self) -> Result<Client> {
+ let mut builder = Client::builder();
+ if let Some(proxy) = self.config.borrow().proxy.as_ref() {
+ builder = builder.proxy(Proxy::all(proxy).with_context(|| "Invalid config.proxy")?);
+ }
+ let client = builder
+ .connect_timeout(CONNECT_TIMEOUT)
+ .build()
+ .with_context(|| "Failed to build http client")?;
+ Ok(client)
+ }
+
fn request_builder(
&self,
content: &str,
prompt: Option<String>,
stream: bool,
- ) -> RequestBuilder {
+ ) -> Result<RequestBuilder> {
let user_message = json!({ "role": "user", "content": content });
let messages = match prompt {
Some(prompt) => {
@@ -153,7 +153,7 @@ impl ChatGptClient {
"messages": messages,
});
- if let Some(v) = self.config.temperature {
+ if let Some(v) = self.config.borrow().temperature {
body.as_object_mut()
.and_then(|m| m.insert("temperature".into(), json!(v)));
}
@@ -163,10 +163,13 @@ impl ChatGptClient {
.and_then(|m| m.insert("stream".into(), json!(true)));
}
- self.client
+ let builder = self
+ .build_client()?
.post(API_URL)
- .bearer_auth(&self.config.api_key)
- .json(&body)
+ .bearer_auth(&self.config.borrow().api_key)
+ .json(&body);
+
+ Ok(builder)
}
}
@@ -176,6 +179,7 @@ fn combine(content: &str, prompt: Option<String>) -> String {
None => content.to_string(),
}
}
+
fn init_runtime() -> Result<Runtime> {
tokio::runtime::Builder::new_current_thread()
.enable_all()
diff --git a/src/config.rs b/src/config.rs
index 802f376..2f0a361 100644
--- a/src/config.rs
+++ b/src/config.rs
@@ -1,9 +1,11 @@
use std::{
+ cell::RefCell,
env,
fs::{create_dir_all, read_to_string, File, OpenOptions},
io::Write,
path::{Path, PathBuf},
process::exit,
+ sync::Arc,
};
use anyhow::{anyhow, Context, Result};
@@ -35,11 +37,24 @@ pub struct Config {
#[serde(default)]
pub dry_run: bool,
/// Predefined roles
- #[serde(default, skip_serializing)]
+ #[serde(default, skip)]
pub roles: Vec<Role>,
+ /// Current selected role
+ #[serde(default, skip)]
+ pub role: Option<Role>,
}
+pub type SharedConfig = Arc<RefCell<Config>>;
+
impl Config {
+ pub const UPDATE_KEYS: [&str; 6] = [
+ "api_key",
+ "temperature",
+ "save",
+ "highlight",
+ "proxy",
+ "dry_run",
+ ];
pub fn init(is_interactive: bool) -> Result<Config> {
let config_path = Config::config_file()?;
if is_interactive && !config_path.exists() {
@@ -99,18 +114,17 @@ impl Config {
Ok(file)
}
- pub fn save_message(
- file: Option<&mut File>,
- input: &str,
- output: &str,
- role_name: &Option<String>,
- ) {
- let role_name = match role_name {
- Some(v) => format!("({v})"),
- None => String::new(),
- };
- let timestamp = format!("[{}]", now());
- if let (false, Some(file)) = (output.is_empty(), file) {
+ pub fn save_message(&self, file: Option<&mut File>, input: &str, output: &str) {
+ if output.is_empty() || !self.save {
+ return;
+ }
+ if let Some(file) = file {
+ let role_name = self
+ .role
+ .as_ref()
+ .map(|v| format!("({})", v.name))
+ .unwrap_or_default();
+ let timestamp = format!("[{}]", now());
let _ = file.write_all(
format!(
"# CHAT:{timestamp} {role_name}\n{}\n\n--------\n{}\n--------\n\n",
@@ -138,6 +152,118 @@ impl Config {
Self::local_file(MESSAGE_FILE_NAME)
}
+ pub fn change_role(&mut self, name: &str) -> String {
+ match self.find_role(name) {
+ Some(role) => {
+ let output = format!("{}>> {}", role.name, role.prompt.trim());
+ self.role = Some(role);
+ output
+ }
+ None => "Unknown role".into(),
+ }
+ }
+
+ pub fn get_prompt(&self) -> Option<String> {
+ self.role.as_ref().and_then(|v| {
+ if v.prompt.is_empty() {
+ None
+ } else {
+ Some(v.prompt.to_string())
+ }
+ })
+ }
+
+ pub fn info(&self) -> Result<String> {
+ let file_info = |path: &Path| {
+ let state = if path.exists() { "" } else { " ⚠️" };
+ format!("{}{state}", path.display())
+ };
+ let proxy = self
+ .proxy
+ .as_ref()
+ .map(|v| v.to_string())
+ .unwrap_or("-".into());
+ let temperature = self
+ .temperature
+ .map(|v| v.to_string())
+ .unwrap_or("-".into());
+ let role_name = self
+ .role
+ .as_ref()
+ .map(|v| v.name.to_string())
+ .unwrap_or("-".into());
+ let items = vec![
+ ("config_file", file_info(&Config::config_file()?)),
+ ("roles_file", file_info(&Config::roles_file()?)),
+ ("messages_file", file_info(&Config::messages_file()?)),
+ ("role", role_name),
+ ("api_key", self.api_key.clone()),
+ ("temperature", temperature),
+ ("save", self.save.to_string()),
+ ("highlight", self.highlight.to_string()),
+ ("proxy", proxy),
+ ("dry_run", self.dry_run.to_string()),
+ ];
+ let mut output = String::new();
+ for (name, value) in items {
+ output.push_str(&format!("{name:<20}{value}\n"));
+ }
+ Ok(output)
+ }
+
+ pub fn update(&mut self, data: &str) -> Result<String> {
+ let parts: Vec<&str> = data.split_whitespace().collect();
+ if parts.len() != 2 {
+ return Ok("Usage: .set <key> <value>. If value is null, unset key.".into());
+ }
+ let key = parts[0];
+ let value = parts[1];
+ let unset = value == "null";
+ match key {
+ "api_key" => {
+ if unset {
+ return Ok("Not allowd".into());
+ } else {
+ self.api_key = value.to_string();
+ }
+ }
+ "temperature" => {
+ if unset {
+ self.temperature = None;
+ } else {
+ let value = value.parse().with_context(|| "Invalid value")?;
+ self.temperature = Some(value);
+ }
+ }
+ "save" => {
+ let value = value.parse().with_context(|| "Invalid value")?;
+ self.save = value;
+ }
+ "highlight" => {
+ let value = value.parse().with_context(|| "Invalid value")?;
+ self.highlight = value;
+ }
+ "proxy" => {
+ if unset {
+ self.proxy = None;
+ } else {
+ self.proxy = Some(value.to_string());
+ }
+ }
+ "dry_run" => {
+ let value = value.parse().with_context(|| "Invalid value")?;
+ self.dry_run = value;
+ }
+ _ => {
+ return Ok(format!(
+ "Unknown key, valid keys are {}",
+ Config::UPDATE_KEYS.join(", ")
+ ))
+ }
+ }
+ Ok("Done".into())
+ }
+
fn load_roles(&mut self) -> Result<()> {
let path = Self::roles_file()?;
if !path.exists() {
diff --git a/src/main.rs b/src/main.rs
index 4f5198c..39a49e4 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -6,13 +6,14 @@ mod repl;
mod term;
mod utils;
+use std::cell::RefCell;
use std::io::{stdin, Read};
use std::sync::Arc;
use std::{io::stdout, process::exit};
use cli::Cli;
use client::ChatGptClient;
-use config::{Config, Role};
+use config::{Config, SharedConfig};
use is_terminal::IsTerminal;
use anyhow::{anyhow, Result};
@@ -23,56 +24,56 @@ use repl::{Repl, ReplCmdHandler};
fn main() -> Result<()> {
let cli = Cli::parse();
let text = cli.text();
- let config = Arc::new(Config::init(text.is_none())?);
+ let config = Arc::new(RefCell::new(Config::init(text.is_none())?));
if cli.list_roles {
- config.roles.iter().for_each(|v| println!("{}", v.name));
+ config
+ .borrow()
+ .roles
+ .iter()
+ .for_each(|v| println!("{}", v.name));
exit(0);
}
let role = match &cli.role {
Some(name) => Some(
config
+ .borrow()
.find_role(name)
.ok_or_else(|| anyhow!("Unknown role '{name}'"))?,
),
None => None,
};
+ config.borrow_mut().role = role;
let client = ChatGptClient::init(config.clone())?;
if atty::isnt(atty::Stream::Stdin) {
let mut text = String::new();
stdin().read_to_string(&mut text)?;
- start_directive(client, config, role, &text)
+ start_directive(client, config, &text)
} else {
match text {
- Some(text) => start_directive(client, config, role, &text),
- None => start_interactive(client, config, role),
+ Some(text) => start_directive(client, config, &text),
+ None => start_interactive(client, config),
}
}
}
-fn start_directive(
- client: ChatGptClient,
- config: Arc<Config>,
- role: Option<Role>,
- input: &str,
-) -> Result<()> {
- let mut file = config.open_message_file()?;
- let prompt = role.as_ref().map(|v| v.prompt.to_string());
- let role_name = role.as_ref().map(|v| v.name.to_string());
+fn start_directive(client: ChatGptClient, config: SharedConfig, input: &str) -> Result<()> {
+ let mut file = config.borrow().open_message_file()?;
+ let prompt = config.borrow().get_prompt();
let output = client.acquire(input, prompt)?;
let output = output.trim();
- if config.highlight && stdout().is_terminal() {
+ if config.borrow().highlight && stdout().is_terminal() {
let markdown_render = MarkdownRender::init()?;
markdown_render.print(output)?;
} else {
println!("{output}");
}
- Config::save_message(file.as_mut(), input, output, &role_name);
+ config.borrow().save_message(file.as_mut(), input, output);
Ok(())
}
-fn start_interactive(client: ChatGptClient, config: Arc<Config>, role: Option<Role>) -> Result<()> {
+fn start_interactive(client: ChatGptClient, config: SharedConfig) -> Result<()> {
let mut repl = Repl::init(config.clone())?;
- let handler = ReplCmdHandler::init(client, config, role)?;
+ let handler = ReplCmdHandler::init(client, config)?;
repl.run(handler)
}
diff --git a/src/repl.rs b/src/repl.rs
index 98f43f1..3c101b7 100644
--- a/src/repl.rs
+++ b/src/repl.rs
@@ -1,5 +1,5 @@
use crate::client::ChatGptClient;
-use crate::config::{Config, Role};
+use crate::config::{Config, SharedConfig};
use crate::render::{self, MarkdownRender};
use crate::term;
use crate::utils::{copy, dump};
@@ -13,12 +13,11 @@ use reedline::{
};
use std::cell::RefCell;
use std::fs::File;
-use std::path::Path;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::thread::spawn;
-const REPL_COMMANDS: [(&str, &str); 10] = [
+const REPL_COMMANDS: [(&str, &str); 11] = [
(".role", "Specifies the role the AI will play"),
(".clear role", "Clear the currently selected role"),
(".history", "Print the history"),
@@ -26,6 +25,7 @@ const REPL_COMMANDS: [(&str, &str); 10] = [
(".multiline", "Enter multiline editor mode"),
(".copy", "Copy last reply message"),
(".info", "Print the information"),
+ (".set", "Modify the configuration temporarily"),
(".help", "Print this help message"),
(".exit", "Exit the REPL"),
(".clear screen", "Clear the screen"),
@@ -39,7 +39,7 @@ pub struct Repl {
}
impl Repl {
- pub fn init(config: Arc<Config>) -> Result<Self> {
+ pub fn init(config: SharedConfig) -> Result<Self> {
let completer = Self::create_completer(config);
let keybindings = Self::create_keybindings();
let history = Self::create_history()?;
@@ -154,6 +154,9 @@ impl Repl {
dump("Copied", 1);
}
}
+ ".set" => {
+ handler.handle(ReplCmd::UpdateConfig(args.unwrap_or_default().to_string()))?
+ }
_ => dump_unknown_command(),
}
} else {
@@ -167,13 +170,21 @@ impl Repl {
DefaultPrompt::new(DefaultPromptSegment::Empty, DefaultPromptSegment::Empty)
}
- fn create_completer(config: Arc<Config>) -> DefaultCompleter {
+ fn create_completer(config: SharedConfig) -> DefaultCompleter {
let mut commands: Vec<String> = REPL_COMMANDS
.into_iter()
.map(|(v, _)| v.to_string())
.collect();
- commands.extend(config.roles.iter().map(|v| format!(".role {}", v.name)));
- let mut completer = DefaultCompleter::with_inclusions(&['.', '-']).set_min_word_len(2);
+ commands.extend(
+ config
+ .as_ref()
+ .borrow()
+ .roles
+ .iter()
+ .map(|v| format!(".role {}", v.name)),
+ );
+ commands.extend(Config::UPDATE_KEYS.map(|v| format!(".set {v}")));
+ let mut completer = DefaultCompleter::with_inclusions(&['.', '-', '_']).set_min_word_len(2);
completer.insert(commands.clone());
completer
}
@@ -243,29 +254,23 @@ fn incomplete_brackets(line: &str) -> bool {
pub struct ReplCmdHandler {
client: ChatGptClient,
- config: Arc<Config>,
+ config: SharedConfig,
state: RefCell<ReplCmdHandlerState>,
ctrlc: Arc<AtomicBool>,
- render: Option<Arc<MarkdownRender>>,
+ render: Arc<MarkdownRender>,
}
struct ReplCmdHandlerState {
reply: String,
- role: Option<Role>,
save_file: Option<File>,
}
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 save_file = config.open_message_file()?;
+ pub fn init(client: ChatGptClient, config: SharedConfig) -> Result<Self> {
+ let render = Arc::new(MarkdownRender::init()?);
+ let save_file = config.as_ref().borrow().open_message_file()?;
let ctrlc = Arc::new(AtomicBool::new(false));
let state = RefCell::new(ReplCmdHandlerState {
- role,
save_file,
reply: String::new(),
});
@@ -284,25 +289,16 @@ impl ReplCmdHandler {
self.state.borrow_mut().reply.clear();
return Ok(());
}
- let prompt = self
- .state
- .borrow()
- .role
- .as_ref()
- .map(|v| v.prompt.to_string())
- .unwrap_or_default();
- let prompt = if prompt.is_empty() {
- None
- } else {
- Some(prompt)
- };
+ let prompt = self.config.borrow().get_prompt();
let wg = WaitGroup::new();
- let mut receiver = if let Some(markdown_render) = self.render.clone() {
+ let highlight = self.config.borrow().highlight;
+ let mut receiver = if highlight {
let (tx, rx) = unbounded();
let ctrlc = self.ctrlc.clone();
let wg = wg.clone();
+ let render = self.render.clone();
spawn(move || {
- let _ = render::render_stream(rx, ctrlc, markdown_render);
+ let _ = render::render_stream(rx, ctrlc, render);
drop(wg);
});
ReplyReceiver::new(Some(tx))
@@ -311,69 +307,29 @@ impl ReplCmdHandler {
};
self.client
.acquire_stream(&input, prompt, &mut receiver, self.ctrlc.clone())?;
- let role = self
- .state
- .borrow_mut()
- .role
- .as_ref()
- .map(|v| v.name.to_string());
- Config::save_message(
+ self.config.borrow().save_message(
self.state.borrow_mut().save_file.as_mut(),
&input,
&receiver.output,
- &role,
);
wg.wait();
self.state.borrow_mut().reply = receiver.output;
}
- ReplCmd::SetRole(name) => match self.config.find_role(&name) {
- Some(role) => {
- let output = format!("{}>> {}", role.name, role.prompt.trim());
- self.state.borrow_mut().role = Some(role);
- dump(output, 2);
- }
- None => {
- dump("Unknown role", 2);
- }
- },
+ ReplCmd::SetRole(name) => {
+ let output = self.config.borrow_mut().change_role(&name);
+ dump(output.trim(), 2);
+ }
ReplCmd::ClearRole => {
- self.state.borrow_mut().role = None;
+ self.config.borrow_mut().role = None;
dump("Done", 2);
}
ReplCmd::Info => {
- let state = self.state.borrow();
- let file_info = |path: &Path| {
- let state = if path.exists() { "" } else { " [not found]" };
- format!("{}{state}", path.display())
- };
- let items = vec![
- ("config file", file_info(&Config::config_file()?)),
- ("roles file", file_info(&Config::roles_file()?)),
- ("messages file", file_info(&Config::messages_file()?)),
- (
- "current role",
- state
- .role
- .as_ref()
- .map(|v| v.name.to_string())
- .unwrap_or_default(),
- ),
- (
- "proxy",
- self.config
- .proxy
- .as_ref()
- .map(|v| v.to_string())
- .unwrap_or_default(),
- ),
- ("save messages", self.config.save.to_string()),
- ("highlight", (self.config.highlight).to_string()),
- ];
- let mut info = String::new();
- for (name, value) in items {
- info.push_str(&format!("{name:<20}{value}\n"));
- }
- dump(info, 1);
+ let output = self.config.borrow().info()?;
+ dump(output.trim(), 2);
+ }
+ ReplCmd::UpdateConfig(input) => {
+ let output = self.config.borrow_mut().update(&input)?;
+ dump(output.trim(), 2);
}
}
Ok(())
@@ -429,6 +385,7 @@ pub enum RenderStreamEvent {
enum ReplCmd {
Submit(String),
SetRole(String),
+ UpdateConfig(String),
ClearRole,
Info,
}