From 4b1d6c16b31cf2605fc928d4f783eb4c49d643ea Mon Sep 17 00:00:00 2001 From: sigoden Date: Sun, 5 Mar 2023 22:51:29 +0800 Subject: feat: add `.set` command (#20) * feat: add `.set` command * Add config.role --- src/client.rs | 58 +++++++++++++++++++++++++++++++--------------------------- 1 file changed, 31 insertions(+), 27 deletions(-) (limited to 'src/client.rs') 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: SharedConfig, runtime: Runtime, } impl ChatGptClient { - pub fn init(config: Arc) -> Result { - 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 { 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) -> Result { @@ -80,10 +68,10 @@ impl ChatGptClient { } async fn acquire_inner(&self, content: &str, prompt: Option) -> Result { - 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, 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 { + 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, stream: bool, - ) -> RequestBuilder { + ) -> Result { 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 { None => content.to_string(), } } + fn init_runtime() -> Result { tokio::runtime::Builder::new_current_thread() .enable_all() -- cgit v1.2.3