diff options
| author | sigoden <sigoden@gmail.com> | 2023-03-09 07:58:44 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-03-09 07:58:44 +0800 |
| commit | 1ec451da896f930829d8908da81e416d891f9aa5 (patch) | |
| tree | a5a99ba939cd67dd735bdd87912bd8cb23926a67 /src | |
| parent | ebd3cb2401739305c6f36c1f7a4cd543a1fdf419 (diff) | |
| download | aichat-1ec451da896f930829d8908da81e416d891f9aa5.tar.gz | |
refactor: replace Arc<Refcell<Config>> with Arc<Mutex<Config>> (#46)
Diffstat (limited to 'src')
| -rw-r--r-- | src/client.rs | 16 | ||||
| -rw-r--r-- | src/config.rs | 5 | ||||
| -rw-r--r-- | src/main.rs | 16 | ||||
| -rw-r--r-- | src/repl/handler.rs | 14 | ||||
| -rw-r--r-- | src/repl/init.rs | 2 |
5 files changed, 27 insertions, 26 deletions
diff --git a/src/client.rs b/src/client.rs index 6d01109..960af0c 100644 --- a/src/client.rs +++ b/src/client.rs @@ -69,8 +69,8 @@ impl ChatGptClient { } async fn send_message_inner(&self, content: &str) -> Result<String> { - if self.config.borrow().dry_run { - return Ok(self.config.borrow().merge_prompt(content)); + if self.config.lock().dry_run { + return Ok(self.config.lock().merge_prompt(content)); } let builder = self.request_builder(content, false)?; @@ -88,8 +88,8 @@ impl ChatGptClient { content: &str, handler: &mut ReplyStreamHandler, ) -> Result<()> { - if self.config.borrow().dry_run { - handler.text(&self.config.borrow().merge_prompt(content))?; + if self.config.lock().dry_run { + handler.text(&self.config.lock().merge_prompt(content))?; return Ok(()); } let builder = self.request_builder(content, true)?; @@ -122,7 +122,7 @@ impl ChatGptClient { fn build_client(&self) -> Result<Client> { let mut builder = Client::builder(); - if let Some(proxy) = self.config.borrow().proxy.as_ref() { + if let Some(proxy) = self.config.lock().proxy.as_ref() { builder = builder.proxy(Proxy::all(proxy).with_context(|| "Invalid config.proxy")?); } let client = builder @@ -134,7 +134,7 @@ impl ChatGptClient { fn request_builder(&self, content: &str, stream: bool) -> Result<RequestBuilder> { let user_message = json!({ "role": "user", "content": content }); - let messages = match self.config.borrow().get_prompt() { + let messages = match self.config.lock().get_prompt() { Some(prompt) => { let system_message = json!({ "role": "system", "content": prompt.trim() }); json!([system_message, user_message]) @@ -148,7 +148,7 @@ impl ChatGptClient { "messages": messages, }); - if let Some(v) = self.config.borrow().get_temperature() { + if let Some(v) = self.config.lock().get_temperature() { body.as_object_mut() .and_then(|m| m.insert("temperature".into(), json!(v))); } @@ -161,7 +161,7 @@ impl ChatGptClient { let builder = self .build_client()? .post(API_URL) - .bearer_auth(&self.config.borrow().api_key) + .bearer_auth(&self.config.lock().api_key) .json(&body); Ok(builder) diff --git a/src/config.rs b/src/config.rs index 1542f9e..fe93da6 100644 --- a/src/config.rs +++ b/src/config.rs @@ -1,5 +1,4 @@ use std::{ - cell::RefCell, env, fs::{create_dir_all, read_to_string, File, OpenOptions}, io::Write, @@ -8,6 +7,8 @@ use std::{ sync::Arc, }; +use parking_lot::Mutex; + use anyhow::{anyhow, Context, Result}; use inquire::{Confirm, Text}; use serde::{Deserialize, Serialize}; @@ -56,7 +57,7 @@ pub struct Config { pub role: Option<Role>, } -pub type SharedConfig = Arc<RefCell<Config>>; +pub type SharedConfig = Arc<Mutex<Config>>; impl Config { pub fn init(is_interactive: bool) -> Result<Config> { diff --git a/src/main.rs b/src/main.rs index a25ab74..a6662e7 100644 --- a/src/main.rs +++ b/src/main.rs @@ -7,7 +7,6 @@ mod term; #[macro_use] mod utils; -use std::cell::RefCell; use std::io::{stdin, Read}; use std::sync::Arc; use std::{io::stdout, process::exit}; @@ -17,6 +16,7 @@ use client::ChatGptClient; use config::{Config, SharedConfig}; use crossbeam::sync::WaitGroup; use is_terminal::IsTerminal; +use parking_lot::Mutex; use anyhow::{anyhow, Result}; use clap::Parser; @@ -26,10 +26,10 @@ use repl::{AbortSignal, Repl}; fn main() -> Result<()> { let cli = Cli::parse(); let text = cli.text(); - let config = Arc::new(RefCell::new(Config::init(text.is_none())?)); + let config = Arc::new(Mutex::new(Config::init(text.is_none())?)); if cli.list_roles { config - .borrow() + .lock() .roles .iter() .for_each(|v| println!("{}", v.name)); @@ -38,15 +38,15 @@ fn main() -> Result<()> { let role = match &cli.role { Some(name) => Some( config - .borrow() + .lock() .find_role(name) .ok_or_else(|| anyhow!("Unknown role '{name}'"))?, ), None => None, }; - config.borrow_mut().role = role; + config.lock().role = role; if cli.no_highlight { - config.borrow_mut().highlight = false; + config.lock().highlight = false; } let no_stream = cli.no_stream; let client = ChatGptClient::init(config.clone())?; @@ -71,7 +71,7 @@ fn start_directive( input: &str, no_stream: bool, ) -> Result<()> { - let highlight = config.borrow().highlight && stdout().is_terminal(); + let highlight = config.lock().highlight && stdout().is_terminal(); let output = if no_stream { let output = client.send_message(input)?; if highlight { @@ -93,7 +93,7 @@ fn start_directive( wg.wait(); output }; - config.borrow().save_message(input, &output) + config.lock().save_message(input, &output) } fn start_interactive(client: ChatGptClient, config: SharedConfig) -> Result<()> { diff --git a/src/repl/handler.rs b/src/repl/handler.rs index 22d8fd0..6a7fc93 100644 --- a/src/repl/handler.rs +++ b/src/repl/handler.rs @@ -48,7 +48,7 @@ impl ReplCmdHandler { self.reply.borrow_mut().clear(); return Ok(()); } - let highlight = self.config.borrow().highlight; + let highlight = self.config.lock().highlight; let wg = WaitGroup::new(); let ret = render_stream( &input, @@ -60,27 +60,27 @@ impl ReplCmdHandler { ); wg.wait(); let buffer = ret?; - self.config.borrow().save_message(&input, &buffer)?; + self.config.lock().save_message(&input, &buffer)?; *self.reply.borrow_mut() = buffer; } ReplCmd::SetRole(name) => { - let output = self.config.borrow_mut().change_role(&name); + let output = self.config.lock().change_role(&name); print_now!("{}\n\n", output.trim_end()); } ReplCmd::ClearRole => { - self.config.borrow_mut().role = None; + self.config.lock().role = None; print_now!("\n"); } ReplCmd::Prompt(prompt) => { - self.config.borrow_mut().create_temp_role(&prompt); + self.config.lock().create_temp_role(&prompt); print_now!("\n"); } ReplCmd::Info => { - let output = self.config.borrow().info()?; + let output = self.config.lock().info()?; print_now!("{}\n\n", output.trim_end()); } ReplCmd::UpdateConfig(input) => { - let output = self.config.borrow_mut().update(&input)?; + let output = self.config.lock().update(&input)?; let output = output.trim(); if output.is_empty() { print_now!("\n"); diff --git a/src/repl/init.rs b/src/repl/init.rs index f640deb..88f66fa 100644 --- a/src/repl/init.rs +++ b/src/repl/init.rs @@ -47,7 +47,7 @@ impl Repl { .into_iter() .map(|(v, _, _)| v.to_string()) .collect(); - completion.extend(config.borrow().repl_completions()); + completion.extend(config.lock().repl_completions()); let mut completer = DefaultCompleter::with_inclusions(&['.', '-', '_']).set_min_word_len(2); completer.insert(completion.clone()); completer |
