From 4f8d895154c602fb51a7694911efd96d52129f82 Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 24 Apr 2024 07:56:24 +0800 Subject: refactor: handling of system message (#432) --- src/client/claude.rs | 8 ++++++-- src/client/cohere.rs | 8 ++++++-- src/client/common.rs | 8 ++++++++ src/client/message.rs | 15 +++++++++++++++ src/client/ollama.rs | 8 +++----- 5 files changed, 38 insertions(+), 9 deletions(-) (limited to 'src/client') diff --git a/src/client/claude.rs b/src/client/claude.rs index 24bffbf..5c2a912 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -1,5 +1,5 @@ use super::{ - patch_system_message, ClaudeClient, Client, ExtraConfig, ImageUrl, MessageContent, + extract_sytem_message, ClaudeClient, Client, ExtraConfig, ImageUrl, MessageContent, MessageContentPart, Model, ModelConfig, PromptType, ReplyHandler, SendData, }; @@ -141,7 +141,7 @@ fn build_body(data: SendData, model: &Model) -> Result { stream, } = data; - patch_system_message(&mut messages); + let system_message = extract_sytem_message(&mut messages); let mut network_image_urls = vec![]; let messages: Vec = messages @@ -196,6 +196,10 @@ fn build_body(data: SendData, model: &Model) -> Result { "messages": messages, }); + if let Some(system) = system_message { + body["system"] = system.into(); + } + if let Some(v) = temperature { body["temperature"] = v.into(); } diff --git a/src/client/cohere.rs b/src/client/cohere.rs index 445c145..db2ad28 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -1,5 +1,5 @@ use super::{ - json_stream, message::*, patch_system_message, Client, CohereClient, ExtraConfig, Model, + extract_sytem_message, json_stream, message::*, Client, CohereClient, ExtraConfig, Model, ModelConfig, PromptType, ReplyHandler, SendData, }; @@ -129,7 +129,7 @@ pub(crate) fn build_body(data: SendData, model: &Model) -> Result { stream, } = data; - patch_system_message(&mut messages); + let system_message = extract_sytem_message(&mut messages); let mut image_urls = vec![]; let mut messages: Vec = messages @@ -174,6 +174,10 @@ pub(crate) fn build_body(data: SendData, model: &Model) -> Result { "message": message, }); + if let Some(preamble) = system_message { + body["preamble"] = preamble.into(); + } + if let Some(max_tokens) = model.max_output_tokens { body["max_tokens"] = max_tokens.into(); } diff --git a/src/client/common.rs b/src/client/common.rs index b6ccce7..62c5f92 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -414,6 +414,14 @@ pub fn patch_system_message(messages: &mut Vec) { } } +pub fn extract_sytem_message(messages: &mut Vec) -> Option { + if messages[0].role.is_system() { + let system_message = messages.remove(0); + return Some(system_message.content.to_text()); + } + None +} + pub async fn json_stream(mut stream: S, mut handle: F) -> Result<()> where S: Stream> + Unpin, diff --git a/src/client/message.rs b/src/client/message.rs index 1a454ab..f9978c5 100644 --- a/src/client/message.rs +++ b/src/client/message.rs @@ -85,6 +85,21 @@ impl MessageContent { } } } + + pub fn to_text(&self) -> String { + match self { + MessageContent::Text(text) => text.to_string(), + MessageContent::Array(list) => { + let mut parts = vec![]; + for item in list { + if let MessageContentPart::Text { text } = item { + parts.push(text.clone()) + } + } + parts.join("\n\n") + } + } + } } #[derive(Debug, Clone, Deserialize, Serialize)] diff --git a/src/client/ollama.rs b/src/client/ollama.rs index de658c1..403712f 100644 --- a/src/client/ollama.rs +++ b/src/client/ollama.rs @@ -1,6 +1,6 @@ use super::{ - message::*, patch_system_message, Client, ExtraConfig, Model, ModelConfig, OllamaClient, - PromptType, ReplyHandler, SendData, + message::*, Client, ExtraConfig, Model, ModelConfig, OllamaClient, PromptType, ReplyHandler, + SendData, }; use crate::utils::PromptKind; @@ -121,13 +121,11 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHand fn build_body(data: SendData, model: &Model) -> Result { let SendData { - mut messages, + messages, temperature, stream, } = data; - patch_system_message(&mut messages); - let mut network_image_urls = vec![]; let messages: Vec = messages .into_iter() -- cgit v1.2.3