fix openapi mcp extra header passthrough

This commit is contained in:
Genmin 2026-04-29 23:04:42 -07:00
parent ebd335da67
commit 67fbb33689
5 changed files with 231 additions and 34 deletions

View file

@ -2264,6 +2264,8 @@ class MCPServerManager:
server: MCPServer,
tool_name: str,
arguments: Dict[str, Any],
mcp_auth_header: Optional[str] = None,
request_extra_headers: Optional[Dict[str, str]] = None,
) -> CallToolResult:
"""
Call an OpenAPI tool handler directly.
@ -2281,6 +2283,10 @@ class MCPServerManager:
"""
from mcp.types import TextContent
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_auth_header,
_request_extra_headers,
)
from litellm.proxy._experimental.mcp_server.tool_registry import (
global_mcp_tool_registry,
)
@ -2297,9 +2303,24 @@ class MCPServerManager:
)
try:
auth_header_value: Optional[str] = None
if mcp_auth_header:
if server.auth_type == MCPAuth.api_key:
auth_header_value = f"ApiKey {mcp_auth_header}"
elif server.auth_type == MCPAuth.basic:
auth_header_value = f"Basic {mcp_auth_header}"
else:
auth_header_value = f"Bearer {mcp_auth_header}"
auth_token = _request_auth_header.set(auth_header_value)
extra_headers_token = _request_extra_headers.set(request_extra_headers)
# Call the tool handler with the arguments
# The handler is an async function that makes the HTTP request
handler_result = await tool.handler(**arguments)
try:
handler_result = await tool.handler(**arguments)
finally:
_request_extra_headers.reset(extra_headers_token)
_request_auth_header.reset(auth_token)
# Convert the handler result (string response) to CallToolResult format
result = CallToolResult(
@ -2474,6 +2495,50 @@ class MCPServerManager:
)
)
def _build_openapi_request_extra_headers(
self,
mcp_server: MCPServer,
oauth2_headers: Optional[Dict[str, str]],
raw_headers: Optional[Dict[str, str]],
hook_extra_headers: Optional[Dict[str, str]],
) -> Optional[Dict[str, str]]:
"""Build per-request headers for OpenAPI-generated MCP tool handlers."""
extra_headers = oauth2_headers.copy() if oauth2_headers else None
if mcp_server.extra_headers and raw_headers:
if extra_headers is None:
extra_headers = {}
normalized_raw_headers = {
str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)
}
for header in mcp_server.extra_headers:
if not isinstance(header, str):
continue
if (
mcp_server.has_client_credentials
and header.lower() == "authorization"
):
continue
header_value = normalized_raw_headers.get(header.lower())
if header_value is None:
continue
extra_headers[header] = header_value
if mcp_server.static_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(mcp_server.static_headers)
if hook_extra_headers:
if extra_headers is None:
extra_headers = {}
extra_headers.update(hook_extra_headers)
if extra_headers is not None and len(extra_headers) == 0:
return None
return extra_headers
async def _call_regular_mcp_tool( # noqa: PLR0915
self,
mcp_server: MCPServer,
@ -2746,17 +2811,21 @@ class MCPServerManager:
verbose_logger.debug(
"Calling OpenAPI tool %s directly via HTTP handler", name
)
if hook_result.get("extra_headers"):
verbose_logger.warning(
"pre_mcp_call hook returned extra_headers for OpenAPI-backed "
"MCP server '%s' — header injection is not supported for "
"OpenAPI servers; headers will be ignored. Use SSE/HTTP "
"transport to enable hook header injection.",
server_name,
)
request_extra_headers = self._build_openapi_request_extra_headers(
mcp_server=mcp_server,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
hook_extra_headers=hook_result.get("extra_headers"),
)
tasks.append(
asyncio.create_task(
self._call_openapi_tool_handler(mcp_server, name, arguments)
self._call_openapi_tool_handler(
mcp_server,
name,
arguments,
mcp_auth_header=mcp_auth_header,
request_extra_headers=request_extra_headers,
)
)
)
else:

View file

@ -30,6 +30,9 @@ HEADERS: Dict[str, str] = {}
_request_auth_header: contextvars.ContextVar[Optional[str]] = contextvars.ContextVar(
"_request_auth_header", default=None
)
_request_extra_headers: contextvars.ContextVar[Optional[Dict[str, str]]] = (
contextvars.ContextVar("_request_extra_headers", default=None)
)
def _sanitize_path_parameter_value(param_value: Any, param_name: str) -> str:
@ -273,6 +276,20 @@ def build_input_schema(operation: Dict[str, Any]) -> Dict[str, Any]:
}
def _get_effective_headers(headers: Dict[str, str]) -> Dict[str, str]:
effective_headers = dict(headers)
request_extra_headers = _request_extra_headers.get()
if request_extra_headers:
effective_headers.update(request_extra_headers)
override_auth = _request_auth_header.get()
if override_auth:
effective_headers["Authorization"] = override_auth
return effective_headers
def create_tool_function(
path: str,
method: str,
@ -314,10 +331,7 @@ def create_tool_function(
# The ContextVar holds the full Authorization header value, including the
# correct prefix (Bearer / ApiKey / Basic) formatted by the caller in
# server.py based on the server's configured auth_type.
effective_headers = dict(headers)
override_auth = _request_auth_header.get()
if override_auth:
effective_headers["Authorization"] = override_auth
effective_headers = _get_effective_headers(headers)
# Build URL from base_url and path
url = base_url + path

View file

@ -158,6 +158,7 @@ if MCP_AVAILABLE:
)
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_auth_header,
_request_extra_headers,
)
from litellm.proxy._experimental.mcp_server.sse_transport import SseServerTransport
from litellm.proxy._experimental.mcp_server.tool_registry import (
@ -1150,6 +1151,32 @@ if MCP_AVAILABLE:
return server_auth_header, extra_headers
def _get_request_extra_headers_for_openapi_tool(
server: Optional[MCPServer],
oauth2_headers: Optional[Dict[str, str]],
raw_headers: Optional[Dict[str, str]],
) -> Optional[Dict[str, str]]:
"""Build per-request headers for a local OpenAPI-generated MCP tool."""
extra_headers = oauth2_headers.copy() if oauth2_headers else None
if server and server.extra_headers and raw_headers:
if extra_headers is None:
extra_headers = {}
normalized_raw_headers = {
str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)
}
for header in server.extra_headers:
if not isinstance(header, str):
continue
header_value = normalized_raw_headers.get(header.lower())
if header_value is None:
continue
extra_headers[header] = header_value
return extra_headers
def _merge_gateway_initialize_instructions(
allowed_mcp_servers: List[MCPServer],
) -> Optional[str]:
@ -2154,10 +2181,17 @@ if MCP_AVAILABLE:
auth_header_value = f"Basic {mcp_auth_header}"
else:
auth_header_value = f"Bearer {mcp_auth_header}"
request_extra_headers = _get_request_extra_headers_for_openapi_tool(
server=mcp_server,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
_auth_token = _request_auth_header.set(auth_header_value)
_extra_headers_token = _request_extra_headers.set(request_extra_headers)
try:
local_content = await _handle_local_mcp_tool(name, arguments)
finally:
_request_extra_headers.reset(_extra_headers_token)
_request_auth_header.reset(_auth_token)
response = CallToolResult(content=cast(Any, local_content), isError=False)

View file

@ -6,7 +6,7 @@ Validates that:
2. pre_call_tool_check returns hook-provided extra_headers AND modified arguments
3. call_tool flows hook headers and modified arguments downstream
4. Hook-provided headers take highest priority (merge after static_headers)
5. OpenAPI-backed servers log a warning and continue (skip injection) when hook headers are present
5. OpenAPI-backed servers merge hook headers into request-scoped generated-tool headers
6. JWT claims are propagated in both standard and virtual-key fast paths
7. Backward compatibility: hooks without extra_headers continue to work
"""
@ -16,7 +16,6 @@ from typing import Any, Dict, Optional
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.proxy._types import UserAPIKeyAuth
@ -422,8 +421,8 @@ class TestCallToolFlowsHookHeaders:
assert call_kwargs.kwargs.get("arguments") == modified_args
@pytest.mark.asyncio
async def test_openapi_server_warns_and_continues_on_hook_headers(self):
"""OpenAPI-backed servers log a warning and continue when hook injects headers."""
async def test_openapi_server_forwards_hook_headers(self):
"""OpenAPI-backed servers forward hook headers through request context."""
manager = MCPServerManager()
server = MCPServer(
server_id="test-id",
@ -454,24 +453,20 @@ class TestCallToolFlowsHookHeaders:
"_call_openapi_tool_handler",
new_callable=AsyncMock,
return_value=MagicMock(),
):
import litellm.proxy._experimental.mcp_server.mcp_server_manager as mgr_mod
) as mock_call:
proxy_logging = MagicMock(spec=ProxyLogging)
with patch.object(mgr_mod, "verbose_logger") as mock_logger:
# Should NOT raise — just warn and proceed
await manager.call_tool(
server_name="openapi_server",
name="test_tool",
arguments={},
proxy_logging_obj=proxy_logging,
)
mock_logger.warning.assert_called_once()
assert (
"header injection is not supported"
in mock_logger.warning.call_args[0][0]
)
await manager.call_tool(
server_name="openapi_server",
name="test_tool",
arguments={},
proxy_logging_obj=proxy_logging,
)
mock_call.assert_called_once()
assert mock_call.call_args.kwargs["request_extra_headers"] == {
"Authorization": "Bearer jwt"
}
@pytest.mark.asyncio
async def test_openapi_server_no_error_without_hook_headers(self):
@ -773,6 +768,28 @@ class TestHookHeaderMergePriority:
assert "Authorization" not in headers
assert headers.get("X-Custom") == "from-client"
def test_openapi_request_headers_merge_oauth_raw_and_hook_headers(self):
"""OpenAPI tools receive the same runtime header sources as MCP transports."""
manager = MCPServerManager()
server = self._make_server(extra_headers_config=["X-TOKEN", "X-Trace"])
headers = manager._build_openapi_request_extra_headers(
mcp_server=server,
oauth2_headers={"Authorization": "Bearer oauth-token"},
raw_headers={
"x-token": "request-token",
"x-trace": "trace-from-request",
"x-ignored": "not-forwarded",
},
hook_extra_headers={"Authorization": "Bearer hook-token"},
)
assert headers == {
"Authorization": "Bearer hook-token",
"X-TOKEN": "request-token",
"X-Trace": "trace-from-request",
}
class TestUserAPIKeyAuthJwtClaims:
"""Tests that UserAPIKeyAuth correctly carries jwt_claims."""

View file

@ -17,6 +17,8 @@ import pytest
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_resolve_param_list,
_resolve_ref,
_request_auth_header,
_request_extra_headers,
build_input_schema,
create_tool_function,
extract_parameters,
@ -77,6 +79,67 @@ class TestCreateToolFunction:
call_args[0][0]
)
@pytest.mark.asyncio
async def test_request_extra_headers_are_forwarded(self):
"""OpenAPI tools should merge per-request header passthrough."""
operation = {}
func = create_tool_function(
path="/protected",
method="get",
operation=operation,
base_url="https://api.example.com",
headers={"X-Static": "static"},
)
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
async_client = _create_mock_client("get", "ok")
mock_client.return_value = async_client
token = _request_extra_headers.set({"X-TOKEN": "request-token"})
try:
result = await func()
finally:
_request_extra_headers.reset(token)
assert result == "ok"
call_args = async_client.get.call_args
assert call_args.kwargs["headers"] == {
"X-Static": "static",
"X-TOKEN": "request-token",
}
@pytest.mark.asyncio
async def test_auth_context_overrides_request_extra_authorization_header(self):
"""BYOK auth must keep highest precedence over generic request headers."""
operation = {}
func = create_tool_function(
path="/protected",
method="get",
operation=operation,
base_url="https://api.example.com",
)
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
async_client = _create_mock_client("get", "ok")
mock_client.return_value = async_client
extra_token = _request_extra_headers.set(
{"Authorization": "Bearer request-token", "X-TOKEN": "request-token"}
)
auth_token = _request_auth_header.set("Bearer byok-token")
try:
result = await func()
finally:
_request_auth_header.reset(auth_token)
_request_extra_headers.reset(extra_token)
assert result == "ok"
call_args = async_client.get.call_args
assert call_args.kwargs["headers"] == {
"Authorization": "Bearer byok-token",
"X-TOKEN": "request-token",
}
@pytest.mark.asyncio
async def test_leading_digit_parameter(self):
"""Test function with parameter starting with digit (e.g., 2fa-code)."""