[router] Support complex assistant and tool messages in /chat/completions (#12860)

Co-authored-by: Chang Su <chang.s.su@oracle.com>
Co-authored-by: Simo Lin <linsimo.mark@gmail.com>
This commit is contained in:
Danylo Vashchilenko
2025-11-12 00:14:15 -08:00
committed by GitHub
co-authored by Chang Su Simo Lin
parent ad8d24c39e
commit d28caaf60a
13 changed files with 127 additions and 104 deletions
+8 -6
View File
@@ -5,7 +5,7 @@ use serde_json::{from_str, to_string, to_value, to_vec};
use sglang_router_rs::{ use sglang_router_rs::{
core::{BasicWorker, BasicWorkerBuilder, Worker, WorkerType}, core::{BasicWorker, BasicWorkerBuilder, Worker, WorkerType},
protocols::{ protocols::{
chat::{ChatCompletionRequest, ChatMessage, UserMessageContent}, chat::{ChatCompletionRequest, ChatMessage, MessageContent},
common::StringOrArray, common::StringOrArray,
completion::CompletionRequest, completion::CompletionRequest,
generate::GenerateRequest, generate::GenerateRequest,
@@ -148,11 +148,11 @@ fn create_sample_chat_completion_request() -> ChatCompletionRequest {
model: "gpt-3.5-turbo".to_string(), model: "gpt-3.5-turbo".to_string(),
messages: vec![ messages: vec![
ChatMessage::System { ChatMessage::System {
content: "You are a helpful assistant".to_string(), content: MessageContent::Text("You are a helpful assistant".to_string()),
name: None, name: None,
}, },
ChatMessage::User { ChatMessage::User {
content: UserMessageContent::Text( content: MessageContent::Text(
"Explain quantum computing in simple terms".to_string(), "Explain quantum computing in simple terms".to_string(),
), ),
name: None, name: None,
@@ -188,18 +188,20 @@ fn create_sample_completion_request() -> CompletionRequest {
#[allow(deprecated)] #[allow(deprecated)]
fn create_large_chat_completion_request() -> ChatCompletionRequest { fn create_large_chat_completion_request() -> ChatCompletionRequest {
let mut messages = vec![ChatMessage::System { let mut messages = vec![ChatMessage::System {
content: "You are a helpful assistant with extensive knowledge.".to_string(), content: MessageContent::Text(
"You are a helpful assistant with extensive knowledge.".to_string(),
),
name: None, name: None,
}]; }];
// Add many user/assistant pairs to simulate a long conversation // Add many user/assistant pairs to simulate a long conversation
for i in 0..50 { for i in 0..50 {
messages.push(ChatMessage::User { messages.push(ChatMessage::User {
content: UserMessageContent::Text(format!("Question {}: What do you think about topic number {} which involves complex reasoning about multiple interconnected systems and their relationships?", i, i)), content: MessageContent::Text(format!("Question {}: What do you think about topic number {} which involves complex reasoning about multiple interconnected systems and their relationships?", i, i)),
name: None, name: None,
}); });
messages.push(ChatMessage::Assistant { messages.push(ChatMessage::Assistant {
content: Some(format!("Answer {}: This is a detailed response about topic {} that covers multiple aspects and provides comprehensive analysis of the interconnected systems you mentioned.", i, i)), content: Some(MessageContent::Text(format!("Answer {}: This is a detailed response about topic {} that covers multiple aspects and provides comprehensive analysis of the interconnected systems you mentioned.", i, i))),
name: None, name: None,
tool_calls: None, tool_calls: None,
reasoning_content: None, reasoning_content: None,
+33 -24
View File
@@ -22,20 +22,20 @@ use crate::protocols::{
pub enum ChatMessage { pub enum ChatMessage {
#[serde(rename = "system")] #[serde(rename = "system")]
System { System {
content: String, content: MessageContent,
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>, name: Option<String>,
}, },
#[serde(rename = "user")] #[serde(rename = "user")]
User { User {
content: UserMessageContent, content: MessageContent,
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>, name: Option<String>,
}, },
#[serde(rename = "assistant")] #[serde(rename = "assistant")]
Assistant { Assistant {
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
content: Option<String>, content: Option<MessageContent>,
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>, name: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
@@ -46,20 +46,38 @@ pub enum ChatMessage {
}, },
#[serde(rename = "tool")] #[serde(rename = "tool")]
Tool { Tool {
content: String, content: MessageContent,
tool_call_id: String, tool_call_id: String,
}, },
#[serde(rename = "function")] #[serde(rename = "function")]
Function { content: String, name: String }, Function { content: String, name: String },
} }
#[derive(Debug, Clone, Deserialize, Serialize)] #[derive(Debug, Clone, Deserialize, Serialize, PartialEq)]
#[serde(untagged)] #[serde(untagged)]
pub enum UserMessageContent { pub enum MessageContent {
Text(String), Text(String),
Parts(Vec<ContentPart>), Parts(Vec<ContentPart>),
} }
impl MessageContent {
pub fn to_simple_string(&self) -> String {
match self {
MessageContent::Text(text) => text.clone(),
MessageContent::Parts(parts) => {
let texts: Vec<String> = parts
.iter()
.filter_map(|part| match part {
ContentPart::Text { text } => Some(text.clone()),
_ => None,
})
.collect();
texts.join(" ")
}
}
}
}
// ============================================================================ // ============================================================================
// Chat Completion Request // Chat Completion Request
// ============================================================================ // ============================================================================
@@ -320,12 +338,12 @@ fn validate_messages(messages: &[ChatMessage]) -> Result<(), validator::Validati
for msg in messages.iter() { for msg in messages.iter() {
if let ChatMessage::User { content, .. } = msg { if let ChatMessage::User { content, .. } = msg {
match content { match content {
UserMessageContent::Text(text) if text.is_empty() => { MessageContent::Text(text) if text.is_empty() => {
return Err(validator::ValidationError::new( return Err(validator::ValidationError::new(
"message content cannot be empty", "message content cannot be empty",
)); ));
} }
UserMessageContent::Parts(parts) if parts.is_empty() => { MessageContent::Parts(parts) if parts.is_empty() => {
return Err(validator::ValidationError::new( return Err(validator::ValidationError::new(
"message content parts cannot be empty", "message content parts cannot be empty",
)); ));
@@ -589,27 +607,18 @@ impl GenerationRequest for ChatCompletionRequest {
self.messages self.messages
.iter() .iter()
.filter_map(|msg| match msg { .filter_map(|msg| match msg {
ChatMessage::System { content, .. } => Some(content.clone()), ChatMessage::System { content, .. } => Some(content.to_simple_string()),
ChatMessage::User { content, .. } => match content { ChatMessage::User { content, .. } => Some(content.to_simple_string()),
UserMessageContent::Text(text) => Some(text.clone()),
UserMessageContent::Parts(parts) => {
let texts: Vec<String> = parts
.iter()
.filter_map(|part| match part {
ContentPart::Text { text } => Some(text.clone()),
_ => None,
})
.collect();
Some(texts.join(" "))
}
},
ChatMessage::Assistant { ChatMessage::Assistant {
content, content,
reasoning_content, reasoning_content,
.. ..
} => { } => {
// Combine content and reasoning content for routing decisions // Combine content and reasoning content for routing decisions
let main_content = content.clone().unwrap_or_default(); let main_content = content
.as_ref()
.map(|c| c.to_simple_string())
.unwrap_or_default();
let reasoning = reasoning_content.clone().unwrap_or_default(); let reasoning = reasoning_content.clone().unwrap_or_default();
if main_content.is_empty() && reasoning.is_empty() { if main_content.is_empty() && reasoning.is_empty() {
None None
@@ -617,7 +626,7 @@ impl GenerationRequest for ChatCompletionRequest {
Some(format!("{} {}", main_content, reasoning).trim().to_string()) Some(format!("{} {}", main_content, reasoning).trim().to_string())
} }
} }
ChatMessage::Tool { content, .. } => Some(content.clone()), ChatMessage::Tool { content, .. } => Some(content.to_simple_string()),
ChatMessage::Function { content, .. } => Some(content.clone()), ChatMessage::Function { content, .. } => Some(content.clone()),
}) })
.collect::<Vec<String>>() .collect::<Vec<String>>()
+2 -2
View File
@@ -77,7 +77,7 @@ impl StringOrArray {
// Content Parts (for multimodal messages) // Content Parts (for multimodal messages)
// ============================================================================ // ============================================================================
#[derive(Debug, Clone, Deserialize, Serialize)] #[derive(Debug, Clone, Deserialize, Serialize, PartialEq)]
#[serde(tag = "type")] #[serde(tag = "type")]
pub enum ContentPart { pub enum ContentPart {
#[serde(rename = "text")] #[serde(rename = "text")]
@@ -86,7 +86,7 @@ pub enum ContentPart {
ImageUrl { image_url: ImageUrl }, ImageUrl { image_url: ImageUrl },
} }
#[derive(Debug, Clone, Deserialize, Serialize)] #[derive(Debug, Clone, Deserialize, Serialize, PartialEq)]
pub struct ImageUrl { pub struct ImageUrl {
pub url: String, pub url: String,
#[serde(skip_serializing_if = "Option::is_none")] #[serde(skip_serializing_if = "Option::is_none")]
+10 -6
View File
@@ -16,7 +16,7 @@ use tracing::debug;
use super::types::HarmonyBuildOutput; use super::types::HarmonyBuildOutput;
use crate::protocols::{ use crate::protocols::{
chat::{ChatCompletionRequest, ChatMessage, UserMessageContent}, chat::{ChatCompletionRequest, ChatMessage, MessageContent},
common::{ContentPart, Tool}, common::{ContentPart, Tool},
responses::{ responses::{
ReasoningEffort as ResponsesReasoningEffort, ResponseContentPart, ResponseInput, ReasoningEffort as ResponsesReasoningEffort, ResponseContentPart, ResponseInput,
@@ -704,7 +704,7 @@ impl HarmonyBuilder {
}, },
recipient: None, recipient: None,
content: vec![Content::Text(TextContent { content: vec![Content::Text(TextContent {
text: content.clone(), text: content.to_simple_string(),
})], })],
channel: None, channel: None,
content_type: None, content_type: None,
@@ -715,8 +715,8 @@ impl HarmonyBuilder {
ChatMessage::User { content, name } => { ChatMessage::User { content, name } => {
// Extract text from user content // Extract text from user content
let text = match content { let text = match content {
UserMessageContent::Text(text) => text.clone(), MessageContent::Text(text) => text.clone(),
UserMessageContent::Parts(parts) => { MessageContent::Parts(parts) => {
// For multimodal content, extract text parts // For multimodal content, extract text parts
parts parts
.iter() .iter()
@@ -772,7 +772,11 @@ impl HarmonyBuilder {
} else { } else {
// Regular assistant message with content // Regular assistant message with content
// Combine content with reasoning if present // Combine content with reasoning if present
let mut text = content.clone().unwrap_or_default(); let mut text = content
.as_ref()
.map(|c| c.to_simple_string())
.unwrap_or_default();
if let Some(reasoning) = reasoning_content { if let Some(reasoning) = reasoning_content {
if !text.is_empty() { if !text.is_empty() {
text.push('\n'); text.push('\n');
@@ -813,7 +817,7 @@ impl HarmonyBuilder {
}, },
recipient: Some("assistant".to_string()), recipient: Some("assistant".to_string()),
content: vec![Content::Text(TextContent { content: vec![Content::Text(TextContent {
text: content.clone(), text: content.to_simple_string(),
})], })],
channel: None, channel: None,
content_type: None, content_type: None,
@@ -9,7 +9,7 @@
use crate::{ use crate::{
protocols::{ protocols::{
chat::{ChatCompletionRequest, ChatCompletionResponse, ChatMessage, UserMessageContent}, chat::{ChatCompletionRequest, ChatCompletionResponse, ChatMessage, MessageContent},
common::{ common::{
FunctionCallResponse, JsonSchemaFormat, ResponseFormat, StreamOptions, ToolCall, FunctionCallResponse, JsonSchemaFormat, ResponseFormat, StreamOptions, ToolCall,
UsageInfo, UsageInfo,
@@ -38,7 +38,7 @@ pub fn responses_to_chat(req: &ResponsesRequest) -> Result<ChatCompletionRequest
// 1. Add system message if instructions provided // 1. Add system message if instructions provided
if let Some(instructions) = &req.instructions { if let Some(instructions) = &req.instructions {
messages.push(ChatMessage::System { messages.push(ChatMessage::System {
content: instructions.clone(), content: MessageContent::Text(instructions.clone()),
name: None, name: None,
}); });
} }
@@ -48,7 +48,7 @@ pub fn responses_to_chat(req: &ResponsesRequest) -> Result<ChatCompletionRequest
ResponseInput::Text(text) => { ResponseInput::Text(text) => {
// Simple text input → user message // Simple text input → user message
messages.push(ChatMessage::User { messages.push(ChatMessage::User {
content: UserMessageContent::Text(text.clone()), content: MessageContent::Text(text.clone()),
name: None, name: None,
}); });
} }
@@ -111,7 +111,7 @@ pub fn responses_to_chat(req: &ResponsesRequest) -> Result<ChatCompletionRequest
// Add tool result message if output exists // Add tool result message if output exists
if let Some(output_text) = output { if let Some(output_text) = output {
messages.push(ChatMessage::Tool { messages.push(ChatMessage::Tool {
content: output_text.clone(), content: MessageContent::Text(output_text.clone()),
tool_call_id: id.clone(), tool_call_id: id.clone(),
}); });
} }
@@ -140,7 +140,7 @@ pub fn responses_to_chat(req: &ResponsesRequest) -> Result<ChatCompletionRequest
// Note: The function name is looked up from prev_outputs in Harmony path // Note: The function name is looked up from prev_outputs in Harmony path
// For Chat path, we just use the call_id // For Chat path, we just use the call_id
messages.push(ChatMessage::Tool { messages.push(ChatMessage::Tool {
content: output.clone(), content: MessageContent::Text(output.clone()),
tool_call_id: call_id.clone(), tool_call_id: call_id.clone(),
}); });
} }
@@ -213,23 +213,23 @@ fn extract_text_from_content(content: &[ResponseContentPart]) -> String {
fn role_to_chat_message(role: &str, text: String) -> ChatMessage { fn role_to_chat_message(role: &str, text: String) -> ChatMessage {
match role { match role {
"user" => ChatMessage::User { "user" => ChatMessage::User {
content: UserMessageContent::Text(text), content: MessageContent::Text(text),
name: None, name: None,
}, },
"assistant" => ChatMessage::Assistant { "assistant" => ChatMessage::Assistant {
content: Some(text), content: Some(MessageContent::Text(text)),
name: None, name: None,
tool_calls: None, tool_calls: None,
reasoning_content: None, reasoning_content: None,
}, },
"system" => ChatMessage::System { "system" => ChatMessage::System {
content: text, content: MessageContent::Text(text),
name: None, name: None,
}, },
_ => { _ => {
// Unknown role, treat as user message // Unknown role, treat as user message
ChatMessage::User { ChatMessage::User {
content: UserMessageContent::Text(text), content: MessageContent::Text(text),
name: None, name: None,
} }
} }
+9 -9
View File
@@ -948,7 +948,7 @@ mod tests {
use super::*; use super::*;
use crate::{ use crate::{
protocols::{ protocols::{
chat::{ChatMessage, UserMessageContent}, chat::{ChatMessage, MessageContent},
common::{ContentPart, ImageUrl}, common::{ContentPart, ImageUrl},
}, },
tokenizer::chat_template::ChatTemplateContentFormat, tokenizer::chat_template::ChatTemplateContentFormat,
@@ -957,7 +957,7 @@ mod tests {
#[test] #[test]
fn test_transform_messages_string_format() { fn test_transform_messages_string_format() {
let messages = vec![ChatMessage::User { let messages = vec![ChatMessage::User {
content: UserMessageContent::Parts(vec![ content: MessageContent::Parts(vec![
ContentPart::Text { ContentPart::Text {
text: "Hello".to_string(), text: "Hello".to_string(),
}, },
@@ -990,7 +990,7 @@ mod tests {
#[test] #[test]
fn test_transform_messages_openai_format() { fn test_transform_messages_openai_format() {
let messages = vec![ChatMessage::User { let messages = vec![ChatMessage::User {
content: UserMessageContent::Parts(vec![ content: MessageContent::Parts(vec![
ContentPart::Text { ContentPart::Text {
text: "Describe this image:".to_string(), text: "Describe this image:".to_string(),
}, },
@@ -1024,7 +1024,7 @@ mod tests {
#[test] #[test]
fn test_transform_messages_simple_string_content() { fn test_transform_messages_simple_string_content() {
let messages = vec![ChatMessage::User { let messages = vec![ChatMessage::User {
content: UserMessageContent::Text("Simple text message".to_string()), content: MessageContent::Text("Simple text message".to_string()),
name: None, name: None,
}]; }];
@@ -1044,11 +1044,11 @@ mod tests {
fn test_transform_messages_multiple_messages() { fn test_transform_messages_multiple_messages() {
let messages = vec![ let messages = vec![
ChatMessage::System { ChatMessage::System {
content: "System prompt".to_string(), content: MessageContent::Text("System prompt".to_string()),
name: None, name: None,
}, },
ChatMessage::User { ChatMessage::User {
content: UserMessageContent::Parts(vec![ content: MessageContent::Parts(vec![
ContentPart::Text { ContentPart::Text {
text: "User message".to_string(), text: "User message".to_string(),
}, },
@@ -1079,7 +1079,7 @@ mod tests {
#[test] #[test]
fn test_transform_messages_empty_text_parts() { fn test_transform_messages_empty_text_parts() {
let messages = vec![ChatMessage::User { let messages = vec![ChatMessage::User {
content: UserMessageContent::Parts(vec![ContentPart::ImageUrl { content: MessageContent::Parts(vec![ContentPart::ImageUrl {
image_url: ImageUrl { image_url: ImageUrl {
url: "https://example.com/image.jpg".to_string(), url: "https://example.com/image.jpg".to_string(),
detail: None, detail: None,
@@ -1101,11 +1101,11 @@ mod tests {
fn test_transform_messages_mixed_content_types() { fn test_transform_messages_mixed_content_types() {
let messages = vec![ let messages = vec![
ChatMessage::User { ChatMessage::User {
content: UserMessageContent::Text("Plain text".to_string()), content: MessageContent::Text("Plain text".to_string()),
name: None, name: None,
}, },
ChatMessage::User { ChatMessage::User {
content: UserMessageContent::Parts(vec![ content: MessageContent::Parts(vec![
ContentPart::Text { ContentPart::Text {
text: "With image".to_string(), text: "With image".to_string(),
}, },
+4 -4
View File
@@ -23,7 +23,7 @@ use crate::{
metrics::RouterMetrics, metrics::RouterMetrics,
policies::{LoadBalancingPolicy, PolicyRegistry}, policies::{LoadBalancingPolicy, PolicyRegistry},
protocols::{ protocols::{
chat::{ChatCompletionRequest, ChatMessage, UserMessageContent}, chat::{ChatCompletionRequest, ChatMessage, MessageContent},
classify::ClassifyRequest, classify::ClassifyRequest,
common::{InputIds, StringOrArray}, common::{InputIds, StringOrArray},
completion::CompletionRequest, completion::CompletionRequest,
@@ -1099,10 +1099,10 @@ impl RouterTrait for PDRouter {
let request_text = if self.policies_need_request_text() { let request_text = if self.policies_need_request_text() {
body.messages.first().and_then(|msg| match msg { body.messages.first().and_then(|msg| match msg {
ChatMessage::User { content, .. } => match content { ChatMessage::User { content, .. } => match content {
UserMessageContent::Text(text) => Some(text.clone()), MessageContent::Text(text) => Some(text.clone()),
UserMessageContent::Parts(_) => None, MessageContent::Parts(_) => None,
}, },
ChatMessage::System { content, .. } => Some(content.clone()), ChatMessage::System { content, .. } => Some(content.to_simple_string()),
_ => None, _ => None,
}) })
} else { } else {
@@ -1,5 +1,5 @@
use sglang_router_rs::{ use sglang_router_rs::{
protocols::chat::{ChatMessage, UserMessageContent}, protocols::chat::{ChatMessage, MessageContent},
tokenizer::chat_template::{ tokenizer::chat_template::{
detect_chat_template_content_format, ChatTemplateContentFormat, ChatTemplateParams, detect_chat_template_content_format, ChatTemplateContentFormat, ChatTemplateParams,
ChatTemplateProcessor, ChatTemplateProcessor,
@@ -176,11 +176,11 @@ assistant:
let messages = [ let messages = [
ChatMessage::System { ChatMessage::System {
content: "You are helpful".to_string(), content: MessageContent::Text("You are helpful".to_string()),
name: None, name: None,
}, },
ChatMessage::User { ChatMessage::User {
content: UserMessageContent::Text("Hello".to_string()), content: MessageContent::Text("Hello".to_string()),
name: None, name: None,
}, },
]; ];
@@ -216,7 +216,7 @@ fn test_chat_template_with_tokens_unit_test() {
let processor = ChatTemplateProcessor::new(template.to_string()); let processor = ChatTemplateProcessor::new(template.to_string());
let messages = [ChatMessage::User { let messages = [ChatMessage::User {
content: UserMessageContent::Text("Test".to_string()), content: MessageContent::Text("Test".to_string()),
name: None, name: None,
}]; }];
+10 -10
View File
@@ -1,6 +1,6 @@
use sglang_router_rs::{ use sglang_router_rs::{
protocols::{ protocols::{
chat::{ChatMessage, UserMessageContent}, chat::{ChatMessage, MessageContent},
common::{ContentPart, ImageUrl}, common::{ContentPart, ImageUrl},
}, },
tokenizer::chat_template::{ tokenizer::chat_template::{
@@ -23,7 +23,7 @@ fn test_simple_chat_template() {
let processor = ChatTemplateProcessor::new(template.to_string()); let processor = ChatTemplateProcessor::new(template.to_string());
let messages = [ChatMessage::User { let messages = [ChatMessage::User {
content: UserMessageContent::Text("Test".to_string()), content: MessageContent::Text("Test".to_string()),
name: None, name: None,
}]; }];
@@ -57,7 +57,7 @@ fn test_chat_template_with_tokens() {
let processor = ChatTemplateProcessor::new(template.to_string()); let processor = ChatTemplateProcessor::new(template.to_string());
let messages = [ChatMessage::User { let messages = [ChatMessage::User {
content: UserMessageContent::Text("Test".to_string()), content: MessageContent::Text("Test".to_string()),
name: None, name: None,
}]; }];
@@ -118,11 +118,11 @@ fn test_llama_style_template() {
let messages = [ let messages = [
ChatMessage::System { ChatMessage::System {
content: "You are a helpful assistant".to_string(), content: MessageContent::Text("You are a helpful assistant".to_string()),
name: None, name: None,
}, },
ChatMessage::User { ChatMessage::User {
content: UserMessageContent::Text("What is 2+2?".to_string()), content: MessageContent::Text("What is 2+2?".to_string()),
name: None, name: None,
}, },
]; ];
@@ -173,17 +173,17 @@ fn test_chatml_template() {
let messages = [ let messages = [
ChatMessage::User { ChatMessage::User {
content: UserMessageContent::Text("Hello".to_string()), content: MessageContent::Text("Hello".to_string()),
name: None, name: None,
}, },
ChatMessage::Assistant { ChatMessage::Assistant {
content: Some("Hi there!".to_string()), content: Some(MessageContent::Text("Hi there!".to_string())),
name: None, name: None,
tool_calls: None, tool_calls: None,
reasoning_content: None, reasoning_content: None,
}, },
ChatMessage::User { ChatMessage::User {
content: UserMessageContent::Text("How are you?".to_string()), content: MessageContent::Text("How are you?".to_string()),
name: None, name: None,
}, },
]; ];
@@ -225,7 +225,7 @@ assistant:
let processor = ChatTemplateProcessor::new(template.to_string()); let processor = ChatTemplateProcessor::new(template.to_string());
let messages = [ChatMessage::User { let messages = [ChatMessage::User {
content: UserMessageContent::Text("Test".to_string()), content: MessageContent::Text("Test".to_string()),
name: None, name: None,
}]; }];
@@ -312,7 +312,7 @@ fn test_template_with_multimodal_content() {
let processor = ChatTemplateProcessor::new(template.to_string()); let processor = ChatTemplateProcessor::new(template.to_string());
let messages = [ChatMessage::User { let messages = [ChatMessage::User {
content: UserMessageContent::Parts(vec![ content: MessageContent::Parts(vec![
ContentPart::Text { ContentPart::Text {
text: "Look at this:".to_string(), text: "Look at this:".to_string(),
}, },
+6 -6
View File
@@ -3,7 +3,7 @@ mod tests {
use std::fs; use std::fs;
use sglang_router_rs::{ use sglang_router_rs::{
protocols::chat::{ChatMessage, UserMessageContent}, protocols::chat::{ChatMessage, MessageContent},
tokenizer::{chat_template::ChatTemplateParams, huggingface::HuggingFaceTokenizer}, tokenizer::{chat_template::ChatTemplateParams, huggingface::HuggingFaceTokenizer},
}; };
use tempfile::TempDir; use tempfile::TempDir;
@@ -61,11 +61,11 @@ mod tests {
let messages = [ let messages = [
ChatMessage::User { ChatMessage::User {
content: UserMessageContent::Text("Hello".to_string()), content: MessageContent::Text("Hello".to_string()),
name: None, name: None,
}, },
ChatMessage::Assistant { ChatMessage::Assistant {
content: Some("Hi there".to_string()), content: Some(MessageContent::Text("Hi there".to_string())),
name: None, name: None,
tool_calls: None, tool_calls: None,
reasoning_content: None, reasoning_content: None,
@@ -143,7 +143,7 @@ mod tests {
.unwrap(); .unwrap();
let messages = [ChatMessage::User { let messages = [ChatMessage::User {
content: UserMessageContent::Text("Test".to_string()), content: MessageContent::Text("Test".to_string()),
name: None, name: None,
}]; }];
@@ -202,11 +202,11 @@ mod tests {
let messages = [ let messages = [
ChatMessage::User { ChatMessage::User {
content: UserMessageContent::Text("Hello".to_string()), content: MessageContent::Text("Hello".to_string()),
name: None, name: None,
}, },
ChatMessage::Assistant { ChatMessage::Assistant {
content: Some("World".to_string()), content: Some(MessageContent::Text("World".to_string())),
name: None, name: None,
tool_calls: None, tool_calls: None,
reasoning_content: None, reasoning_content: None,
+17 -17
View File
@@ -1,6 +1,6 @@
use serde_json::json; use serde_json::json;
use sglang_router_rs::protocols::{ use sglang_router_rs::protocols::{
chat::{ChatCompletionRequest, ChatMessage, UserMessageContent}, chat::{ChatCompletionRequest, ChatMessage, MessageContent},
common::{ common::{
Function, FunctionCall, FunctionChoice, StreamOptions, Tool, ToolChoice, ToolChoiceValue, Function, FunctionCall, FunctionChoice, StreamOptions, Tool, ToolChoice, ToolChoiceValue,
ToolReference, ToolReference,
@@ -17,7 +17,7 @@ fn test_max_tokens_normalizes_to_max_completion_tokens() {
let mut req = ChatCompletionRequest { let mut req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
max_tokens: Some(100), max_tokens: Some(100),
@@ -50,7 +50,7 @@ fn test_max_completion_tokens_takes_precedence() {
let mut req = ChatCompletionRequest { let mut req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
max_tokens: Some(100), max_tokens: Some(100),
@@ -76,7 +76,7 @@ fn test_functions_normalizes_to_tools() {
let mut req = ChatCompletionRequest { let mut req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
functions: Some(vec![Function { functions: Some(vec![Function {
@@ -112,7 +112,7 @@ fn test_function_call_normalizes_to_tool_choice() {
let mut req = ChatCompletionRequest { let mut req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
function_call: Some(FunctionCall::None), function_call: Some(FunctionCall::None),
@@ -148,7 +148,7 @@ fn test_function_call_function_variant_normalizes() {
let mut req = ChatCompletionRequest { let mut req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
function_call: Some(FunctionCall::Function { function_call: Some(FunctionCall::Function {
@@ -198,7 +198,7 @@ fn test_stream_options_requires_stream_enabled() {
let req = ChatCompletionRequest { let req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
stream: false, stream: false,
@@ -226,7 +226,7 @@ fn test_stream_options_valid_when_stream_enabled() {
let req = ChatCompletionRequest { let req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
stream: true, stream: true,
@@ -248,7 +248,7 @@ fn test_no_stream_options_valid_when_stream_disabled() {
let req = ChatCompletionRequest { let req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
stream: false, stream: false,
@@ -269,7 +269,7 @@ fn test_tool_choice_function_not_found() {
let req = ChatCompletionRequest { let req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
tools: Some(vec![Tool { tools: Some(vec![Tool {
@@ -305,7 +305,7 @@ fn test_tool_choice_function_exists_valid() {
let req = ChatCompletionRequest { let req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
tools: Some(vec![Tool { tools: Some(vec![Tool {
@@ -335,7 +335,7 @@ fn test_tool_choice_allowed_tools_invalid_mode() {
let req = ChatCompletionRequest { let req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
tools: Some(vec![Tool { tools: Some(vec![Tool {
@@ -372,7 +372,7 @@ fn test_tool_choice_allowed_tools_valid_mode_auto() {
let req = ChatCompletionRequest { let req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
tools: Some(vec![Tool { tools: Some(vec![Tool {
@@ -403,7 +403,7 @@ fn test_tool_choice_allowed_tools_valid_mode_required() {
let req = ChatCompletionRequest { let req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
tools: Some(vec![Tool { tools: Some(vec![Tool {
@@ -434,7 +434,7 @@ fn test_tool_choice_allowed_tools_tool_not_found() {
let req = ChatCompletionRequest { let req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
tools: Some(vec![Tool { tools: Some(vec![Tool {
@@ -471,7 +471,7 @@ fn test_tool_choice_allowed_tools_multiple_tools_valid() {
let req = ChatCompletionRequest { let req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
tools: Some(vec![ tools: Some(vec![
@@ -518,7 +518,7 @@ fn test_tool_choice_allowed_tools_one_invalid_among_valid() {
let req = ChatCompletionRequest { let req = ChatCompletionRequest {
model: "test-model".to_string(), model: "test-model".to_string(),
messages: vec![ChatMessage::User { messages: vec![ChatMessage::User {
content: UserMessageContent::Text("hello".to_string()), content: MessageContent::Text("hello".to_string()),
name: None, name: None,
}], }],
tools: Some(vec![ tools: Some(vec![
+13 -5
View File
@@ -1,5 +1,5 @@
use serde_json::json; use serde_json::json;
use sglang_router_rs::protocols::chat::{ChatMessage, UserMessageContent}; use sglang_router_rs::protocols::chat::{ChatMessage, MessageContent};
#[test] #[test]
fn test_chat_message_tagged_by_role_system() { fn test_chat_message_tagged_by_role_system() {
@@ -11,7 +11,10 @@ fn test_chat_message_tagged_by_role_system() {
let msg: ChatMessage = serde_json::from_value(json).unwrap(); let msg: ChatMessage = serde_json::from_value(json).unwrap();
match msg { match msg {
ChatMessage::System { content, .. } => { ChatMessage::System { content, .. } => {
assert_eq!(content, "You are a helpful assistant"); assert_eq!(
content,
MessageContent::Text("You are a helpful assistant".to_string())
)
} }
_ => panic!("Expected System variant"), _ => panic!("Expected System variant"),
} }
@@ -27,7 +30,7 @@ fn test_chat_message_tagged_by_role_user() {
let msg: ChatMessage = serde_json::from_value(json).unwrap(); let msg: ChatMessage = serde_json::from_value(json).unwrap();
match msg { match msg {
ChatMessage::User { content, .. } => match content { ChatMessage::User { content, .. } => match content {
UserMessageContent::Text(text) => assert_eq!(text, "Hello"), MessageContent::Text(text) => assert_eq!(text, "Hello"),
_ => panic!("Expected text content"), _ => panic!("Expected text content"),
}, },
_ => panic!("Expected User variant"), _ => panic!("Expected User variant"),
@@ -44,7 +47,7 @@ fn test_chat_message_tagged_by_role_assistant() {
let msg: ChatMessage = serde_json::from_value(json).unwrap(); let msg: ChatMessage = serde_json::from_value(json).unwrap();
match msg { match msg {
ChatMessage::Assistant { content, .. } => { ChatMessage::Assistant { content, .. } => {
assert_eq!(content, Some("Hi there!".to_string())); assert_eq!(content, Some(MessageContent::Text("Hi there!".to_string())));
} }
_ => panic!("Expected Assistant variant"), _ => panic!("Expected Assistant variant"),
} }
@@ -64,7 +67,12 @@ fn test_chat_message_tagged_by_role_tool() {
content, content,
tool_call_id, tool_call_id,
} => { } => {
assert_eq!(content, "Tool result"); match content {
MessageContent::Text(text) => {
assert_eq!(text, "Tool result");
}
_ => panic!("Expected content to be a string"),
}
assert_eq!(tool_call_id, "call_123"); assert_eq!(tool_call_id, "call_123");
} }
_ => panic!("Expected Tool variant"), _ => panic!("Expected Tool variant"),
+2 -2
View File
@@ -23,7 +23,7 @@ use sglang_router_rs::{
}, },
data_connector::{ResponseId, StoredResponse}, data_connector::{ResponseId, StoredResponse},
protocols::{ protocols::{
chat::{ChatCompletionRequest, ChatMessage, UserMessageContent}, chat::{ChatCompletionRequest, ChatMessage, MessageContent},
common::StringOrArray, common::StringOrArray,
completion::CompletionRequest, completion::CompletionRequest,
generate::GenerateRequest, generate::GenerateRequest,
@@ -661,7 +661,7 @@ async fn test_openai_router_chat_completion_with_mock() {
// Create a minimal chat completion request // Create a minimal chat completion request
let mut chat_request = create_minimal_chat_request(); let mut chat_request = create_minimal_chat_request();
chat_request.messages = vec![ChatMessage::User { chat_request.messages = vec![ChatMessage::User {
content: UserMessageContent::Text("Hello, how are you?".to_string()), content: MessageContent::Text("Hello, how are you?".to_string()),
name: None, name: None,
}]; }];
chat_request.temperature = Some(0.7); chat_request.temperature = Some(0.7);