mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(mcp): forward OpenAPI extra headers on call_tool
This commit is contained in:
parent
7a4a68f022
commit
60c3a418f8
2 changed files with 106 additions and 1 deletions
|
|
@ -241,6 +241,34 @@ def _without_authorization(
|
|||
return filtered or None
|
||||
|
||||
|
||||
def _openapi_forwarded_extra_headers(
|
||||
mcp_server: MCPServer,
|
||||
raw_headers: Optional[dict[str, str]],
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
) -> Optional[dict[str, str]]:
|
||||
if not mcp_server.extra_headers or not raw_headers:
|
||||
return None
|
||||
|
||||
normalized_raw_headers = {
|
||||
str(header_name).lower(): header_value
|
||||
for header_name, header_value in raw_headers.items()
|
||||
if isinstance(header_name, str)
|
||||
}
|
||||
strip_caller_authorization = _should_strip_caller_authorization(
|
||||
mcp_server=mcp_server,
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
forwarded_headers = {
|
||||
header_name: normalized_raw_headers[header_name.lower()]
|
||||
for header_name in mcp_server.extra_headers
|
||||
if isinstance(header_name, str)
|
||||
and not (strip_caller_authorization and header_name.lower() == "authorization")
|
||||
and header_name.lower() in normalized_raw_headers
|
||||
}
|
||||
return forwarded_headers or None
|
||||
|
||||
|
||||
def _extract_upstream_auth_failure(
|
||||
exc: BaseException,
|
||||
) -> Optional[Tuple[int, Optional[str]]]:
|
||||
|
|
@ -3523,7 +3551,24 @@ class MCPServerManager:
|
|||
"transport to enable hook header injection.",
|
||||
server_name,
|
||||
)
|
||||
tasks.append(asyncio.create_task(self._call_openapi_tool_handler(mcp_server, name, arguments)))
|
||||
forwarded_headers = _openapi_forwarded_extra_headers(
|
||||
mcp_server=mcp_server,
|
||||
raw_headers=raw_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
|
||||
async def _call_openapi_via_handler() -> CallToolResult:
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_extra_headers,
|
||||
)
|
||||
|
||||
extra_headers_token = _request_extra_headers.set(forwarded_headers)
|
||||
try:
|
||||
return await self._call_openapi_tool_handler(mcp_server, name, arguments)
|
||||
finally:
|
||||
_request_extra_headers.reset(extra_headers_token)
|
||||
|
||||
tasks.append(asyncio.create_task(_call_openapi_via_handler()))
|
||||
else:
|
||||
return await self._call_regular_mcp_tool(
|
||||
mcp_server=mcp_server,
|
||||
|
|
|
|||
|
|
@ -515,6 +515,66 @@ class TestCallToolFlowsHookHeaders:
|
|||
proxy_logging_obj=proxy_logging,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openapi_server_forwards_allowlisted_client_headers(self):
|
||||
from mcp.types import TextContent
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
||||
_request_extra_headers,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="test-id",
|
||||
name="openapi_server",
|
||||
server_name="openapi_server",
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
spec_path="/path/to/spec.yaml",
|
||||
extra_headers=["Authorization", "X-Tenant-ID"],
|
||||
)
|
||||
manager.registry[server.server_id] = server
|
||||
manager.tool_name_to_mcp_server_name_mapping["test_tool"] = server.name
|
||||
manager.tool_name_to_mcp_server_name_mapping["openapi_server-test_tool"] = (
|
||||
server.name
|
||||
)
|
||||
|
||||
async def capture_headers() -> Optional[Dict[str, str]]:
|
||||
return _request_extra_headers.get()
|
||||
|
||||
registered_name = "openapi_server-test_tool"
|
||||
global_mcp_tool_registry.register_tool(
|
||||
name=registered_name,
|
||||
description="test",
|
||||
input_schema={},
|
||||
handler=capture_headers,
|
||||
)
|
||||
|
||||
try:
|
||||
forwarded_result = await manager.call_tool(
|
||||
server_name="openapi_server",
|
||||
name="test_tool",
|
||||
arguments={},
|
||||
raw_headers={
|
||||
"x-litellm-api-key": "sk-proxy",
|
||||
"authorization": "Bearer user-token",
|
||||
"x-tenant-id": "tenant-001",
|
||||
"x-unlisted": "not-forwarded",
|
||||
},
|
||||
)
|
||||
finally:
|
||||
global_mcp_tool_registry.tools.pop(registered_name)
|
||||
|
||||
assert forwarded_result.isError is False
|
||||
assert isinstance(forwarded_result.content[0], TextContent)
|
||||
assert forwarded_result.content[0].text == (
|
||||
"{'Authorization': 'Bearer user-token', 'X-Tenant-ID': 'tenant-001'}"
|
||||
)
|
||||
|
||||
|
||||
class TestHookHeaderMergePriority:
|
||||
"""Tests that hook-provided headers have highest priority in _call_regular_mcp_tool."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue