diff options
| author | sigoden <sigoden@gmail.com> | 2023-03-05 07:52:36 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-03-05 07:52:36 +0800 |
| commit | 598c31c3d946e8772c4d1bc450b5d2b1bdf587c0 (patch) | |
| tree | df12513b90f3b559b25f93a9f9d26f54e58d1118 /src/client.rs | |
| parent | 822f0eed9429f71524e739352da3f76d604071d5 (diff) | |
| download | aichat-598c31c3d946e8772c4d1bc450b5d2b1bdf587c0.tar.gz | |
refactor: optimize role prompt (#13)
Diffstat (limited to 'src/client.rs')
| -rw-r--r-- | src/client.rs | 32 |
1 files changed, 22 insertions, 10 deletions
diff --git a/src/client.rs b/src/client.rs index 1f42346..3fe75dd 100644 --- a/src/client.rs +++ b/src/client.rs @@ -81,11 +81,10 @@ impl ChatGptClient { } async fn acquire_inner(&self, content: &str, prompt: Option<String>) -> Result<String> { - let content = combine(content, prompt); if self.config.dry_run { - return Ok(content); + return Ok(combine(content, prompt)); } - let builder = self.request_builder(&content, false); + let builder = self.request_builder(content, prompt, false); let data: Value = builder.send().await?.json().await?; @@ -102,12 +101,11 @@ impl ChatGptClient { prompt: Option<String>, receiver: &mut ReplyReceiver, ) -> Result<()> { - let content = combine(content, prompt); if self.config.dry_run { - receiver.text(&content); + receiver.text(&combine(content, prompt)); return Ok(()); } - let builder = self.request_builder(&content, 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 { @@ -135,10 +133,25 @@ impl ChatGptClient { Ok(()) } - fn request_builder(&self, content: &str, stream: bool) -> RequestBuilder { + fn request_builder( + &self, + content: &str, + prompt: Option<String>, + stream: bool, + ) -> RequestBuilder { + let user_message = json!({ "role": "user", "content": content }); + let messages = match prompt { + Some(prompt) => { + let system_message = json!({ "role": "system", "content": prompt.trim() }); + json!([system_message, user_message]) + } + None => { + json!([user_message]) + } + }; let mut body = json!({ "model": MODEL, - "messages": [{"role": "user", "content": content}], + "messages": messages, }); if let Some(v) = self.config.temperature { @@ -160,11 +173,10 @@ impl ChatGptClient { fn combine(content: &str, prompt: Option<String>) -> String { match prompt { - Some(v) => format!("{v} {content}"), + Some(prompt) => format!("{}\n{content}", prompt.trim()), None => content.to_string(), } } - fn init_runtime() -> Result<Runtime> { tokio::runtime::Builder::new_current_thread() .enable_all() |
