[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:
co-authored by
Chang Su
Simo Lin
parent
ad8d24c39e
commit
d28caaf60a
@@ -22,20 +22,20 @@ use crate::protocols::{
|
||||
pub enum ChatMessage {
|
||||
#[serde(rename = "system")]
|
||||
System {
|
||||
content: String,
|
||||
content: MessageContent,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
name: Option<String>,
|
||||
},
|
||||
#[serde(rename = "user")]
|
||||
User {
|
||||
content: UserMessageContent,
|
||||
content: MessageContent,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
name: Option<String>,
|
||||
},
|
||||
#[serde(rename = "assistant")]
|
||||
Assistant {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
content: Option<String>,
|
||||
content: Option<MessageContent>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
name: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
@@ -46,20 +46,38 @@ pub enum ChatMessage {
|
||||
},
|
||||
#[serde(rename = "tool")]
|
||||
Tool {
|
||||
content: String,
|
||||
content: MessageContent,
|
||||
tool_call_id: String,
|
||||
},
|
||||
#[serde(rename = "function")]
|
||||
Function { content: String, name: String },
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Deserialize, Serialize)]
|
||||
#[derive(Debug, Clone, Deserialize, Serialize, PartialEq)]
|
||||
#[serde(untagged)]
|
||||
pub enum UserMessageContent {
|
||||
pub enum MessageContent {
|
||||
Text(String),
|
||||
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
|
||||
// ============================================================================
|
||||
@@ -320,12 +338,12 @@ fn validate_messages(messages: &[ChatMessage]) -> Result<(), validator::Validati
|
||||
for msg in messages.iter() {
|
||||
if let ChatMessage::User { content, .. } = msg {
|
||||
match content {
|
||||
UserMessageContent::Text(text) if text.is_empty() => {
|
||||
MessageContent::Text(text) if text.is_empty() => {
|
||||
return Err(validator::ValidationError::new(
|
||||
"message content cannot be empty",
|
||||
));
|
||||
}
|
||||
UserMessageContent::Parts(parts) if parts.is_empty() => {
|
||||
MessageContent::Parts(parts) if parts.is_empty() => {
|
||||
return Err(validator::ValidationError::new(
|
||||
"message content parts cannot be empty",
|
||||
));
|
||||
@@ -589,27 +607,18 @@ impl GenerationRequest for ChatCompletionRequest {
|
||||
self.messages
|
||||
.iter()
|
||||
.filter_map(|msg| match msg {
|
||||
ChatMessage::System { content, .. } => Some(content.clone()),
|
||||
ChatMessage::User { content, .. } => match content {
|
||||
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::System { content, .. } => Some(content.to_simple_string()),
|
||||
ChatMessage::User { content, .. } => Some(content.to_simple_string()),
|
||||
ChatMessage::Assistant {
|
||||
content,
|
||||
reasoning_content,
|
||||
..
|
||||
} => {
|
||||
// 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();
|
||||
if main_content.is_empty() && reasoning.is_empty() {
|
||||
None
|
||||
@@ -617,7 +626,7 @@ impl GenerationRequest for ChatCompletionRequest {
|
||||
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()),
|
||||
})
|
||||
.collect::<Vec<String>>()
|
||||
|
||||
@@ -77,7 +77,7 @@ impl StringOrArray {
|
||||
// Content Parts (for multimodal messages)
|
||||
// ============================================================================
|
||||
|
||||
#[derive(Debug, Clone, Deserialize, Serialize)]
|
||||
#[derive(Debug, Clone, Deserialize, Serialize, PartialEq)]
|
||||
#[serde(tag = "type")]
|
||||
pub enum ContentPart {
|
||||
#[serde(rename = "text")]
|
||||
@@ -86,7 +86,7 @@ pub enum ContentPart {
|
||||
ImageUrl { image_url: ImageUrl },
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Deserialize, Serialize)]
|
||||
#[derive(Debug, Clone, Deserialize, Serialize, PartialEq)]
|
||||
pub struct ImageUrl {
|
||||
pub url: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
|
||||
@@ -16,7 +16,7 @@ use tracing::debug;
|
||||
|
||||
use super::types::HarmonyBuildOutput;
|
||||
use crate::protocols::{
|
||||
chat::{ChatCompletionRequest, ChatMessage, UserMessageContent},
|
||||
chat::{ChatCompletionRequest, ChatMessage, MessageContent},
|
||||
common::{ContentPart, Tool},
|
||||
responses::{
|
||||
ReasoningEffort as ResponsesReasoningEffort, ResponseContentPart, ResponseInput,
|
||||
@@ -704,7 +704,7 @@ impl HarmonyBuilder {
|
||||
},
|
||||
recipient: None,
|
||||
content: vec![Content::Text(TextContent {
|
||||
text: content.clone(),
|
||||
text: content.to_simple_string(),
|
||||
})],
|
||||
channel: None,
|
||||
content_type: None,
|
||||
@@ -715,8 +715,8 @@ impl HarmonyBuilder {
|
||||
ChatMessage::User { content, name } => {
|
||||
// Extract text from user content
|
||||
let text = match content {
|
||||
UserMessageContent::Text(text) => text.clone(),
|
||||
UserMessageContent::Parts(parts) => {
|
||||
MessageContent::Text(text) => text.clone(),
|
||||
MessageContent::Parts(parts) => {
|
||||
// For multimodal content, extract text parts
|
||||
parts
|
||||
.iter()
|
||||
@@ -772,7 +772,11 @@ impl HarmonyBuilder {
|
||||
} else {
|
||||
// Regular assistant message with content
|
||||
// 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 !text.is_empty() {
|
||||
text.push('\n');
|
||||
@@ -813,7 +817,7 @@ impl HarmonyBuilder {
|
||||
},
|
||||
recipient: Some("assistant".to_string()),
|
||||
content: vec![Content::Text(TextContent {
|
||||
text: content.clone(),
|
||||
text: content.to_simple_string(),
|
||||
})],
|
||||
channel: None,
|
||||
content_type: None,
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
|
||||
use crate::{
|
||||
protocols::{
|
||||
chat::{ChatCompletionRequest, ChatCompletionResponse, ChatMessage, UserMessageContent},
|
||||
chat::{ChatCompletionRequest, ChatCompletionResponse, ChatMessage, MessageContent},
|
||||
common::{
|
||||
FunctionCallResponse, JsonSchemaFormat, ResponseFormat, StreamOptions, ToolCall,
|
||||
UsageInfo,
|
||||
@@ -38,7 +38,7 @@ pub fn responses_to_chat(req: &ResponsesRequest) -> Result<ChatCompletionRequest
|
||||
// 1. Add system message if instructions provided
|
||||
if let Some(instructions) = &req.instructions {
|
||||
messages.push(ChatMessage::System {
|
||||
content: instructions.clone(),
|
||||
content: MessageContent::Text(instructions.clone()),
|
||||
name: None,
|
||||
});
|
||||
}
|
||||
@@ -48,7 +48,7 @@ pub fn responses_to_chat(req: &ResponsesRequest) -> Result<ChatCompletionRequest
|
||||
ResponseInput::Text(text) => {
|
||||
// Simple text input → user message
|
||||
messages.push(ChatMessage::User {
|
||||
content: UserMessageContent::Text(text.clone()),
|
||||
content: MessageContent::Text(text.clone()),
|
||||
name: None,
|
||||
});
|
||||
}
|
||||
@@ -111,7 +111,7 @@ pub fn responses_to_chat(req: &ResponsesRequest) -> Result<ChatCompletionRequest
|
||||
// Add tool result message if output exists
|
||||
if let Some(output_text) = output {
|
||||
messages.push(ChatMessage::Tool {
|
||||
content: output_text.clone(),
|
||||
content: MessageContent::Text(output_text.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
|
||||
// For Chat path, we just use the call_id
|
||||
messages.push(ChatMessage::Tool {
|
||||
content: output.clone(),
|
||||
content: MessageContent::Text(output.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 {
|
||||
match role {
|
||||
"user" => ChatMessage::User {
|
||||
content: UserMessageContent::Text(text),
|
||||
content: MessageContent::Text(text),
|
||||
name: None,
|
||||
},
|
||||
"assistant" => ChatMessage::Assistant {
|
||||
content: Some(text),
|
||||
content: Some(MessageContent::Text(text)),
|
||||
name: None,
|
||||
tool_calls: None,
|
||||
reasoning_content: None,
|
||||
},
|
||||
"system" => ChatMessage::System {
|
||||
content: text,
|
||||
content: MessageContent::Text(text),
|
||||
name: None,
|
||||
},
|
||||
_ => {
|
||||
// Unknown role, treat as user message
|
||||
ChatMessage::User {
|
||||
content: UserMessageContent::Text(text),
|
||||
content: MessageContent::Text(text),
|
||||
name: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -948,7 +948,7 @@ mod tests {
|
||||
use super::*;
|
||||
use crate::{
|
||||
protocols::{
|
||||
chat::{ChatMessage, UserMessageContent},
|
||||
chat::{ChatMessage, MessageContent},
|
||||
common::{ContentPart, ImageUrl},
|
||||
},
|
||||
tokenizer::chat_template::ChatTemplateContentFormat,
|
||||
@@ -957,7 +957,7 @@ mod tests {
|
||||
#[test]
|
||||
fn test_transform_messages_string_format() {
|
||||
let messages = vec![ChatMessage::User {
|
||||
content: UserMessageContent::Parts(vec![
|
||||
content: MessageContent::Parts(vec![
|
||||
ContentPart::Text {
|
||||
text: "Hello".to_string(),
|
||||
},
|
||||
@@ -990,7 +990,7 @@ mod tests {
|
||||
#[test]
|
||||
fn test_transform_messages_openai_format() {
|
||||
let messages = vec![ChatMessage::User {
|
||||
content: UserMessageContent::Parts(vec![
|
||||
content: MessageContent::Parts(vec![
|
||||
ContentPart::Text {
|
||||
text: "Describe this image:".to_string(),
|
||||
},
|
||||
@@ -1024,7 +1024,7 @@ mod tests {
|
||||
#[test]
|
||||
fn test_transform_messages_simple_string_content() {
|
||||
let messages = vec![ChatMessage::User {
|
||||
content: UserMessageContent::Text("Simple text message".to_string()),
|
||||
content: MessageContent::Text("Simple text message".to_string()),
|
||||
name: None,
|
||||
}];
|
||||
|
||||
@@ -1044,11 +1044,11 @@ mod tests {
|
||||
fn test_transform_messages_multiple_messages() {
|
||||
let messages = vec![
|
||||
ChatMessage::System {
|
||||
content: "System prompt".to_string(),
|
||||
content: MessageContent::Text("System prompt".to_string()),
|
||||
name: None,
|
||||
},
|
||||
ChatMessage::User {
|
||||
content: UserMessageContent::Parts(vec![
|
||||
content: MessageContent::Parts(vec![
|
||||
ContentPart::Text {
|
||||
text: "User message".to_string(),
|
||||
},
|
||||
@@ -1079,7 +1079,7 @@ mod tests {
|
||||
#[test]
|
||||
fn test_transform_messages_empty_text_parts() {
|
||||
let messages = vec![ChatMessage::User {
|
||||
content: UserMessageContent::Parts(vec![ContentPart::ImageUrl {
|
||||
content: MessageContent::Parts(vec![ContentPart::ImageUrl {
|
||||
image_url: ImageUrl {
|
||||
url: "https://example.com/image.jpg".to_string(),
|
||||
detail: None,
|
||||
@@ -1101,11 +1101,11 @@ mod tests {
|
||||
fn test_transform_messages_mixed_content_types() {
|
||||
let messages = vec![
|
||||
ChatMessage::User {
|
||||
content: UserMessageContent::Text("Plain text".to_string()),
|
||||
content: MessageContent::Text("Plain text".to_string()),
|
||||
name: None,
|
||||
},
|
||||
ChatMessage::User {
|
||||
content: UserMessageContent::Parts(vec![
|
||||
content: MessageContent::Parts(vec![
|
||||
ContentPart::Text {
|
||||
text: "With image".to_string(),
|
||||
},
|
||||
|
||||
@@ -23,7 +23,7 @@ use crate::{
|
||||
metrics::RouterMetrics,
|
||||
policies::{LoadBalancingPolicy, PolicyRegistry},
|
||||
protocols::{
|
||||
chat::{ChatCompletionRequest, ChatMessage, UserMessageContent},
|
||||
chat::{ChatCompletionRequest, ChatMessage, MessageContent},
|
||||
classify::ClassifyRequest,
|
||||
common::{InputIds, StringOrArray},
|
||||
completion::CompletionRequest,
|
||||
@@ -1099,10 +1099,10 @@ impl RouterTrait for PDRouter {
|
||||
let request_text = if self.policies_need_request_text() {
|
||||
body.messages.first().and_then(|msg| match msg {
|
||||
ChatMessage::User { content, .. } => match content {
|
||||
UserMessageContent::Text(text) => Some(text.clone()),
|
||||
UserMessageContent::Parts(_) => None,
|
||||
MessageContent::Text(text) => Some(text.clone()),
|
||||
MessageContent::Parts(_) => None,
|
||||
},
|
||||
ChatMessage::System { content, .. } => Some(content.clone()),
|
||||
ChatMessage::System { content, .. } => Some(content.to_simple_string()),
|
||||
_ => None,
|
||||
})
|
||||
} else {
|
||||
|
||||
Reference in New Issue
Block a user