Merge pull request #18940 from BerriAI/litellm_fix_extra_headers

[fix] forward MCP extra headers case-insensitively
This commit is contained in:
YutaSaito 2026-01-13 06:03:19 +09:00 committed by GitHub
commit 9caf685f1e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 93 additions and 4 deletions

View file

@ -1825,9 +1825,17 @@ class MCPServerManager:
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 isinstance(header, str) and header in raw_headers:
extra_headers[header] = raw_headers[header]
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
if mcp_server.static_headers:
if extra_headers is None:

View file

@ -715,9 +715,18 @@ if MCP_AVAILABLE:
if 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 header in raw_headers:
extra_headers[header] = raw_headers[header]
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
if server_auth_header is None:
server_auth_header = mcp_auth_header

View file

@ -79,6 +79,33 @@ async def test_mcp_server_tool_call_body_contains_request_data():
assert body["arguments"] == tool_arguments
def test_prepare_mcp_server_headers_case_insensitive_extra_headers():
try:
from litellm.proxy._experimental.mcp_server.server import (
_prepare_mcp_server_headers,
)
except ImportError:
pytest.skip("MCP server not available")
server = MCPServer(
server_id="server-case",
name="server",
transport=MCPTransport.http,
extra_headers=["Authorization"],
)
server_auth_header, extra_headers = _prepare_mcp_server_headers(
server=server,
mcp_server_auth_headers=None,
mcp_auth_header=None,
oauth2_headers=None,
raw_headers={"authorization": "Bearer token"},
)
assert server_auth_header is None
assert extra_headers == {"Authorization": "Bearer token"}
@pytest.mark.asyncio
async def test_get_prompts_from_mcp_servers_success():
try:

View file

@ -12,6 +12,7 @@ sys.path.insert(0, "../../../../../")
import httpx
from mcp import ReadResourceResult, Resource
from mcp.types import (
CallToolResult,
GetPromptResult,
Prompt,
ResourceTemplate,
@ -286,6 +287,50 @@ class TestMCPServerManager:
assert len(result) == 1
assert result[0].name == "github_tool_1"
@pytest.mark.asyncio
async def test_call_regular_mcp_tool_case_insensitive_extra_headers(self):
"""_call_regular_mcp_tool should forward headers regardless of original casing."""
manager = MCPServerManager()
server = MCPServer(
server_id="server-case-call",
name="case-call-server",
url="https://example.com",
transport=MCPTransport.http,
auth_type=MCPAuth.authorization,
extra_headers=["Authorization"],
)
mock_client = AsyncMock()
mock_client.call_tool = AsyncMock(
return_value=CallToolResult(content=[], isError=False)
)
captured_extra_headers = None
def capture_create_mcp_client(
server, mcp_auth_header, extra_headers, stdio_env
): # pragma: no cover - helper
nonlocal captured_extra_headers
captured_extra_headers = extra_headers
return mock_client
manager._create_mcp_client = MagicMock(side_effect=capture_create_mcp_client)
result = await manager._call_regular_mcp_tool(
mcp_server=server,
original_tool_name="tool",
arguments={},
tasks=[],
mcp_auth_header=None,
mcp_server_auth_headers=None,
oauth2_headers=None,
raw_headers={"authorization": "Bearer token"},
proxy_logging_obj=None,
)
assert captured_extra_headers == {"Authorization": "Bearer token"}
assert isinstance(result, CallToolResult)
@pytest.mark.asyncio
async def test_get_prompts_from_server_success(self):
"""Ensure prompts are fetched and prefixed when requested."""