summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2025-02-10 14:57:08 +0800
committerGitHub <noreply@github.com>2025-02-10 14:57:08 +0800
commit69974365ee786af5b643c335e2a12dd7817b84be (patch)
treecdb43b4911907fc918b0719379ffdad4bb880f8d
parentbaf0c1ed07fc185fd237fbdf51aecaa8ed926ba6 (diff)
downloadaichat-69974365ee786af5b643c335e2a12dd7817b84be.tar.gz
feat: new model field `system_prompt_prefix` (#1163)
-rw-r--r--models.yaml6
-rw-r--r--src/client/message.rs59
-rw-r--r--src/client/model.rs6
-rw-r--r--src/config/input.rs8
-rw-r--r--src/serve.rs4
5 files changed, 55 insertions, 28 deletions
diff --git a/models.yaml b/models.yaml
index 2b6b62e..bce8990 100644
--- a/models.yaml
+++ b/models.yaml
@@ -53,6 +53,7 @@
supports_vision: true
supports_function_calling: true
supports_reasoning: true
+ system_prompt_prefix: Formatting re-enabled
- name: o1
max_input_tokens: 200000
input_price: 15
@@ -60,6 +61,7 @@
supports_vision: true
supports_function_calling: true
supports_reasoning: true
+ system_prompt_prefix: Formatting re-enabled
- name: o1-preview
max_input_tokens: 128000
max_output_tokens: 32768
@@ -1172,6 +1174,7 @@
supports_vision: true
supports_function_calling: true
supports_reasoning: true
+ system_prompt_prefix: Formatting re-enabled
- name: openai/o1
max_input_tokens: 128000
input_price: 15
@@ -1179,6 +1182,7 @@
supports_vision: true
supports_function_calling: true
supports_reasoning: true
+ system_prompt_prefix: Formatting re-enabled
- name: openai/o1-preview
max_input_tokens: 128000
input_price: 15
@@ -1477,11 +1481,13 @@
supports_function_calling: true
supports_vision: true
supports_reasoning: true
+ system_prompt_prefix: Formatting re-enabled
- name: o1
max_input_tokens: 200000
supports_function_calling: true
supports_vision: true
supports_reasoning: true
+ system_prompt_prefix: Formatting re-enabled
- name: o1-preview
max_input_tokens: 128000
supports_reasoning: true
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")]
diff --git a/src/config/input.rs b/src/config/input.rs
index c16e1e2..1067b2c 100644
--- a/src/config/input.rs
+++ b/src/config/input.rs
@@ -1,8 +1,8 @@
use super::*;
use crate::client::{
- init_client, patch_system_message, ChatCompletionsData, Client, ImageUrl, Message,
- MessageContent, MessageContentPart, MessageContentToolCalls, MessageRole, Model,
+ init_client, patch_messages, ChatCompletionsData, Client, ImageUrl, Message, MessageContent,
+ MessageContentPart, MessageContentToolCalls, MessageRole, Model,
};
use crate::function::ToolResult;
use crate::utils::{base64_encode, is_loader_protocol, sha256, AbortSignal};
@@ -238,9 +238,7 @@ impl Input {
stream: bool,
) -> Result<ChatCompletionsData> {
let mut messages = self.build_messages()?;
- if model.no_system_message() {
- patch_system_message(&mut messages);
- }
+ patch_messages(&mut messages, model);
model.guard_max_input_tokens(&messages)?;
let temperature = self.role().temperature();
let top_p = self.role().top_p();
diff --git a/src/serve.rs b/src/serve.rs
index b3fe5e0..f51e56f 100644
--- a/src/serve.rs
+++ b/src/serve.rs
@@ -309,9 +309,7 @@ impl Server {
let completion_id = generate_completion_id();
let created = Utc::now().timestamp();
- if client.model().no_system_message() {
- patch_system_message(&mut messages);
- }
+ patch_messages(&mut messages, client.model());
let data: ChatCompletionsData = ChatCompletionsData {
messages,
temperature,