diff options
| author | sigoden <sigoden@gmail.com> | 2025-02-10 14:57:08 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2025-02-10 14:57:08 +0800 |
| commit | 69974365ee786af5b643c335e2a12dd7817b84be (patch) | |
| tree | cdb43b4911907fc918b0719379ffdad4bb880f8d | |
| parent | baf0c1ed07fc185fd237fbdf51aecaa8ed926ba6 (diff) | |
| download | aichat-69974365ee786af5b643c335e2a12dd7817b84be.tar.gz | |
feat: new model field `system_prompt_prefix` (#1163)
| -rw-r--r-- | models.yaml | 6 | ||||
| -rw-r--r-- | src/client/message.rs | 59 | ||||
| -rw-r--r-- | src/client/model.rs | 6 | ||||
| -rw-r--r-- | src/config/input.rs | 8 | ||||
| -rw-r--r-- | src/serve.rs | 4 |
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, |
