From 69974365ee786af5b643c335e2a12dd7817b84be Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 10 Feb 2025 14:57:08 +0800 Subject: feat: new model field `system_prompt_prefix` (#1163) --- src/client/message.rs | 59 ++++++++++++++++++++++++++++++++++----------------- src/client/model.rs | 6 ++++++ 2 files changed, 45 insertions(+), 20 deletions(-) (limited to 'src/client') diff --git a/src/client/message.rs b/src/client/message.rs index 2f7517c..5ca7f78 100644 --- a/src/client/message.rs +++ b/src/client/message.rs @@ -1,3 +1,5 @@ +use super::Model; + use crate::{function::ToolResult, multiline_text, utils::dimmed_text}; use serde::{Deserialize, Serialize}; @@ -22,25 +24,28 @@ impl Message { Self { role, content } } - pub fn merge_system(&mut self, system: &str) { - match &mut self.content { - MessageContent::Text(text) => { + pub fn merge_system(&mut self, system: MessageContent) { + match (&mut self.content, system) { + (MessageContent::Text(text), MessageContent::Text(system_text)) => { self.content = MessageContent::Array(vec![ - MessageContentPart::Text { - text: system.to_string(), - }, + MessageContentPart::Text { text: system_text }, MessageContentPart::Text { text: text.to_string(), }, - ]); + ]) } - MessageContent::Array(list) => { - list.insert( - 0, - MessageContentPart::Text { - text: system.to_string(), - }, - ); + (MessageContent::Array(list), MessageContent::Text(system_text)) => { + list.insert(0, MessageContentPart::Text { text: system_text }) + } + (MessageContent::Text(text), MessageContent::Array(mut system_list)) => { + system_list.push(MessageContentPart::Text { + text: text.to_string(), + }); + self.content = MessageContent::Array(system_list); + } + (MessageContent::Array(list), MessageContent::Array(mut system_list)) => { + system_list.append(list); + self.content = MessageContent::Array(system_list); } _ => {} } @@ -196,13 +201,27 @@ impl MessageContentToolCalls { } } -pub fn patch_system_message(messages: &mut Vec) { - if messages[0].role.is_system() { +pub fn patch_messages(messages: &mut Vec, model: &Model) { + if messages.is_empty() { + return; + } + if let Some(prefix) = model.system_prompt_prefix() { + if messages[0].role.is_system() { + messages[0].merge_system(MessageContent::Text(prefix.to_string())); + } else { + messages.insert( + 0, + Message { + role: MessageRole::System, + content: MessageContent::Text(prefix.to_string()), + }, + ); + } + } + if model.no_system_message() && messages[0].role.is_system() { let system_message = messages.remove(0); - if let (Some(message), MessageContent::Text(system)) = - (messages.get_mut(0), system_message.content) - { - message.merge_system(&system); + if let (Some(message), system) = (messages.get_mut(0), system_message.content) { + message.merge_system(system); } } } diff --git a/src/client/model.rs b/src/client/model.rs index d80a1d4..8fcc6db 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -198,6 +198,10 @@ impl Model { self.data.no_system_message } + pub fn system_prompt_prefix(&self) -> Option<&str> { + self.data.system_prompt_prefix.as_deref() + } + pub fn max_tokens_per_chunk(&self) -> Option { self.data.max_tokens_per_chunk } @@ -321,6 +325,8 @@ pub struct ModelData { no_stream: bool, #[serde(default, skip_serializing_if = "std::ops::Not::not")] no_system_message: bool, + #[serde(skip_serializing_if = "Option::is_none")] + system_prompt_prefix: Option, // embedding-only properties #[serde(skip_serializing_if = "Option::is_none")] -- cgit v1.2.3