mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
fix openapi mcp extra header passthrough
This commit is contained in:
parent
ebd335da67
commit
67fbb33689
5 changed files with 231 additions and 34 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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)."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue