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/openai/context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@ pub struct SharedComponents {
}

pub struct ResponsesComponents {
pub shared: SharedComponents,
pub shared: Arc<SharedComponents>,
pub mcp_orchestrator: Arc<McpOrchestrator>,
pub response_storage: Arc<dyn ResponseStorage>,
pub conversation_storage: Arc<dyn ConversationStorage>,
Expand Down
41 changes: 13 additions & 28 deletions model_gateway/src/routers/openai/mcp/tool_loop.rs
Original file line number Diff line number Diff line change
Expand Up @@ -480,18 +480,12 @@ pub(crate) fn inject_mcp_metadata_streaming(
item.get("type").and_then(|t| t.as_str()) != Some(ItemType::MCP_LIST_TOOLS)
});

for binding in mcp_servers.iter().rev() {
let list_tools_item =
session.build_mcp_list_tools_json(&binding.label, &binding.server_key);
output_array.insert(0, list_tools_item);
}

// Use stored transformed items (no reconstruction needed)
let mut insert_pos = mcp_servers.len();
for item in &state.mcp_call_items {
output_array.insert(insert_pos, item.clone());
insert_pos += 1;
let mut prefix = Vec::with_capacity(mcp_servers.len() + state.mcp_call_items.len());
for binding in mcp_servers {
prefix.push(session.build_mcp_list_tools_json(&binding.label, &binding.server_key));
}
prefix.extend(state.mcp_call_items.iter().cloned());
output_array.splice(0..0, prefix);
} else if let Some(obj) = response.as_object_mut() {
let mut output_items = Vec::new();
for binding in mcp_servers {
Expand Down Expand Up @@ -728,24 +722,15 @@ fn build_incomplete_response(

// Add mcp_list_tools and executed mcp_call items at the beginning
if state.total_calls > 0 || !incomplete_items.is_empty() {
for binding in mcp_servers.iter().rev() {
let list_tools_item =
session.build_mcp_list_tools_json(&binding.label, &binding.server_key);
output_array.insert(0, list_tools_item);
}

// Insert stored transformed items for executed calls (no reconstruction needed)
let mut insert_pos = mcp_servers.len();
for item in &state.mcp_call_items {
output_array.insert(insert_pos, item.clone());
insert_pos += 1;
}

// Add incomplete mcp_call items (never executed, so no stored item)
for item in incomplete_items {
output_array.insert(insert_pos, item);
insert_pos += 1;
let mut prefix = Vec::with_capacity(
mcp_servers.len() + state.mcp_call_items.len() + incomplete_items.len(),
);
for binding in mcp_servers {
prefix.push(session.build_mcp_list_tools_json(&binding.label, &binding.server_key));
}
prefix.extend(state.mcp_call_items.iter().cloned());
prefix.extend(incomplete_items);
output_array.splice(0..0, prefix);
}
}

Expand Down
29 changes: 13 additions & 16 deletions model_gateway/src/routers/openai/responses/history.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,9 @@
//! input before forwarding to the upstream provider.

use axum::response::Response;
use openai_protocol::responses::{
ResponseContentPart, ResponseInput, ResponseInputOutputItem, ResponsesRequest,
use openai_protocol::{
event_types::ItemType,
responses::{ResponseContentPart, ResponseInput, ResponseInputOutputItem, ResponsesRequest},
};
use serde_json::Value;
use smg_data_connector::{ConversationId, ListParams, ResponseId, SortOrder};
Expand All @@ -25,22 +26,18 @@ const MAX_CONVERSATION_HISTORY_ITEMS: usize = 100;
/// Returns `Ok(original_previous_response_id)` on success, or `Err(response)` on validation failure.
pub(crate) async fn load_input_history(
components: &ResponsesComponents,
body: &ResponsesRequest,
conversation: Option<&str>,
request_body: &mut ResponsesRequest,
model: &str,
) -> Result<Option<String>, Response> {
let original_previous_response_id = request_body
let previous_response_id = request_body
.previous_response_id
.clone()
.take()
.filter(|id| !id.is_empty());

// Load items from previous response chain if specified
let mut chain_items: Option<Vec<ResponseInputOutputItem>> = None;
if let Some(prev_id_str) = request_body
.previous_response_id
.take()
.filter(|id| !id.is_empty())
{
if let Some(prev_id_str) = &previous_response_id {
let prev_id = ResponseId::from(prev_id_str.as_str());
match components
.response_storage
Expand Down Expand Up @@ -74,8 +71,8 @@ pub(crate) async fn load_input_history(
}

// Load conversation history if specified
if let Some(conv_id_str) = body.conversation.clone().filter(|id| !id.is_empty()) {
let conv_id = ConversationId::from(conv_id_str.as_str());
if let Some(conv_id_str) = conversation {
let conv_id = ConversationId::from(conv_id_str);

if let Ok(None) = components
.conversation_storage
Expand Down Expand Up @@ -129,7 +126,7 @@ pub(crate) async fn load_input_history(
}
}
}
"function_call" => {
ItemType::FUNCTION_CALL => {
match serde_json::from_value::<ResponseInputOutputItem>(item.content) {
Ok(func_call) => items.push(func_call),
Err(e) => {
Expand Down Expand Up @@ -164,7 +161,7 @@ pub(crate) async fn load_input_history(
}
}

append_current_input(&mut items, &request_body.input, &conv_id.0);
append_current_input(&mut items, &request_body.input, conv_id_str);
request_body.input = ResponseInput::Items(items);
}
Err(e) => {
Expand All @@ -178,12 +175,12 @@ pub(crate) async fn load_input_history(
// (enforced by the caller in route_responses), so this branch and the
// conversation branch above never both modify request_body.input.
if let Some(mut items) = chain_items {
let id_suffix = original_previous_response_id.as_deref().unwrap_or("new");
let id_suffix = previous_response_id.as_deref().unwrap_or("new");
append_current_input(&mut items, &request_body.input, id_suffix);
request_body.input = ResponseInput::Items(items);
}

Ok(original_previous_response_id)
Ok(previous_response_id)
}

/// Deserialize ResponseInputOutputItems from a JSON array value
Expand Down
26 changes: 14 additions & 12 deletions model_gateway/src/routers/openai/responses/route.rs
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,6 @@ pub(in crate::routers::openai) struct ResponsesRouterContext<'a> {
pub worker_registry: &'a WorkerRegistry,
pub provider_registry: &'a ProviderRegistry,
pub responses_components: &'a Arc<ResponsesComponents>,
pub client: &'a reqwest::Client,
}

/// Route a responses API request to the appropriate upstream worker.
Expand All @@ -57,14 +56,17 @@ pub(in crate::routers::openai) async fn route_responses(
bool_to_static_str(streaming),
);

let worker = match WorkerSelector::new(deps.worker_registry, deps.client)
.select_worker(&SelectWorkerRequest {
model_id: model,
headers,
provider: Some(ProviderType::OpenAI),
..Default::default()
})
.await
let worker = match WorkerSelector::new(
deps.worker_registry,
&deps.responses_components.shared.client,
)
.select_worker(&SelectWorkerRequest {
model_id: model,
headers,
provider: Some(ProviderType::OpenAI),
..Default::default()
})
.await
{
Ok(w) => w,
Err(response) => {
Expand All @@ -82,12 +84,12 @@ pub(in crate::routers::openai) async fn route_responses(

// Validate mutual exclusivity of conversation and previous_response_id
// Treat empty strings as unset to match other metadata paths
let has_conversation = body.conversation.as_ref().is_some_and(|s| !s.is_empty());
let conversation = body.conversation.as_ref().filter(|s| !s.is_empty());
let has_previous_response = body
.previous_response_id
.as_ref()
.is_some_and(|s| !s.is_empty());
if has_conversation && has_previous_response {
if conversation.is_some() && has_previous_response {
Metrics::record_router_error(
metrics_labels::ROUTER_OPENAI,
metrics_labels::BACKEND_EXTERNAL,
Expand All @@ -110,7 +112,7 @@ pub(in crate::routers::openai) async fn route_responses(

let original_previous_response_id = match super::history::load_input_history(
deps.responses_components,
body,
conversation.map(String::as_str),
&mut request_body,
model,
)
Expand Down
5 changes: 1 addition & 4 deletions model_gateway/src/routers/openai/router.rs
Original file line number Diff line number Diff line change
Expand Up @@ -90,9 +90,7 @@ impl OpenAIRouter {
});

let responses_components = Arc::new(ResponsesComponents {
shared: SharedComponents {
client: ctx.client.clone(),
},
shared: Arc::clone(&shared_components),
mcp_orchestrator: mcp_orchestrator.clone(),
response_storage: ctx.response_storage.clone(),
conversation_storage: ctx.conversation_storage.clone(),
Expand Down Expand Up @@ -166,7 +164,6 @@ impl crate::routers::RouterTrait for OpenAIRouter {
worker_registry: &self.worker_registry,
provider_registry: &self.provider_registry,
responses_components: &self.responses_components,
client: &self.responses_components.shared.client,
};
responses_route::route_responses(&deps, headers, body, model_id).await
}
Expand Down
Loading