From 52a847743efccb72bcd031cb1d419c84ed9d9afa Mon Sep 17 00:00:00 2001 From: sigoden Date: Sun, 23 Jun 2024 06:00:58 +0800 Subject: refactor: improve system message handling (#634) --- src/client/ernie.rs | 6 +++++- src/client/message.rs | 30 ++++++++++++++++++++++++++---- src/client/vertexai.rs | 11 ++++++++++- 3 files changed, 41 insertions(+), 6 deletions(-) (limited to 'src') diff --git a/src/client/ernie.rs b/src/client/ernie.rs index 21362c1..428d263 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -238,7 +238,7 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Valu stream, } = data; - patch_system_message(&mut messages); + let system_message = extract_system_message(&mut messages); let messages: Vec = messages .into_iter() @@ -269,6 +269,10 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Valu "messages": messages, }); + if let Some(v) = system_message { + body["system"] = v.into(); + } + if let Some(v) = model.max_tokens_param() { body["max_output_tokens"] = v.into(); } 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) { 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); } } } diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index 16910b6..cc7e464 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -262,7 +262,12 @@ pub fn gemini_build_chat_completions_body( stream: _, } = data; - patch_system_message(&mut messages); + let system_message = if model.name().starts_with("gemini-1.5-") { + extract_system_message(&mut messages) + } else { + patch_system_message(&mut messages); + None + }; let mut network_image_urls = vec![]; let contents: Vec = messages @@ -333,6 +338,10 @@ pub fn gemini_build_chat_completions_body( let mut body = json!({ "contents": contents, "generationConfig": {} }); + if let Some(v) = system_message { + body["systemInstruction"] = json!({ "parts": [{"text": v }] }); + } + if let Some(v) = model.max_tokens_param() { body["generationConfig"]["maxOutputTokens"] = v.into(); } -- cgit v1.2.3