From 5843ebbf2f6bf92663df1759056d7244a88d9df7 Mon Sep 17 00:00:00 2001 From: Ziwen Zhao Date: Tue, 24 Feb 2026 13:57:16 -0800 Subject: [PATCH 1/3] refactor(protocol): model ResponseTool as tagged enum to match Responses spec and tighten MCP validation Signed-off-by: Ziwen Zhao --- mcp/src/responses_bridge.rs | 26 +-- .../routers/grpc/common/responses/utils.rs | 49 ++--- .../src/routers/grpc/harmony/builder.rs | 40 ++-- .../routers/grpc/harmony/responses/common.rs | 17 +- .../grpc/harmony/responses/non_streaming.rs | 11 +- .../grpc/harmony/responses/streaming.rs | 28 +-- .../grpc/harmony/stages/preparation.rs | 4 +- .../src/routers/grpc/harmony/streaming.rs | 42 ++-- .../grpc/regular/responses/conversions.rs | 8 +- model_gateway/src/routers/mcp_utils.rs | 133 ++++++------ .../src/routers/openai/responses/streaming.rs | 10 +- .../src/routers/openai/responses/utils.rs | 20 +- model_gateway/tests/api/responses_api_test.rs | 75 +++---- model_gateway/tests/spec/responses.rs | 197 ++++++------------ protocols/src/responses.rs | 153 +++++++------- 15 files changed, 338 insertions(+), 475 deletions(-) diff --git a/mcp/src/responses_bridge.rs b/mcp/src/responses_bridge.rs index 6002f96a7e..42c456a2ed 100644 --- a/mcp/src/responses_bridge.rs +++ b/mcp/src/responses_bridge.rs @@ -8,7 +8,7 @@ use openai_protocol::{ common::{Function, Tool}, - responses::{generate_id, McpToolInfo, ResponseOutputItem, ResponseTool, ResponseToolType}, + responses::{generate_id, FunctionTool, McpToolInfo, ResponseOutputItem, ResponseTool}, }; use serde_json::{json, Value}; @@ -88,21 +88,15 @@ pub fn build_response_tools_with_names( ) -> Vec { entries .iter() - .map(|entry| ResponseTool { - r#type: ResponseToolType::Mcp, - function: Some(Function { - name: resolved_name_for_entry(entry, exposed_names).to_string(), - description: entry.tool.description.as_ref().map(|d| d.to_string()), - parameters: Value::Object((*entry.tool.input_schema).clone()), - strict: None, - }), - server_url: None, - authorization: None, - headers: None, - server_label: Some(entry.server_key().to_string()), - server_description: None, - require_approval: None, - allowed_tools: None, + .map(|entry| { + ResponseTool::Function(FunctionTool { + function: Function { + name: resolved_name_for_entry(entry, exposed_names).to_string(), + description: entry.tool.description.as_ref().map(|d| d.to_string()), + parameters: Value::Object((*entry.tool.input_schema).clone()), + strict: None, + }, + }) }) .collect() } diff --git a/model_gateway/src/routers/grpc/common/responses/utils.rs b/model_gateway/src/routers/grpc/common/responses/utils.rs index 6c00af4bdf..df4ffa81dc 100644 --- a/model_gateway/src/routers/grpc/common/responses/utils.rs +++ b/model_gateway/src/routers/grpc/common/responses/utils.rs @@ -5,7 +5,7 @@ use std::sync::Arc; use axum::response::Response; use openai_protocol::{ common::Tool, - responses::{ResponseTool, ResponseToolType, ResponsesRequest, ResponsesResponse}, + responses::{ResponseTool, ResponsesRequest, ResponsesResponse}, }; use serde_json::to_value; use smg_data_connector::{ConversationItemStorage, ConversationStorage, ResponseStorage}; @@ -32,10 +32,7 @@ pub(crate) async fn ensure_mcp_connection( ) -> Result<(bool, Vec), Response> { // Check for explicit MCP tools (must error if connection fails) let has_explicit_mcp_tools = tools - .map(|t| { - t.iter() - .any(|tool| matches!(tool.r#type, ResponseToolType::Mcp)) - }) + .map(|t| t.iter().any(|tool| matches!(tool, ResponseTool::Mcp(_)))) .unwrap_or(false); // Check for builtin tools that MAY have MCP routing configured @@ -43,8 +40,8 @@ pub(crate) async fn ensure_mcp_connection( .map(|t| { t.iter().any(|tool| { matches!( - tool.r#type, - ResponseToolType::WebSearchPreview | ResponseToolType::CodeInterpreter + tool, + ResponseTool::WebSearchPreview(_) | ResponseTool::CodeInterpreter(_) ) }) }) @@ -107,23 +104,19 @@ pub(crate) fn validate_worker_availability( None } -/// Extract function tools (and optionally MCP tools) from ResponseTools +/// Extract function tools from ResponseTools /// /// This utility consolidates the logic for extracting tools with schemas from ResponseTools. /// It's used by both Harmony and Regular routers for different purposes: /// -/// - **Harmony router**: Extracts both Function and MCP tools (with `include_mcp: true`) -/// because MCP schemas are populated by convert_mcp_tools_to_response_tools() before the -/// pipeline runs. These tools are used to generate structural constraints in the -/// Harmony preparation stage. +/// - **Harmony router**: Extracts function tools because MCP tools are exposed to the model as +/// function tools (via `convert_mcp_tools_to_response_tools()`), and those are used to +/// generate structural constraints in the Harmony preparation stage. /// -/// - **Regular router**: Extracts only Function tools (with `include_mcp: false`) during -/// the initial conversion from ResponsesRequest to ChatCompletionRequest. MCP tools -/// are merged later by the tool loop before being sent to the chat pipeline, where -/// tool_choice constraints are generated for ALL tools (function + MCP combined). +/// - **Regular router**: Extracts function tools during the initial conversion from +/// ResponsesRequest to ChatCompletionRequest. MCP tools are merged later by the tool loop. pub(crate) fn extract_tools_from_response_tools( response_tools: Option<&[ResponseTool]>, - include_mcp: bool, ) -> Vec { let Some(tools) = response_tools else { return Vec::new(); @@ -131,22 +124,12 @@ pub(crate) fn extract_tools_from_response_tools( tools .iter() - .filter_map(|rt| { - match rt.r#type { - // Function tools: Schema in request - ResponseToolType::Function => rt.function.as_ref().map(|f| Tool { - tool_type: "function".to_string(), - function: f.clone(), - }), - // MCP tools: Schema populated by convert_mcp_tools_to_response_tools() - // Only include if requested (Harmony case) - ResponseToolType::Mcp if include_mcp => rt.function.as_ref().map(|f| Tool { - tool_type: "function".to_string(), - function: f.clone(), - }), - // Hosted tools: No schema available, skip - _ => None, - } + .filter_map(|rt| match rt { + ResponseTool::Function(ft) => Some(Tool { + tool_type: "function".to_string(), + function: ft.function.clone(), + }), + _ => None, }) .collect() } diff --git a/model_gateway/src/routers/grpc/harmony/builder.rs b/model_gateway/src/routers/grpc/harmony/builder.rs index 1220f5c48c..6a95a281c4 100644 --- a/model_gateway/src/routers/grpc/harmony/builder.rs +++ b/model_gateway/src/routers/grpc/harmony/builder.rs @@ -17,8 +17,8 @@ use openai_protocol::{ common::{ChatLogProbs, ContentPart, Tool}, responses::{ ReasoningEffort as ResponsesReasoningEffort, ResponseContentPart, ResponseInput, - ResponseInputOutputItem, ResponseReasoningContent, ResponseTool, ResponseToolType, - ResponsesRequest, StringOrContentParts, + ResponseInputOutputItem, ResponseReasoningContent, ResponseTool, ResponsesRequest, + StringOrContentParts, }, }; use tracing::{debug, trace, warn}; @@ -85,7 +85,7 @@ impl ToolLike for Tool { } fn is_custom(&self) -> bool { - matches!(self.tool_type.as_str(), "mcp" | "function") + matches!(self.tool_type.as_str(), "function") } fn to_tool_description(&self) -> Option { @@ -101,26 +101,24 @@ impl ToolLike for Tool { impl ToolLike for ResponseTool { fn is_builtin(&self) -> bool { matches!( - self.r#type, - ResponseToolType::WebSearchPreview | ResponseToolType::CodeInterpreter + self, + ResponseTool::WebSearchPreview(_) | ResponseTool::CodeInterpreter(_) ) } fn is_custom(&self) -> bool { - matches!( - self.r#type, - ResponseToolType::Mcp | ResponseToolType::Function - ) + matches!(self, ResponseTool::Function(_)) } fn to_tool_description(&self) -> Option { - self.function.as_ref().map(|func| { - ToolDescription::new( - func.name.clone(), - func.description.clone().unwrap_or_default(), - Some(func.parameters.clone()), - ) - }) + match self { + ResponseTool::Function(ft) => Some(ToolDescription::new( + ft.function.name.clone(), + ft.function.description.clone().unwrap_or_default(), + Some(ft.function.parameters.clone()), + )), + _ => None, + } } } @@ -423,11 +421,11 @@ impl HarmonyBuilder { .map(|tools| { tools .iter() - .map(|tool| match tool.r#type { - ResponseToolType::Function => "function", - ResponseToolType::WebSearchPreview => "web_search_preview", - ResponseToolType::CodeInterpreter => "code_interpreter", - ResponseToolType::Mcp => "mcp", + .map(|tool| match tool { + ResponseTool::Function(_) => "function", + ResponseTool::WebSearchPreview(_) => "web_search_preview", + ResponseTool::CodeInterpreter(_) => "code_interpreter", + ResponseTool::Mcp(_) => "mcp", }) .collect() }) diff --git a/model_gateway/src/routers/grpc/harmony/responses/common.rs b/model_gateway/src/routers/grpc/harmony/responses/common.rs index f8167260d6..feb8c58f4e 100644 --- a/model_gateway/src/routers/grpc/harmony/responses/common.rs +++ b/model_gateway/src/routers/grpc/harmony/responses/common.rs @@ -5,8 +5,7 @@ use openai_protocol::{ common::{ToolCall, ToolChoice, ToolChoiceValue}, responses::{ ResponseContentPart, ResponseInput, ResponseInputOutputItem, ResponseOutputItem, - ResponseReasoningContent, ResponseTool, ResponseToolType, ResponsesRequest, - ResponsesResponse, StringOrContentParts, + ResponseReasoningContent, ResponsesRequest, ResponsesResponse, StringOrContentParts, }, }; use serde_json::{from_value, to_string, Value}; @@ -54,20 +53,6 @@ impl McpCallTracking { } } -/// Build a HashSet of MCP tool names for O(1) lookup -/// -/// Creates a HashSet containing the names of all MCP tools in the request, -/// allowing for efficient O(1) lookups when partitioning tool calls. -pub(super) fn build_mcp_tool_names_set( - request_tools: &[ResponseTool], -) -> std::collections::HashSet<&str> { - request_tools - .iter() - .filter(|t| t.r#type == ResponseToolType::Mcp) - .filter_map(|t| t.function.as_ref().map(|f| f.name.as_str())) - .collect() -} - /// Build next request with tool results appended to history /// /// Constructs a new ResponsesRequest with: diff --git a/model_gateway/src/routers/grpc/harmony/responses/non_streaming.rs b/model_gateway/src/routers/grpc/harmony/responses/non_streaming.rs index 53a796db33..95953bafd6 100644 --- a/model_gateway/src/routers/grpc/harmony/responses/non_streaming.rs +++ b/model_gateway/src/routers/grpc/harmony/responses/non_streaming.rs @@ -19,8 +19,7 @@ use tracing::{debug, error, warn}; use super::{ common::{ - build_mcp_tool_names_set, build_next_request_with_tools, inject_mcp_metadata, - load_previous_messages, McpCallTracking, + build_next_request_with_tools, inject_mcp_metadata, load_previous_messages, McpCallTracking, }, execution::{convert_mcp_tools_to_response_tools, execute_mcp_tools, ToolResult}, }; @@ -168,12 +167,12 @@ async fn execute_with_mcp_loop( "Tool calls found - separating MCP and function tools" ); - // Separate MCP and function tool calls based on tool type - let request_tools = current_request.tools.as_deref().unwrap_or(&[]); - let mcp_tool_names = build_mcp_tool_names_set(request_tools); + // Separate MCP and function tool calls based on session exposure. + // MCP tools are exposed to the model as function tools, so the only reliable + // discriminator is whether the name belongs to the MCP session. let (mcp_tool_calls, function_tool_calls): (Vec<_>, Vec<_>) = tool_calls .into_iter() - .partition(|tc| mcp_tool_names.contains(tc.function.name.as_str())); + .partition(|tc| session.has_exposed_tool(&tc.function.name)); debug!( mcp_calls = mcp_tool_calls.len(), diff --git a/model_gateway/src/routers/grpc/harmony/responses/streaming.rs b/model_gateway/src/routers/grpc/harmony/responses/streaming.rs index 09bcbcac06..df28515fc2 100644 --- a/model_gateway/src/routers/grpc/harmony/responses/streaming.rs +++ b/model_gateway/src/routers/grpc/harmony/responses/streaming.rs @@ -4,7 +4,7 @@ use std::time::{SystemTime, UNIX_EPOCH}; use axum::response::Response; use bytes::Bytes; -use openai_protocol::responses::{ResponseToolType, ResponsesRequest}; +use openai_protocol::responses::ResponsesRequest; use serde_json::json; use smg_mcp::{McpServerBinding, McpToolSession}; use tokio::sync::mpsc; @@ -12,10 +12,7 @@ use tracing::{debug, warn}; use uuid::Uuid; use super::{ - common::{ - build_mcp_tool_names_set, build_next_request_with_tools, load_previous_messages, - McpCallTracking, - }, + common::{build_next_request_with_tools, load_previous_messages, McpCallTracking}, execution::{convert_mcp_tools_to_response_tools, execute_mcp_tools}, }; use crate::{ @@ -153,19 +150,6 @@ async fn execute_mcp_tool_loop_streaming( ); } - // Build HashSet of MCP tool names for O(1) lookup during streaming - let mcp_tool_names: std::collections::HashSet = current_request - .tools - .as_ref() - .map(|tools| { - tools - .iter() - .filter(|t| t.r#type == ResponseToolType::Mcp) - .filter_map(|t| t.function.as_ref().map(|f| f.name.clone())) - .collect() - }) - .unwrap_or_default(); - let mut mcp_tracking = McpCallTracking::new(); // Emit mcp_list_tools on first iteration @@ -231,7 +215,6 @@ async fn execute_mcp_tool_loop_streaming( emitter, tx, Some(&session), - Some(&mcp_tool_names), ) .await { @@ -258,12 +241,10 @@ async fn execute_mcp_tool_loop_streaming( "Tool calls found - separating MCP and function tools" ); - // Separate MCP and function tool calls based on tool type - let request_tools = current_request.tools.as_deref().unwrap_or(&[]); - let mcp_tool_names = build_mcp_tool_names_set(request_tools); + // Separate MCP and function tool calls based on session exposure. let (mcp_tool_calls, function_tool_calls): (Vec<_>, Vec<_>) = tool_calls .into_iter() - .partition(|tc| mcp_tool_names.contains(tc.function.name.as_str())); + .partition(|tc| session.has_exposed_tool(&tc.function.name)); debug!( mcp_calls = mcp_tool_calls.len(), @@ -434,7 +415,6 @@ async fn execute_without_mcp_streaming( emitter, tx, None, - None, ) .await { diff --git a/model_gateway/src/routers/grpc/harmony/stages/preparation.rs b/model_gateway/src/routers/grpc/harmony/stages/preparation.rs index 031254e00c..a2c67f478b 100644 --- a/model_gateway/src/routers/grpc/harmony/stages/preparation.rs +++ b/model_gateway/src/routers/grpc/harmony/stages/preparation.rs @@ -141,8 +141,8 @@ impl HarmonyPreparationStage { ctx: &mut RequestContext, request: &ResponsesRequest, ) -> Result, Response> { - // Step 1: Extract function and MCP tools with schemas from ResponseTools - let mut function_tools = extract_tools_from_response_tools(request.tools.as_deref(), true); + // Step 1: Extract function tools with schemas from ResponseTools + let mut function_tools = extract_tools_from_response_tools(request.tools.as_deref()); // Step 2: Filter tools based on tool_choice (AllowedTools or Function) // Note: Tool existence is already validated in ResponsesRequest::validate() diff --git a/model_gateway/src/routers/grpc/harmony/streaming.rs b/model_gateway/src/routers/grpc/harmony/streaming.rs index 782af3dbe2..4eb6ada78d 100644 --- a/model_gateway/src/routers/grpc/harmony/streaming.rs +++ b/model_gateway/src/routers/grpc/harmony/streaming.rs @@ -1,7 +1,7 @@ //! Harmony streaming response processor use std::{ - collections::{hash_map::Entry::Vacant, HashMap, HashSet}, + collections::{hash_map::Entry::Vacant, HashMap}, io, sync::Arc, time::Instant, @@ -552,10 +552,10 @@ impl HarmonyStreamingProcessor { /// Process streaming chunks for Responses API iteration. /// - /// When MCP context is provided (session, mcp_tool_names): + /// When MCP context is provided (session): /// - MCP tools with `ResponseFormat::WebSearchCall` → `web_search_call.*` events /// - Other MCP tools → `mcp_call.*` events - /// - Function tools (not in mcp_tool_names) → `function_call.*` events + /// - Other tools → `function_call.*` events /// /// When no MCP context is provided, all tool calls are treated as function calls. pub async fn process_responses_iteration_stream( @@ -563,24 +563,15 @@ impl HarmonyStreamingProcessor { emitter: &mut ResponseStreamEventEmitter, tx: &mpsc::UnboundedSender>, session: Option<&McpToolSession<'_>>, - mcp_tool_names: Option<&HashSet>, ) -> Result { match execution_result { context::ExecutionResult::Single { stream } => { debug!("Processing Responses API single stream mode"); - Self::process_decode_stream(stream, emitter, tx, session, mcp_tool_names).await + Self::process_decode_stream(stream, emitter, tx, session).await } context::ExecutionResult::Dual { prefill, decode } => { debug!("Processing Responses API dual stream mode"); - Self::process_responses_dual_stream( - prefill, - *decode, - emitter, - tx, - session, - mcp_tool_names, - ) - .await + Self::process_responses_dual_stream(prefill, *decode, emitter, tx, session).await } context::ExecutionResult::Embedding { .. } => { Err("Embeddings not supported in Responses API streaming".to_string()) @@ -594,7 +585,6 @@ impl HarmonyStreamingProcessor { emitter: &mut ResponseStreamEventEmitter, tx: &mpsc::UnboundedSender>, session: Option<&McpToolSession<'_>>, - mcp_tool_names: Option<&HashSet>, ) -> Result { // Phase 1: Drain prefill stream while let Some(result) = prefill_stream.next().await { @@ -602,8 +592,7 @@ impl HarmonyStreamingProcessor { } // Phase 2: Process decode stream - let result = - Self::process_decode_stream(decode_stream, emitter, tx, session, mcp_tool_names).await; + let result = Self::process_decode_stream(decode_stream, emitter, tx, session).await; prefill_stream.mark_completed(); result @@ -615,7 +604,6 @@ impl HarmonyStreamingProcessor { emitter: &mut ResponseStreamEventEmitter, tx: &mpsc::UnboundedSender>, session: Option<&McpToolSession<'_>>, - mcp_tool_names: Option<&HashSet>, ) -> Result { let mut parser = HarmonyParserAdapter::new().map_err(|e| format!("Failed to create parser: {e}"))?; @@ -743,20 +731,14 @@ impl HarmonyStreamingProcessor { .map(|n| n.as_str()) .unwrap_or(""); - // Determine response_format based on MCP context - let response_format = if let Some(names) = mcp_tool_names { - if names.contains(tool_name) { - Some( - session - .map(|s| s.tool_response_format(tool_name)) - .unwrap_or(ResponseFormat::Passthrough), - ) + // Determine response_format based on MCP context. + let response_format = session.and_then(|s| { + if s.has_exposed_tool(tool_name) { + Some(s.tool_response_format(tool_name)) } else { - None // Function tool + None } - } else { - None // No MCP context, treat as function tool - }; + }); // Determine output item type and JSON type string let output_item_type = diff --git a/model_gateway/src/routers/grpc/regular/responses/conversions.rs b/model_gateway/src/routers/grpc/regular/responses/conversions.rs index db41d4e1e5..9307397c5e 100644 --- a/model_gateway/src/routers/grpc/regular/responses/conversions.rs +++ b/model_gateway/src/routers/grpc/regular/responses/conversions.rs @@ -162,11 +162,9 @@ pub(crate) fn responses_to_chat(req: &ResponsesRequest) -> Result BuiltinToolType::WebSearchPreview, - ResponseToolType::CodeInterpreter => BuiltinToolType::CodeInterpreter, - // FileSearch is not in ResponseToolType yet, but we handle it if added + let builtin_type = match tool { + ResponseTool::WebSearchPreview(_) => BuiltinToolType::WebSearchPreview, + ResponseTool::CodeInterpreter(_) => BuiltinToolType::CodeInterpreter, _ => continue, }; @@ -192,9 +191,9 @@ pub fn collect_builtin_routing( pub fn extract_builtin_types(tools: &[ResponseTool]) -> Vec { tools .iter() - .filter_map(|t| match t.r#type { - ResponseToolType::WebSearchPreview => Some(BuiltinToolType::WebSearchPreview), - ResponseToolType::CodeInterpreter => Some(BuiltinToolType::CodeInterpreter), + .filter_map(|t| match t { + ResponseTool::WebSearchPreview(_) => Some(BuiltinToolType::WebSearchPreview), + ResponseTool::CodeInterpreter(_) => Some(BuiltinToolType::CodeInterpreter), _ => None, }) .collect() @@ -254,16 +253,15 @@ pub async fn ensure_request_mcp_client( ) -> Option> { let inputs: Vec = tools .iter() - .filter(|t| matches!(t.r#type, ResponseToolType::Mcp)) - .map(|tool| McpServerInput { - label: tool - .server_label - .clone() - .unwrap_or_else(|| "mcp".to_string()), - url: tool.server_url.clone(), - authorization: tool.authorization.clone(), - headers: tool.headers.clone().unwrap_or_default(), - allowed_tools: tool.allowed_tools.clone(), + .filter_map(|tool| match tool { + ResponseTool::Mcp(mcp) => Some(McpServerInput { + label: mcp.server_label.clone(), + url: mcp.server_url.clone(), + authorization: mcp.authorization.clone(), + headers: mcp.headers.clone().unwrap_or_default(), + allowed_tools: mcp.allowed_tools.clone(), + }), + _ => None, }) .collect(); @@ -276,7 +274,13 @@ pub async fn ensure_request_mcp_client( mod tests { use std::{collections::HashMap, sync::Arc}; - use openai_protocol::responses::ResponseTool; + use openai_protocol::{ + common::Function, + responses::{ + CodeInterpreterTool, FunctionTool, McpTool, ResponseTool, WebSearchPreviewTool, + }, + }; + use serde_json::json; use smg_mcp::{McpConfig, ResponseFormatConfig, ToolConfig}; use super::*; @@ -335,10 +339,9 @@ mod tests { async fn test_collect_builtin_routing_with_configured_server() { let orchestrator = create_test_orchestrator_with_builtin().await; - let tools = vec![ResponseTool { - r#type: ResponseToolType::WebSearchPreview, - ..Default::default() - }]; + let tools = vec![ResponseTool::WebSearchPreview( + WebSearchPreviewTool::default(), + )]; let routing = collect_builtin_routing(&orchestrator, Some(&tools)); @@ -353,10 +356,9 @@ mod tests { async fn test_collect_builtin_routing_no_configured_server() { let orchestrator = create_test_orchestrator_no_builtin().await; - let tools = vec![ResponseTool { - r#type: ResponseToolType::WebSearchPreview, - ..Default::default() - }]; + let tools = vec![ResponseTool::WebSearchPreview( + WebSearchPreviewTool::default(), + )]; let routing = collect_builtin_routing(&orchestrator, Some(&tools)); @@ -368,11 +370,15 @@ mod tests { async fn test_collect_builtin_routing_ignores_mcp_tools() { let orchestrator = create_test_orchestrator_with_builtin().await; - let tools = vec![ResponseTool { - r#type: ResponseToolType::Mcp, + let tools = vec![ResponseTool::Mcp(McpTool { server_url: Some("http://example.com/mcp".to_string()), - ..Default::default() - }]; + authorization: None, + headers: None, + server_label: "mcp".to_string(), + server_description: None, + require_approval: None, + allowed_tools: None, + })]; let routing = collect_builtin_routing(&orchestrator, Some(&tools)); @@ -384,10 +390,14 @@ mod tests { async fn test_collect_builtin_routing_ignores_function_tools() { let orchestrator = create_test_orchestrator_with_builtin().await; - let tools = vec![ResponseTool { - r#type: ResponseToolType::Function, - ..Default::default() - }]; + let tools = vec![ResponseTool::Function(FunctionTool { + function: Function { + name: "dummy".to_string(), + description: None, + parameters: json!({}), + strict: None, + }, + })]; let routing = collect_builtin_routing(&orchestrator, Some(&tools)); @@ -464,14 +474,8 @@ mod tests { let orchestrator = Arc::new(McpOrchestrator::new(config).await.unwrap()); let tools = vec![ - ResponseTool { - r#type: ResponseToolType::WebSearchPreview, - ..Default::default() - }, - ResponseTool { - r#type: ResponseToolType::CodeInterpreter, - ..Default::default() - }, + ResponseTool::WebSearchPreview(WebSearchPreviewTool::default()), + ResponseTool::CodeInterpreter(CodeInterpreterTool::default()), ]; let routing = collect_builtin_routing(&orchestrator, Some(&tools)); @@ -510,10 +514,9 @@ mod tests { let orchestrator = create_test_orchestrator_with_builtin().await; // Request has web_search_preview tool (no server_url, not MCP type) - let tools = vec![ResponseTool { - r#type: ResponseToolType::WebSearchPreview, - ..Default::default() - }]; + let tools = vec![ResponseTool::WebSearchPreview( + WebSearchPreviewTool::default(), + )]; let result = ensure_request_mcp_client(&orchestrator, &tools).await; @@ -534,10 +537,9 @@ mod tests { let orchestrator = create_test_orchestrator_no_builtin().await; // Request has web_search_preview tool - let tools = vec![ResponseTool { - r#type: ResponseToolType::WebSearchPreview, - ..Default::default() - }]; + let tools = vec![ResponseTool::WebSearchPreview( + WebSearchPreviewTool::default(), + )]; let result = ensure_request_mcp_client(&orchestrator, &tools).await; @@ -550,10 +552,14 @@ mod tests { let orchestrator = create_test_orchestrator_with_builtin().await; // Request has only function tools (no MCP, no built-in) - let tools = vec![ResponseTool { - r#type: ResponseToolType::Function, - ..Default::default() - }]; + let tools = vec![ResponseTool::Function(FunctionTool { + function: Function { + name: "dummy".to_string(), + description: None, + parameters: json!({}), + strict: None, + }, + })]; let result = ensure_request_mcp_client(&orchestrator, &tools).await; @@ -568,14 +574,15 @@ mod tests { // Request has mixed tools: function + web_search_preview let tools = vec![ - ResponseTool { - r#type: ResponseToolType::Function, - ..Default::default() - }, - ResponseTool { - r#type: ResponseToolType::WebSearchPreview, - ..Default::default() - }, + ResponseTool::Function(FunctionTool { + function: Function { + name: "dummy".to_string(), + description: None, + parameters: json!({}), + strict: None, + }, + }), + ResponseTool::WebSearchPreview(WebSearchPreviewTool::default()), ]; let result = ensure_request_mcp_client(&orchestrator, &tools).await; diff --git a/model_gateway/src/routers/openai/responses/streaming.rs b/model_gateway/src/routers/openai/responses/streaming.rs index d58d961f19..f1dc9dbba8 100644 --- a/model_gateway/src/routers/openai/responses/streaming.rs +++ b/model_gateway/src/routers/openai/responses/streaming.rs @@ -21,7 +21,7 @@ use openai_protocol::{ is_function_call_type, is_response_event, CodeInterpreterCallEvent, FileSearchCallEvent, FunctionCallEvent, ItemType, McpEvent, OutputItemEvent, ResponseEvent, WebSearchCallEvent, }, - responses::{ResponseToolType, ResponsesRequest}, + responses::{ResponseTool, ResponsesRequest}, }; use serde_json::{json, Value}; use smg_mcp::{McpOrchestrator, McpServerBinding, McpToolSession, ResponseFormat}; @@ -103,11 +103,7 @@ pub(super) fn apply_event_transformations_inplace( .original_request .tools .as_ref() - .map(|tools| { - tools - .iter() - .any(|t| matches!(t.r#type, ResponseToolType::Mcp)) - }) + .map(|tools| tools.iter().any(|t| matches!(t, ResponseTool::Mcp(_)))) .unwrap_or(false); if requested_mcp { @@ -192,7 +188,7 @@ fn build_mcp_tools_value(original_body: &ResponsesRequest) -> Option { let tools = original_body.tools.as_ref()?; let mcp_tools: Vec = tools .iter() - .filter(|t| matches!(t.r#type, ResponseToolType::Mcp)) + .filter(|t| matches!(t, ResponseTool::Mcp(_))) .filter_map(response_tool_to_value) .collect(); diff --git a/model_gateway/src/routers/openai/responses/utils.rs b/model_gateway/src/routers/openai/responses/utils.rs index fb4f992f85..5e297c5592 100644 --- a/model_gateway/src/routers/openai/responses/utils.rs +++ b/model_gateway/src/routers/openai/responses/utils.rs @@ -2,7 +2,7 @@ use openai_protocol::{ event_types::is_response_event, - responses::{ResponseTool, ResponseToolType, ResponsesRequest}, + responses::{ResponseTool, ResponsesRequest}, }; use serde::Serialize; use serde_json::{json, Map, Value}; @@ -208,19 +208,19 @@ pub(super) fn insert_optional_value( /// Handles MCP tools (with server metadata), web_search_preview, and code_interpreter. /// Returns None for function tools and other types that don't need restoration. pub(super) fn response_tool_to_value(tool: &ResponseTool) -> Option { - match tool.r#type { - ResponseToolType::Mcp if tool.server_url.is_some() => { + match tool { + ResponseTool::Mcp(mcp) if mcp.server_url.is_some() => { let mut m = Map::new(); m.insert("type".to_string(), json!("mcp")); - insert_optional_value(&mut m, "server_label", tool.server_label.as_ref()); - insert_optional_value(&mut m, "server_url", tool.server_url.as_ref()); + m.insert("server_label".to_string(), json!(&mcp.server_label)); + insert_optional_value(&mut m, "server_url", mcp.server_url.as_ref()); insert_optional_value( &mut m, "server_description", - tool.server_description.as_ref(), + mcp.server_description.as_ref(), ); - insert_optional_value(&mut m, "require_approval", tool.require_approval.as_ref()); - if let Some(allowed) = &tool.allowed_tools { + insert_optional_value(&mut m, "require_approval", mcp.require_approval.as_ref()); + if let Some(allowed) = &mcp.allowed_tools { m.insert( "allowed_tools".to_string(), Value::Array(allowed.iter().map(|s| json!(s)).collect()), @@ -228,8 +228,8 @@ pub(super) fn response_tool_to_value(tool: &ResponseTool) -> Option { } Some(Value::Object(m)) } - ResponseToolType::WebSearchPreview => Some(json!({"type": "web_search_preview"})), - ResponseToolType::CodeInterpreter => Some(json!({"type": "code_interpreter"})), + ResponseTool::WebSearchPreview(_) => serde_json::to_value(tool).ok(), + ResponseTool::CodeInterpreter(_) => serde_json::to_value(tool).ok(), _ => None, } } diff --git a/model_gateway/tests/api/responses_api_test.rs b/model_gateway/tests/api/responses_api_test.rs index 4be1f682ed..7dc6b0307a 100644 --- a/model_gateway/tests/api/responses_api_test.rs +++ b/model_gateway/tests/api/responses_api_test.rs @@ -4,8 +4,9 @@ use axum::http::StatusCode; use openai_protocol::{ common::{GenerationRequest, ToolChoice, ToolChoiceValue, UsageInfo}, responses::{ - ReasoningEffort, RequireApproval, ResponseInput, ResponseReasoningParam, ResponseTool, - ResponseToolType, ResponsesRequest, ServiceTier, Truncation, + CodeInterpreterTool, McpTool, ReasoningEffort, RequireApproval, ResponseInput, + ResponseReasoningParam, ResponseTool, ResponsesRequest, ServiceTier, Truncation, + WebSearchPreviewTool, }, }; use smg::{ @@ -81,17 +82,15 @@ async fn test_non_streaming_mcp_minimal_e2e_with_persistence() { stream: Some(false), temperature: Some(0.2), tool_choice: Some(ToolChoice::default()), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Mcp, - function: None, + tools: Some(vec![ResponseTool::Mcp(McpTool { server_url: Some(mcp.url()), authorization: None, headers: None, - server_label: Some("mock".to_string()), + server_label: "mock".to_string(), server_description: None, require_approval: None, allowed_tools: None, - }]), + })]), top_logprobs: Some(0), top_p: None, truncation: Some(Truncation::Disabled), @@ -311,10 +310,9 @@ fn test_responses_request_creation() { stream: Some(false), temperature: Some(0.7), tool_choice: Some(ToolChoice::Value(ToolChoiceValue::Auto)), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::WebSearchPreview, - ..Default::default() - }]), + tools: Some(vec![ResponseTool::WebSearchPreview( + WebSearchPreviewTool::default(), + )]), top_logprobs: Some(5), top_p: Some(0.9), truncation: Some(Truncation::Disabled), @@ -470,10 +468,9 @@ fn test_json_serialization() { stream: Some(true), temperature: Some(0.9), tool_choice: Some(ToolChoice::Value(ToolChoiceValue::Required)), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::CodeInterpreter, - ..Default::default() - }]), + tools: Some(vec![ResponseTool::CodeInterpreter( + CodeInterpreterTool::default(), + )]), top_logprobs: Some(10), top_p: Some(0.8), truncation: Some(Truncation::Auto), @@ -573,14 +570,15 @@ async fn test_multi_turn_loop_with_mcp() { stream: Some(false), temperature: Some(0.7), tool_choice: Some(ToolChoice::Value(ToolChoiceValue::Auto)), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Mcp, + tools: Some(vec![ResponseTool::Mcp(McpTool { server_url: Some(mcp.url()), - server_label: Some("mock".to_string()), + authorization: None, + headers: None, + server_label: "mock".to_string(), server_description: Some("Mock MCP server for testing".to_string()), require_approval: Some(RequireApproval::Never), - ..Default::default() - }]), + allowed_tools: None, + })]), top_logprobs: Some(0), top_p: Some(1.0), truncation: Some(Truncation::Disabled), @@ -725,12 +723,15 @@ async fn test_max_tool_calls_limit() { stream: Some(false), temperature: Some(0.7), tool_choice: Some(ToolChoice::Value(ToolChoiceValue::Auto)), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Mcp, + tools: Some(vec![ResponseTool::Mcp(McpTool { server_url: Some(mcp.url()), - server_label: Some("mock".to_string()), - ..Default::default() - }]), + authorization: None, + headers: None, + server_label: "mock".to_string(), + server_description: None, + require_approval: None, + allowed_tools: None, + })]), top_logprobs: Some(0), top_p: Some(1.0), truncation: Some(Truncation::Disabled), @@ -902,14 +903,15 @@ async fn test_streaming_with_mcp_tool_calls() { stream: Some(true), // KEY: Enable streaming temperature: Some(0.7), tool_choice: Some(ToolChoice::Value(ToolChoiceValue::Auto)), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Mcp, + tools: Some(vec![ResponseTool::Mcp(McpTool { server_url: Some(mcp.url()), - server_label: Some("mock".to_string()), + authorization: None, + headers: None, + server_label: "mock".to_string(), server_description: Some("Mock MCP for streaming test".to_string()), require_approval: Some(RequireApproval::Never), - ..Default::default() - }]), + allowed_tools: None, + })]), top_logprobs: Some(0), top_p: Some(1.0), truncation: Some(Truncation::Disabled), @@ -1179,12 +1181,15 @@ async fn test_streaming_multi_turn_with_mcp() { stream: Some(true), temperature: Some(0.8), tool_choice: Some(ToolChoice::Value(ToolChoiceValue::Auto)), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Mcp, + tools: Some(vec![ResponseTool::Mcp(McpTool { server_url: Some(mcp.url()), - server_label: Some("mock".to_string()), - ..Default::default() - }]), + authorization: None, + headers: None, + server_label: "mock".to_string(), + server_description: None, + require_approval: None, + allowed_tools: None, + })]), top_logprobs: Some(0), top_p: Some(1.0), truncation: Some(Truncation::Disabled), diff --git a/model_gateway/tests/spec/responses.rs b/model_gateway/tests/spec/responses.rs index 168d5d388b..eb2dcdea45 100644 --- a/model_gateway/tests/spec/responses.rs +++ b/model_gateway/tests/spec/responses.rs @@ -1,7 +1,7 @@ use openai_protocol::{ common::{Function, StringOrArray, ToolChoice, ToolChoiceValue}, responses::{ - IncludeField, ResponseInput, ResponseInputOutputItem, ResponseTool, ResponseToolType, + FunctionTool, IncludeField, McpTool, ResponseInput, ResponseInputOutputItem, ResponseTool, ResponsesRequest, StringOrContentParts, TextConfig, TextFormat, }, }; @@ -601,26 +601,32 @@ fn test_validate_stop_sequences_non_empty() { /// Test tools validation (function tool must have function) #[test] fn test_validate_tools_function_missing() { - let request = ResponsesRequest { - input: ResponseInput::Text("test".to_string()), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Function, - function: None, // Missing function definition - server_url: None, - authorization: None, - headers: None, - server_label: None, - server_description: None, - require_approval: None, - allowed_tools: None, - }]), - ..Default::default() - }; - let result = request.validate(); - assert!( - result.is_err(), - "Function tool without function definition should be invalid" - ); + // With type-discriminated tools, a function tool without the required fields + // should fail deserialization. + let v = json!({ + "input": "test", + "tools": [{ "type": "function" }] + }); + let parsed: Result = serde_json::from_value(v); + assert!(parsed.is_err(), "Expected deserialization to fail"); +} + +#[test] +fn test_deserialize_function_tool_rejects_unknown_fields() { + let v = json!({ + "input": "test", + "tools": [ + { + "type": "function", + "name": "test_func", + "parameters": {}, + "extra_field": 1 + } + ] + }); + + let parsed: Result = serde_json::from_value(v); + assert!(parsed.is_err(), "Expected deserialization to fail"); } /// Test tools validation (valid MCP tool should pass) @@ -628,17 +634,15 @@ fn test_validate_tools_function_missing() { fn test_validate_tools_mcp_valid_ok() { let request = ResponsesRequest { input: ResponseInput::Text("test".to_string()), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Mcp, - function: None, + tools: Some(vec![ResponseTool::Mcp(McpTool { server_url: None, // server_url is optional when server_label is provided authorization: None, headers: None, - server_label: Some("mock".to_string()), + server_label: "mock".to_string(), server_description: None, require_approval: None, allowed_tools: None, - }]), + })]), ..Default::default() }; @@ -651,37 +655,12 @@ fn test_validate_tools_mcp_valid_ok() { /// Test tools validation (MCP tool must have server_label; server_url is optional) #[test] fn test_validate_tools_mcp_missing_server_label() { - let request = ResponsesRequest { - input: ResponseInput::Text("test".to_string()), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Mcp, - function: None, - server_url: None, - authorization: None, - headers: None, - server_label: None, // Missing required server_label - server_description: None, - require_approval: None, - allowed_tools: None, - }]), - ..Default::default() - }; - let result = request.validate(); - assert!( - result.is_err(), - "MCP tool without server_label should be invalid" - ); - - let err = result.unwrap_err(); - let s = format!("{err:?}"); - assert!( - s.contains("missing_required_parameter"), - "Expected error code missing_required_parameter, got: {err:?}", - ); - assert!( - s.contains("tools[0].server_label"), - "Expected error to reference tools[0].server_label, got: {err:?}", - ); + let v = json!({ + "input": "test", + "tools": [{ "type": "mcp" }] + }); + let parsed: Result = serde_json::from_value(v); + assert!(parsed.is_err(), "Expected deserialization to fail"); } /// Test tools validation (MCP tool server_label must be unique; case-insensitive) @@ -690,28 +669,24 @@ fn test_validate_tools_mcp_duplicate_server_label() { let request = ResponsesRequest { input: ResponseInput::Text("test".to_string()), tools: Some(vec![ - ResponseTool { - r#type: ResponseToolType::Mcp, - function: None, + ResponseTool::Mcp(McpTool { server_url: None, authorization: None, headers: None, - server_label: Some("Foo".to_string()), + server_label: "Foo".to_string(), server_description: None, require_approval: None, allowed_tools: None, - }, - ResponseTool { - r#type: ResponseToolType::Mcp, - function: None, + }), + ResponseTool::Mcp(McpTool { server_url: None, authorization: None, headers: None, - server_label: Some("foo".to_string()), + server_label: "foo".to_string(), server_description: None, require_approval: None, allowed_tools: None, - }, + }), ]), ..Default::default() }; @@ -737,17 +712,15 @@ fn test_validate_tools_mcp_server_label_invalid_cases() { for label in invalid_labels { let request = ResponsesRequest { input: ResponseInput::Text("test".to_string()), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Mcp, - function: None, + tools: Some(vec![ResponseTool::Mcp(McpTool { server_url: Some("https://example.com/mcp".to_string()), authorization: None, headers: None, - server_label: Some(label.to_string()), + server_label: label.to_string(), server_description: None, require_approval: None, allowed_tools: None, - }]), + })]), ..Default::default() }; @@ -797,22 +770,14 @@ fn test_validate_tool_choice_requires_tools() { // Valid: tool_choice with tools let request = ResponsesRequest { input: ResponseInput::Text("test".to_string()), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Function, - function: Some(Function { + tools: Some(vec![ResponseTool::Function(FunctionTool { + function: Function { name: "test_func".to_string(), description: None, parameters: json!({}), strict: None, - }), - server_url: None, - authorization: None, - headers: None, - server_label: None, - server_description: None, - require_approval: None, - allowed_tools: None, - }]), + }, + })]), tool_choice: Some(ToolChoice::Value(ToolChoiceValue::Auto)), ..Default::default() }; @@ -1056,22 +1021,14 @@ fn test_normalize_tool_choice_auto() { let mut request = ResponsesRequest { input: ResponseInput::Text("test".to_string()), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Function, - function: Some(Function { + tools: Some(vec![ResponseTool::Function(FunctionTool { + function: Function { name: "test_func".to_string(), description: None, parameters: json!({}), strict: None, - }), - server_url: None, - authorization: None, - headers: None, - server_label: None, - server_description: None, - require_approval: None, - allowed_tools: None, - }]), + }, + })]), tool_choice: None, ..Default::default() }; @@ -1125,22 +1082,14 @@ fn test_normalize_tool_choice_no_override() { let mut request = ResponsesRequest { input: ResponseInput::Text("test".to_string()), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Function, - function: Some(Function { + tools: Some(vec![ResponseTool::Function(FunctionTool { + function: Function { name: "test_func".to_string(), description: None, parameters: json!({}), strict: None, - }), - server_url: None, - authorization: None, - headers: None, - server_label: None, - server_description: None, - require_approval: None, - allowed_tools: None, - }]), + }, + })]), tool_choice: Some(ToolChoice::Value(ToolChoiceValue::Required)), ..Default::default() }; @@ -1163,22 +1112,14 @@ fn test_normalize_parallel_tool_calls() { let mut request = ResponsesRequest { input: ResponseInput::Text("test".to_string()), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Function, - function: Some(Function { + tools: Some(vec![ResponseTool::Function(FunctionTool { + function: Function { name: "test_func".to_string(), description: None, parameters: json!({}), strict: None, - }), - server_url: None, - authorization: None, - headers: None, - server_label: None, - server_description: None, - require_approval: None, - allowed_tools: None, - }]), + }, + })]), parallel_tool_calls: None, ..Default::default() }; @@ -1223,22 +1164,14 @@ fn test_normalize_parallel_tool_calls_no_override() { let mut request = ResponsesRequest { input: ResponseInput::Text("test".to_string()), - tools: Some(vec![ResponseTool { - r#type: ResponseToolType::Function, - function: Some(Function { + tools: Some(vec![ResponseTool::Function(FunctionTool { + function: Function { name: "test_func".to_string(), description: None, parameters: json!({}), strict: None, - }), - server_url: None, - authorization: None, - headers: None, - server_label: None, - server_description: None, - require_approval: None, - allowed_tools: None, - }]), + }, + })]), parallel_tool_calls: Some(false), ..Default::default() }; diff --git a/protocols/src/responses.rs b/protocols/src/responses.rs index 2ce8f6e7f4..4762c8e8a3 100644 --- a/protocols/src/responses.rs +++ b/protocols/src/responses.rs @@ -20,41 +20,64 @@ use crate::{builders::ResponsesResponseBuilder, validated::Normalizable}; // Response Tools (MCP and others) // ============================================================================ +#[derive(Debug, Clone, Deserialize, Serialize)] +#[serde(tag = "type")] +#[serde(rename_all = "snake_case")] +pub enum ResponseTool { + /// Function tool. + #[serde(rename = "function")] + Function(FunctionTool), + + /// Built-in tool. + #[serde(rename = "web_search_preview")] + WebSearchPreview(WebSearchPreviewTool), + + /// Built-in tool. + #[serde(rename = "code_interpreter")] + CodeInterpreter(CodeInterpreterTool), + + /// MCP server tool. + #[serde(rename = "mcp")] + Mcp(McpTool), +} + #[serde_with::skip_serializing_none] #[derive(Debug, Clone, Deserialize, Serialize)] -pub struct ResponseTool { - #[serde(rename = "type")] - pub r#type: ResponseToolType, - // Function tool fields (used when type == "function") - // In Responses API, function fields are flattened at the top level +#[serde(deny_unknown_fields)] +pub struct FunctionTool { + /// Flatten to match Responses API tool JSON shape. #[serde(flatten)] - pub function: Option, - // MCP-specific fields (used when type == "mcp") + pub function: Function, +} + +#[serde_with::skip_serializing_none] +#[derive(Debug, Clone, Deserialize, Serialize)] +#[serde(deny_unknown_fields)] +pub struct McpTool { pub server_url: Option, pub authorization: Option, /// Custom headers to send to MCP server (from request payload, not HTTP headers) pub headers: Option>, - pub server_label: Option, + pub server_label: String, pub server_description: Option, /// Approval requirement configuration for MCP tools. pub require_approval: Option, pub allowed_tools: Option>, } -impl Default for ResponseTool { - fn default() -> Self { - Self { - r#type: ResponseToolType::WebSearchPreview, - function: None, - server_url: None, - authorization: None, - headers: None, - server_label: None, - server_description: None, - require_approval: None, - allowed_tools: None, - } - } +#[serde_with::skip_serializing_none] +#[derive(Debug, Clone, Deserialize, Serialize, Default)] +#[serde(deny_unknown_fields)] +pub struct WebSearchPreviewTool { + pub search_context_size: Option, + pub user_location: Option, +} + +#[serde_with::skip_serializing_none] +#[derive(Debug, Clone, Deserialize, Serialize, Default)] +#[serde(deny_unknown_fields)] +pub struct CodeInterpreterTool { + pub container: Option, } /// `require_approval` values. @@ -65,15 +88,6 @@ pub enum RequireApproval { Never, } -#[derive(Debug, Clone, Deserialize, Serialize, PartialEq)] -#[serde(rename_all = "snake_case")] -pub enum ResponseToolType { - Function, - WebSearchPreview, - CodeInterpreter, - Mcp, -} - // ============================================================================ // Reasoning Parameters // ============================================================================ @@ -965,8 +979,8 @@ fn validate_tool_choice_with_tools(request: &ResponsesRequest) -> Result<(), Val }; let function_tool_names: Vec<&str> = tools .iter() - .filter_map(|t| match t.r#type { - ResponseToolType::Function => t.function.as_ref().map(|f| f.name.as_str()), + .filter_map(|t| match t { + ResponseTool::Function(ft) => Some(ft.function.name.as_str()), _ => None, }) .collect(); @@ -1162,53 +1176,42 @@ fn validate_response_tools(tools: &[ResponseTool]) -> Result<(), ValidationError let mut seen_mcp_labels: HashSet = HashSet::new(); for (idx, tool) in tools.iter().enumerate() { - match tool.r#type { - ResponseToolType::Function => { - if tool.function.is_none() { - let mut e = ValidationError::new("function_tool_missing_function"); - e.message = Some("Function tool must have a function definition".into()); - return Err(e); - } + if let ResponseTool::Mcp(mcp) = tool { + let raw_label = mcp.server_label.as_str(); + if raw_label.is_empty() { + let mut e = ValidationError::new("missing_required_parameter"); + e.message = Some( + format!("Missing required parameter: 'tools[{idx}].server_label'.").into(), + ); + return Err(e); } - ResponseToolType::Mcp => { - let Some(raw_label) = tool.server_label.as_deref().filter(|s| !s.is_empty()) else { - let mut e = ValidationError::new("missing_required_parameter"); - e.message = Some( - format!("Missing required parameter: 'tools[{idx}].server_label'.").into(), - ); - return Err(e); - }; - // OpenAI spec-compatible validation: require a non-empty label that starts with a - // letter and contains only letters, digits, '-' and '_'. - let valid = raw_label.starts_with(|c: char| c.is_ascii_alphabetic()) - && raw_label - .chars() - .all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_'); - if !valid { - let mut e = ValidationError::new("invalid_server_label"); - e.message = Some( - format!( - "Invalid input {raw_label}: 'server_label' must start with a letter and consist of only letters, digits, '-' and '_'" - ) - .into(), - ); - return Err(e); - } + // OpenAI spec-compatible validation: require a non-empty label that starts with a + // letter and contains only letters, digits, '-' and '_'. + let valid = raw_label.starts_with(|c: char| c.is_ascii_alphabetic()) + && raw_label + .chars() + .all(|c| c.is_ascii_alphanumeric() || c == '-' || c == '_'); + if !valid { + let mut e = ValidationError::new("invalid_server_label"); + e.message = Some( + format!( + "Invalid input {raw_label}: 'server_label' must start with a letter and consist of only letters, digits, '-' and '_'" + ) + .into(), + ); + return Err(e); + } - let normalized = raw_label.to_lowercase(); - if !seen_mcp_labels.insert(normalized) { - let mut e = ValidationError::new("mcp_tool_duplicate_server_label"); - e.message = Some( - format!( - "Duplicate MCP server_label '{raw_label}' found in 'tools' parameter." - ) + let normalized = raw_label.to_lowercase(); + if !seen_mcp_labels.insert(normalized) { + let mut e = ValidationError::new("mcp_tool_duplicate_server_label"); + e.message = Some( + format!("Duplicate MCP server_label '{raw_label}' found in 'tools' parameter.") .into(), - ); - return Err(e); - } + ); + return Err(e); } - _ => {} } } Ok(()) From 04cb880bd6013163ea9b08a024ed3062ba6ab8fe Mon Sep 17 00:00:00 2001 From: Ziwen Zhao Date: Tue, 24 Feb 2026 15:09:38 -0800 Subject: [PATCH 2/3] refactor(protocol): model ResponseTool as tagged enum to match Responses spec and tighten MCP validation Signed-off-by: Ziwen Zhao --- mcp/src/responses_bridge.rs | 8 ++++---- model_gateway/src/routers/openai/responses/utils.rs | 2 +- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/mcp/src/responses_bridge.rs b/mcp/src/responses_bridge.rs index 42c456a2ed..702478d05a 100644 --- a/mcp/src/responses_bridge.rs +++ b/mcp/src/responses_bridge.rs @@ -73,15 +73,15 @@ pub fn build_chat_function_tools_with_names( .collect() } -/// Build Responses API MCP tools from MCP tool entries. +/// Build Responses API function tools from MCP tool entries. /// -/// These tools are exposed in Responses requests where MCP tools are represented -/// as `{"type": "mcp", ...}` tool entries. +/// MCP tools are exposed to the model as function tools in the Responses API, +/// so these serialize as `{"type": "function", ...}` tool entries. pub fn build_response_tools(entries: &[ToolEntry]) -> Vec { build_response_tools_with_names(entries, None) } -/// Build Responses API MCP tools from MCP tool entries with optional exposed names. +/// Build Responses API function tools from MCP tool entries with optional exposed names. pub fn build_response_tools_with_names( entries: &[ToolEntry], exposed_names: Option<&std::collections::HashMap>, diff --git a/model_gateway/src/routers/openai/responses/utils.rs b/model_gateway/src/routers/openai/responses/utils.rs index 5e297c5592..550416ba68 100644 --- a/model_gateway/src/routers/openai/responses/utils.rs +++ b/model_gateway/src/routers/openai/responses/utils.rs @@ -209,7 +209,7 @@ pub(super) fn insert_optional_value( /// Returns None for function tools and other types that don't need restoration. pub(super) fn response_tool_to_value(tool: &ResponseTool) -> Option { match tool { - ResponseTool::Mcp(mcp) if mcp.server_url.is_some() => { + ResponseTool::Mcp(mcp) => { let mut m = Map::new(); m.insert("type".to_string(), json!("mcp")); m.insert("server_label".to_string(), json!(&mcp.server_label)); From 8c84c40212dc5a0340e245477ada46b26e59efec Mon Sep 17 00:00:00 2001 From: Ziwen Zhao Date: Tue, 24 Feb 2026 16:10:24 -0800 Subject: [PATCH 3/3] refactor(protocol): model ResponseTool as tagged enum to match Responses spec and tighten MCP validation Signed-off-by: Ziwen Zhao --- model_gateway/src/routers/openai/responses/utils.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/model_gateway/src/routers/openai/responses/utils.rs b/model_gateway/src/routers/openai/responses/utils.rs index 550416ba68..0ea0bd50da 100644 --- a/model_gateway/src/routers/openai/responses/utils.rs +++ b/model_gateway/src/routers/openai/responses/utils.rs @@ -230,7 +230,7 @@ pub(super) fn response_tool_to_value(tool: &ResponseTool) -> Option { } ResponseTool::WebSearchPreview(_) => serde_json::to_value(tool).ok(), ResponseTool::CodeInterpreter(_) => serde_json::to_value(tool).ok(), - _ => None, + ResponseTool::Function(_) => None, } }