mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
Merge pull request #18940 from BerriAI/litellm_fix_extra_headers
[fix] forward MCP extra headers case-insensitively
This commit is contained in:
commit
9caf685f1e
4 changed files with 93 additions and 4 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue