summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
Diffstat (limited to 'src/client')
-rw-r--r--src/client/message.rs59
-rw-r--r--src/client/model.rs6
2 files changed, 45 insertions, 20 deletions
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<Message>) {
- if messages[0].role.is_system() {
+pub fn patch_messages(messages: &mut Vec<Message>, 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<usize> {
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<String>,
// embedding-only properties
#[serde(skip_serializing_if = "Option::is_none")]