summaryrefslogtreecommitdiffstats
path: root/src/client/message.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-23 06:00:58 +0800
committerGitHub <noreply@github.com>2024-06-23 06:00:58 +0800
commit52a847743efccb72bcd031cb1d419c84ed9d9afa (patch)
tree859c93d356fb2455f39b303b75981fb44a054095 /src/client/message.rs
parent3826d808d81c3ddc996a987bf6a749e8df6d8c48 (diff)
downloadaichat-52a847743efccb72bcd031cb1d419c84ed9d9afa.tar.gz
refactor: improve system message handling (#634)
Diffstat (limited to 'src/client/message.rs')
-rw-r--r--src/client/message.rs30
1 files changed, 26 insertions, 4 deletions
diff --git a/src/client/message.rs b/src/client/message.rs
index d7ba698..77adf47 100644
--- a/src/client/message.rs
+++ b/src/client/message.rs
@@ -21,6 +21,30 @@ impl Message {
pub fn new(role: MessageRole, content: MessageContent) -> Self {
Self { role, content }
}
+
+ pub fn merge_system(&mut self, system: &str) {
+ match &mut self.content {
+ MessageContent::Text(text) => {
+ self.content = MessageContent::Array(vec![
+ MessageContentPart::Text {
+ text: system.to_string(),
+ },
+ MessageContentPart::Text {
+ text: text.to_string(),
+ },
+ ]);
+ }
+ MessageContent::Array(list) => {
+ list.insert(
+ 0,
+ MessageContentPart::Text {
+ text: system.to_string(),
+ },
+ );
+ }
+ _ => {}
+ }
+ }
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize)]
@@ -124,12 +148,10 @@ pub struct ImageUrl {
pub fn patch_system_message(messages: &mut Vec<Message>) {
if messages[0].role.is_system() {
let system_message = messages.remove(0);
- if let (Some(message), MessageContent::Text(system_text)) =
+ if let (Some(message), MessageContent::Text(system)) =
(messages.get_mut(0), system_message.content)
{
- if let MessageContent::Text(text) = message.content.clone() {
- message.content = MessageContent::Text(format!("{}\n\n{}", system_text, text))
- }
+ message.merge_system(&system);
}
}
}