diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/client.rs | 58 | ||||
| -rw-r--r-- | src/config.rs | 152 | ||||
| -rw-r--r-- | src/main.rs | 39 | ||||
| -rw-r--r-- | src/repl.rs | 125 |
4 files changed, 231 insertions, 143 deletions
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, } |
