Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion model_gateway/src/routers/grpc/common/responses/utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,7 @@ pub(crate) async fn ensure_mcp_connection(

if let Some(tools) = tools {
match ensure_request_mcp_client(mcp_orchestrator, tools).await {
Some((_orchestrator, mcp_servers)) => {
Some(mcp_servers) => {
return Ok((true, mcp_servers));
}
None => {
Expand Down
28 changes: 5 additions & 23 deletions model_gateway/src/routers/mcp_utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,15 +12,6 @@ use crate::{
/// Default maximum tool loop iterations (safety limit).
pub const DEFAULT_MAX_ITERATIONS: usize = 10;

/// Configuration for MCP tool calling loops.
#[derive(Debug, Clone)]
pub struct McpLoopConfig {
/// Maximum iterations (default: DEFAULT_MAX_ITERATIONS).
pub max_iterations: usize,
/// MCP servers for this request (label, server_key).
pub mcp_servers: Vec<(String, String)>,
}

/// Routing information for a built-in tool type.
///
/// When a built-in tool type (web_search_preview, code_interpreter, file_search)
Expand All @@ -37,15 +28,6 @@ pub struct BuiltinToolRouting {
pub response_format: ResponseFormat,
}

impl Default for McpLoopConfig {
fn default() -> Self {
Self {
max_iterations: DEFAULT_MAX_ITERATIONS,
mcp_servers: Vec::new(),
}
}
}

/// Collect routing information for built-in tools in a request.
///
/// Scans request tools for built-in types (web_search_preview, code_interpreter, file_search)
Expand Down Expand Up @@ -115,12 +97,12 @@ pub fn collect_builtin_routing(
///
/// Headers for MCP servers come from the tool payload (`tool.headers`), not HTTP request headers.
///
/// Returns `Some((orchestrator, mcp_servers))` if MCP tools or built-in routing is available,
/// Returns `Some(mcp_servers)` if MCP tools or built-in routing is available,
/// `None` otherwise.
pub async fn ensure_request_mcp_client(
mcp_orchestrator: &Arc<McpOrchestrator>,
tools: &[ResponseTool],
) -> Option<(Arc<McpOrchestrator>, Vec<(String, String)>)> {
) -> Option<Vec<(String, String)>> {
Comment thread
CatherineSue marked this conversation as resolved.
let mut mcp_servers = Vec::new();

// 1. Process explicit MCP tools (dynamic via `server_url`, or static via `server_label`)
Expand Down Expand Up @@ -217,7 +199,7 @@ pub async fn ensure_request_mcp_client(
if mcp_servers.is_empty() {
None
} else {
Some((mcp_orchestrator.clone(), mcp_servers))
Some(mcp_servers)
}
}

Expand Down Expand Up @@ -470,7 +452,7 @@ mod tests {
// Should return Some because built-in routing is configured
assert!(result.is_some());

let (_, mcp_servers) = result.unwrap();
let mcp_servers = result.unwrap();
assert_eq!(mcp_servers.len(), 1);

// The server key should be the static server name
Expand Down Expand Up @@ -534,7 +516,7 @@ mod tests {
// Should return Some because web_search_preview has built-in routing
assert!(result.is_some());

let (_, mcp_servers) = result.unwrap();
let mcp_servers = result.unwrap();
assert_eq!(mcp_servers.len(), 1);
assert_eq!(mcp_servers[0].0, "search-server");
}
Expand Down
25 changes: 10 additions & 15 deletions model_gateway/src/routers/openai/responses/mcp.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
//! - Payload transformation for MCP tool interception
//! - Metadata injection for MCP operations

use std::{io, slice, sync::Arc};
use std::{io, slice};

use axum::http::HeaderMap;
use bytes::Bytes;
Expand All @@ -17,15 +17,15 @@ use tokio::sync::mpsc;
use tracing::{debug, info, warn};

use crate::{
mcp::{McpOrchestrator, McpToolSession, ResponseFormat, ResponseTransformer},
mcp::{McpToolSession, ResponseFormat, ResponseTransformer},
protocols::{
event_types::{
is_function_call_type, CodeInterpreterCallEvent, FileSearchCallEvent, ItemType,
McpEvent, OutputItemEvent, WebSearchCallEvent,
},
responses::{generate_id, ResponseInput, ResponsesRequest},
},
routers::{header_utils::apply_request_headers, mcp_utils::McpLoopConfig},
routers::{header_utils::apply_request_headers, mcp_utils::DEFAULT_MAX_ITERATIONS},
};

// ============================================================================
Expand Down Expand Up @@ -278,11 +278,7 @@ pub(super) async fn execute_streaming_tool_calls(
///
/// Retains existing function tools from the request, removes non-function tools
/// (MCP, builtin), and appends function tools for discovered MCP server tools.
pub(super) fn prepare_mcp_tools_as_functions(
payload: &mut Value,
orchestrator: &Arc<McpOrchestrator>,
server_keys: &[String],
) {
pub(super) fn prepare_mcp_tools_as_functions(payload: &mut Value, session: &McpToolSession<'_>) {
let Some(obj) = payload.as_object_mut() else {
return;
};
Expand All @@ -302,7 +298,7 @@ pub(super) fn prepare_mcp_tools_as_functions(
}
}

let mcp_tools = orchestrator.list_tools_for_servers(server_keys);
let mcp_tools = session.mcp_tools();
let mut tools_json = Vec::with_capacity(retained_tools.len() + mcp_tools.len());
tools_json.append(&mut retained_tools);

Expand Down Expand Up @@ -634,7 +630,6 @@ pub(super) async fn execute_tool_loop(
initial_payload: Value,
original_body: &ResponsesRequest,
session: &McpToolSession<'_>,
config: &McpLoopConfig,
) -> Result<Value, String> {
let mut state = ToolLoopState::new(original_body.input.clone());

Expand All @@ -648,7 +643,7 @@ pub(super) async fn execute_tool_loop(

info!(
"Starting tool loop: max_tool_calls={:?}, max_iterations={}",
max_tool_calls, config.max_iterations
max_tool_calls, DEFAULT_MAX_ITERATIONS
);

loop {
Expand Down Expand Up @@ -688,8 +683,8 @@ pub(super) async fn execute_tool_loop(

// Check combined limit: use minimum of user's max_tool_calls (if set) and safety max_iterations
let effective_limit = match max_tool_calls {
Some(user_max) => user_max.min(config.max_iterations),
None => config.max_iterations,
Some(user_max) => user_max.min(DEFAULT_MAX_ITERATIONS),
None => DEFAULT_MAX_ITERATIONS,
};

if state.total_calls > effective_limit {
Expand All @@ -699,13 +694,13 @@ pub(super) async fn execute_tool_loop(
} else {
warn!(
"Reached safety max_iterations limit: {}",
config.max_iterations
DEFAULT_MAX_ITERATIONS
);
}
} else {
warn!(
"Reached safety max_iterations limit: {}",
config.max_iterations
DEFAULT_MAX_ITERATIONS
);
}

Expand Down
19 changes: 6 additions & 13 deletions model_gateway/src/routers/openai/responses/non_streaming.rs
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ use crate::{
mcp::McpToolSession,
routers::{
header_utils::{apply_provider_headers, extract_auth_header},
mcp_utils::{ensure_request_mcp_client, McpLoopConfig},
mcp_utils::ensure_request_mcp_client,
openai::context::{PayloadState, RequestContext},
persistence_utils::persist_conversation_items,
},
Expand Down Expand Up @@ -57,28 +57,22 @@ pub async fn handle_non_streaming_response(mut ctx: RequestContext) -> Response
}
};

// Check for MCP tools and create request context if needed
let mcp_result = if let Some(tools) = original_body.tools.as_deref() {
// Check for MCP tools and create session if needed
let mcp_servers = if let Some(tools) = original_body.tools.as_deref() {
ensure_request_mcp_client(mcp_orchestrator, tools).await
} else {
None
};

let mut response_json: Value;

if let Some((orchestrator, mcp_servers)) = mcp_result {
let server_keys: Vec<String> = mcp_servers.iter().map(|(_, key)| key.clone()).collect();
let config = McpLoopConfig {
mcp_servers: mcp_servers.clone(),
..McpLoopConfig::default()
};
prepare_mcp_tools_as_functions(&mut payload, &orchestrator, &server_keys);

if let Some(mcp_servers) = mcp_servers {
let session_request_id = original_body
.request_id
.clone()
.unwrap_or_else(|| format!("req_{}", uuid::Uuid::new_v4()));
let session = McpToolSession::new(&orchestrator, mcp_servers, &session_request_id);
let session = McpToolSession::new(mcp_orchestrator, mcp_servers, &session_request_id);
prepare_mcp_tools_as_functions(&mut payload, &session);

match execute_tool_loop(
ctx.components.client(),
Expand All @@ -87,7 +81,6 @@ pub async fn handle_non_streaming_response(mut ctx: RequestContext) -> Response
payload,
original_body,
&session,
&config,
)
.await
{
Expand Down
54 changes: 24 additions & 30 deletions model_gateway/src/routers/openai/responses/streaming.rs
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,7 @@ use crate::{
},
routers::{
header_utils::{apply_request_headers, preserve_response_headers},
mcp_utils::McpLoopConfig,
mcp_utils::DEFAULT_MAX_ITERATIONS,
openai::context::{RequestContext, StreamingEventContext, StreamingRequest},
persistence_utils::persist_conversation_items,
},
Expand Down Expand Up @@ -685,16 +685,9 @@ pub(super) async fn handle_streaming_with_tool_interception(
headers: Option<&HeaderMap>,
req: StreamingRequest,
orchestrator: &Arc<McpOrchestrator>,
loop_config: McpLoopConfig,
mcp_servers: Vec<(String, String)>,
) -> Response {
let server_keys: Vec<String> = loop_config
.mcp_servers
.iter()
.map(|(_, key)| key.clone())
.collect();
// Transform MCP tools to function tools in payload
let mut payload = req.payload;
prepare_mcp_tools_as_functions(&mut payload, orchestrator, &server_keys);
let payload = req.payload;

let (tx, rx) = mpsc::unbounded_channel::<Result<Bytes, io::Error>>();
let should_store = req.original_body.store.unwrap_or(false);
Expand All @@ -714,23 +707,26 @@ pub(super) async fn handle_streaming_with_tool_interception(
tokio::spawn(async move {
let mut state = ToolLoopState::new(original_request.input.clone());
let max_tool_calls = original_request.max_tool_calls.map(|n| n as usize);
let tools_json = payload_clone.get("tools").cloned().unwrap_or(json!([]));
let base_payload = payload_clone.clone();
let mut current_payload = payload_clone;
let mut mcp_list_tools_sent = false;
let mut is_first_iteration = true;
let mut sequence_number: u64 = 0;
let mut next_output_index: usize = 0;
let mut preserved_response_id: Option<String> = None;

// Create session inside spawned task (borrows from orchestrator_clone which lives in closure)
let session_request_id = format!("resp_{}", uuid::Uuid::new_v4());
let session = McpToolSession::new(
&orchestrator_clone,
loop_config.mcp_servers.clone(),
mcp_servers.clone(),
&session_request_id,
);

// Transform MCP tools to function tools in payload
let mut current_payload = payload_clone;
prepare_mcp_tools_as_functions(&mut current_payload, &session);
let tools_json = current_payload.get("tools").cloned().unwrap_or(json!([]));
let base_payload = current_payload.clone();
let mut mcp_list_tools_sent = false;
let mut is_first_iteration = true;
let mut sequence_number: u64 = 0;
let mut next_output_index: usize = 0;
let mut preserved_response_id: Option<String> = None;

let streaming_ctx = StreamingEventContext {
original_request: &original_request,
previous_response_id: previous_response_id.as_deref(),
Expand Down Expand Up @@ -959,8 +955,8 @@ pub(super) async fn handle_streaming_with_tool_interception(
state.total_calls += pending_calls.len();

let effective_limit = match max_tool_calls {
Some(user_max) => user_max.min(loop_config.max_iterations),
None => loop_config.max_iterations,
Some(user_max) => user_max.min(DEFAULT_MAX_ITERATIONS),
None => DEFAULT_MAX_ITERATIONS,
};

if state.total_calls > effective_limit {
Expand Down Expand Up @@ -1032,11 +1028,12 @@ pub async fn handle_streaming_response(ctx: RequestContext) -> Response {
let mcp_orchestrator = ctx
.components
.mcp_orchestrator()
.expect("MCP orchestrator required");
.expect("MCP orchestrator required")
.clone();

// Check for MCP tools and create request context if needed
let mcp_result = if let Some(tools) = original_body.tools.as_deref() {
ensure_request_mcp_client(mcp_orchestrator, tools).await
let mcp_servers = if let Some(tools) = original_body.tools.as_deref() {
ensure_request_mcp_client(&mcp_orchestrator, tools).await
} else {
None
};
Expand All @@ -1045,7 +1042,7 @@ pub async fn handle_streaming_response(ctx: RequestContext) -> Response {
let req = ctx.into_streaming_context();

// If no MCP tools, use simple passthrough
let Some((orchestrator, mcp_servers)) = mcp_result else {
let Some(mcp_servers) = mcp_servers else {
return handle_simple_streaming_passthrough(
&client,
circuit_breaker,
Expand All @@ -1060,11 +1057,8 @@ pub async fn handle_streaming_response(ctx: RequestContext) -> Response {
&client,
headers.as_ref(),
req,
&orchestrator,
McpLoopConfig {
mcp_servers,
..McpLoopConfig::default()
},
&mcp_orchestrator,
mcp_servers,
)
.await
}
Loading