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
66 changes: 20 additions & 46 deletions model_gateway/src/routers/header_utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -80,52 +80,6 @@ fn should_forward_header_no_alloc(name: &str) -> bool {
|| name.eq_ignore_ascii_case("host"))
}

/// Apply headers to a reqwest request builder, filtering out headers that shouldn't be forwarded
/// or that will be set automatically by reqwest
pub fn apply_request_headers(
headers: &HeaderMap,
mut request_builder: reqwest::RequestBuilder,
skip_content_headers: bool,
) -> reqwest::RequestBuilder {
// Always forward Authorization header first if present
if let Some(auth) = headers
.get("authorization")
.or_else(|| headers.get("Authorization"))
{
request_builder = request_builder.header("Authorization", auth.clone());
}

// Forward other headers, filtering out problematic ones
// Use eq_ignore_ascii_case to avoid to_lowercase() allocation per header
for (key, value) in headers {
let key_str = key.as_str();

// Skip headers that:
// - Are set automatically by reqwest (content-type, content-length for POST/PUT)
// - We already handled (authorization)
// - Are hop-by-hop headers (connection, transfer-encoding)
// - Should not be forwarded (host)
let should_skip = key_str.eq_ignore_ascii_case("authorization") // Already handled above
|| key_str.eq_ignore_ascii_case("host")
|| key_str.eq_ignore_ascii_case("connection")
|| key_str.eq_ignore_ascii_case("transfer-encoding")
|| key_str.eq_ignore_ascii_case("keep-alive")
|| key_str.eq_ignore_ascii_case("te")
|| key_str.eq_ignore_ascii_case("trailers")
|| key_str.eq_ignore_ascii_case("accept-encoding")
|| key_str.eq_ignore_ascii_case("upgrade")
|| (skip_content_headers
&& (key_str.eq_ignore_ascii_case("content-type")
|| key_str.eq_ignore_ascii_case("content-length")));

if !should_skip {
request_builder = request_builder.header(key.clone(), value.clone());
}
}

request_builder
}

/// API provider types for provider-specific header handling
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ApiProvider {
Expand Down Expand Up @@ -375,4 +329,24 @@ mod tests {
assert!(!should_forward_request_header("x-custom-header"));
assert!(!should_forward_request_header("x-api-key"));
}

#[test]
fn test_extract_auth_header_falls_back_with_non_auth_headers_present() {
let mut headers = HeaderMap::new();
headers.insert("openai-project", "project-123".parse().unwrap());

let auth = extract_auth_header(Some(&headers), Some(&"worker-secret".to_string()));

assert_eq!(auth.unwrap(), "Bearer worker-secret");
}

#[test]
fn test_provider_extract_auth_header_prefers_anthropic_key() {
let mut headers = HeaderMap::new();
headers.insert("x-api-key", "anthropic-key".parse().unwrap());

let auth = ApiProvider::Anthropic.extract_auth_header(Some(&headers), None);

assert_eq!(auth.unwrap(), "anthropic-key");
}
}
11 changes: 5 additions & 6 deletions model_gateway/src/routers/openai/mcp/tool_loop.rs
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@ use tracing::{debug, info, warn};
use super::tool_handler::FunctionCallInProgress;
use crate::{
observability::metrics::{metrics_labels, Metrics},
routers::{error, header_utils::apply_request_headers, mcp_utils::DEFAULT_MAX_ITERATIONS},
routers::{error, header_utils::ApiProvider, mcp_utils::DEFAULT_MAX_ITERATIONS},
};

/// State for tracking multi-turn tool calling loop
Expand Down Expand Up @@ -498,6 +498,7 @@ pub(crate) async fn execute_tool_loop(
client: &reqwest::Client,
url: &str,
headers: Option<&HeaderMap>,
worker_api_key: Option<&String>,
initial_payload: Value,
original_body: &ResponsesRequest,
session: &McpToolSession<'_>,
Expand All @@ -512,14 +513,12 @@ pub(crate) async fn execute_tool_loop(
"Starting tool loop: max_tool_calls={:?}, max_iterations={}",
max_tool_calls, DEFAULT_MAX_ITERATIONS
);
let provider = ApiProvider::from_url(url);
let auth_header = provider.extract_auth_header(headers, worker_api_key);
Comment thread
zhaowenzi marked this conversation as resolved.

loop {
let request_builder = client.post(url).json(&current_payload);
let request_builder = if let Some(headers) = headers {
apply_request_headers(headers, request_builder, true)
} else {
request_builder
};
let request_builder = provider.apply_headers(request_builder, auth_header.as_ref());

let response = request_builder
.send()
Expand Down
8 changes: 5 additions & 3 deletions model_gateway/src/routers/openai/responses/non_streaming.rs
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@ use tracing::warn;
use super::utils::{patch_response_with_request_metadata, restore_original_tools};
use crate::routers::{
error,
header_utils::{apply_provider_headers, extract_auth_header},
header_utils::ApiProvider,
mcp_utils::ensure_request_mcp_client,
openai::{
context::{PayloadState, RequestContext},
Expand Down Expand Up @@ -78,6 +78,7 @@ pub async fn handle_non_streaming_response(mut ctx: RequestContext) -> Response
ctx.components.client(),
&url,
ctx.headers(),
worker.api_key(),
payload,
original_body,
&session,
Expand All @@ -95,8 +96,9 @@ pub async fn handle_non_streaming_response(mut ctx: RequestContext) -> Response
}
} else {
let mut request_builder = ctx.components.client().post(&url).json(&payload);
let auth_header = extract_auth_header(ctx.headers(), worker.api_key());
request_builder = apply_provider_headers(request_builder, &url, auth_header.as_ref());
let provider = ApiProvider::from_url(&url);
let auth_header = provider.extract_auth_header(ctx.headers(), worker.api_key());
request_builder = provider.apply_headers(request_builder, auth_header.as_ref());

let response = match request_builder.send().await {
Ok(r) => r,
Expand Down
18 changes: 10 additions & 8 deletions model_gateway/src/routers/openai/responses/streaming.rs
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@ use crate::{
observability::metrics::Metrics,
routers::{
error,
header_utils::{apply_request_headers, preserve_response_headers},
header_utils::{preserve_response_headers, ApiProvider},
mcp_utils::DEFAULT_MAX_ITERATIONS,
openai::{
context::{RequestContext, StreamingEventContext, StreamingRequest},
Expand Down Expand Up @@ -519,10 +519,9 @@ pub(super) async fn handle_simple_streaming_passthrough(
req: StreamingRequest,
) -> Response {
let mut request_builder = client.post(&req.url).json(&req.payload);

if let Some(headers) = headers {
request_builder = apply_request_headers(headers, request_builder, true);
}
let provider = ApiProvider::from_url(&req.url);
let auth_header = provider.extract_auth_header(headers, worker.api_key());
Comment thread
zhaowenzi marked this conversation as resolved.
request_builder = provider.apply_headers(request_builder, auth_header.as_ref());

request_builder = request_builder.header("Accept", "text/event-stream");

Expand Down Expand Up @@ -669,6 +668,7 @@ pub(super) async fn handle_simple_streaming_passthrough(
/// Handle streaming WITH MCP tool call interception and execution
pub(super) fn handle_streaming_with_tool_interception(
client: &reqwest::Client,
worker_api_key: Option<String>,
headers: Option<&HeaderMap>,
req: StreamingRequest,
orchestrator: &Arc<McpOrchestrator>,
Expand Down Expand Up @@ -719,13 +719,14 @@ pub(super) fn handle_streaming_with_tool_interception(
previous_response_id: previous_response_id.as_deref(),
session: Some(&session),
};
let provider = ApiProvider::from_url(&url_clone);
let auth_header =
provider.extract_auth_header(headers_opt.as_ref(), worker_api_key.as_ref());

loop {
// Make streaming request
let mut request_builder = client_clone.post(&url_clone).json(&current_payload);
if let Some(ref h) = headers_opt {
request_builder = apply_request_headers(h, request_builder, true);
}
request_builder = provider.apply_headers(request_builder, auth_header.as_ref());
request_builder = request_builder.header("Accept", "text/event-stream");

let response = match request_builder.send().await {
Expand Down Expand Up @@ -1060,6 +1061,7 @@ pub async fn handle_streaming_response(ctx: RequestContext) -> Response {

handle_streaming_with_tool_interception(
&client,
worker.api_key().cloned(),
headers.as_ref(),
req,
&mcp_orchestrator,
Expand Down
Loading