From 2539e24fe9eadf69cb20b01d2303a99fa2e8a363 Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 8 Mar 2023 17:13:11 +0800 Subject: chore: improve code by getting prompt from config (#39) --- src/client.rs | 34 ++++++++++------------------------ 1 file changed, 10 insertions(+), 24 deletions(-) (limited to 'src/client.rs') diff --git a/src/client.rs b/src/client.rs index 11a5f3b..862fa0e 100644 --- a/src/client.rs +++ b/src/client.rs @@ -28,9 +28,9 @@ impl ChatGptClient { Ok(s) } - pub fn send_message(&self, input: &str, prompt: Option) -> Result { + pub fn send_message(&self, input: &str) -> Result { self.runtime.block_on(async { - self.send_message_inner(input, prompt) + self.send_message_inner(input) .await .with_context(|| "Failed to send message") }) @@ -39,7 +39,6 @@ impl ChatGptClient { pub fn send_message_streaming( &self, input: &str, - prompt: Option, handler: &mut ReplyStreamHandler, ) -> Result<()> { async fn watch_abort(abort: SharedAbortSignal) { @@ -53,7 +52,7 @@ impl ChatGptClient { let abort = handler.get_abort(); self.runtime.block_on(async { tokio::select! { - ret = self.send_message_streaming_inner(input, prompt, handler) => { + ret = self.send_message_streaming_inner(input, handler) => { handler.done()?; ret.with_context(|| "Failed to send message streaming") } @@ -69,11 +68,11 @@ impl ChatGptClient { }) } - async fn send_message_inner(&self, content: &str, prompt: Option) -> Result { + async fn send_message_inner(&self, content: &str) -> Result { if self.config.borrow().dry_run { - return Ok(combine(content, prompt)); + return Ok(self.config.borrow().merge_prompt(content)); } - let builder = self.request_builder(content, prompt, false)?; + let builder = self.request_builder(content, false)?; let data: Value = builder.send().await?.json().await?; @@ -87,14 +86,13 @@ impl ChatGptClient { async fn send_message_streaming_inner( &self, content: &str, - prompt: Option, handler: &mut ReplyStreamHandler, ) -> Result<()> { if self.config.borrow().dry_run { - handler.text(&combine(content, prompt))?; + handler.text(&self.config.borrow().merge_prompt(content))?; return Ok(()); } - let builder = self.request_builder(content, prompt, true)?; + let builder = self.request_builder(content, true)?; let mut stream = builder.send().await?.bytes_stream().eventsource(); let mut virgin = true; while let Some(part) = stream.next().await { @@ -134,14 +132,9 @@ impl ChatGptClient { Ok(client) } - fn request_builder( - &self, - content: &str, - prompt: Option, - stream: bool, - ) -> Result { + fn request_builder(&self, content: &str, stream: bool) -> Result { let user_message = json!({ "role": "user", "content": content }); - let messages = match prompt { + let messages = match self.config.borrow().get_prompt() { Some(prompt) => { let system_message = json!({ "role": "system", "content": prompt.trim() }); json!([system_message, user_message]) @@ -175,13 +168,6 @@ impl ChatGptClient { } } -fn combine(content: &str, prompt: Option) -> String { - match prompt { - Some(prompt) => format!("{}\n{content}", prompt.trim()), - None => content.to_string(), - } -} - fn init_runtime() -> Result { tokio::runtime::Builder::new_current_thread() .enable_all() -- cgit v1.2.3