From 1bcabdf4833c5f5343920f20a710346450edd2cd Mon Sep 17 00:00:00 2001 From: Raj Nagulapalle Date: Tue, 5 May 2026 16:04:11 -0700 Subject: [PATCH] fix(mcp): cover OpenAPI hook extra headers Co-authored-by: Cursor --- .../mcp_server/mcp_server_manager.py | 4 ++- .../mcp_server/test_mcp_hook_extra_headers.py | 34 ++++++++----------- .../test_openapi_to_mcp_generator.py | 33 ++++++++++++++++++ 3 files changed, 51 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 89cc005d292..b6f4ebf889f 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2351,7 +2351,9 @@ class MCPServerManager: extra_headers = {} normalized_raw_headers = { - str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) + 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): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index 84c556b8ddc..fe2d701d45e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -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 forward hook-provided headers to generated HTTP tools 6. JWT claims are propagated in both standard and virtual-key fast paths 7. Backward compatibility: hooks without extra_headers continue to work """ @@ -422,8 +422,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-injected headers to the HTTP handler.""" manager = MCPServerManager() server = MCPServer( server_id="test-id", @@ -454,24 +454,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_openapi_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_openapi_call.assert_called_once() + assert mock_openapi_call.call_args.kwargs[ + "hook_extra_headers" + ] == {"Authorization": "Bearer jwt"} @pytest.mark.asyncio async def test_openapi_server_no_error_without_hook_headers(self): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py index efe100a11dc..0713f5d14de 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py @@ -15,6 +15,7 @@ from unittest.mock import AsyncMock, patch import pytest from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + _request_extra_headers, _resolve_param_list, _resolve_ref, build_input_schema, @@ -369,6 +370,38 @@ class TestCreateToolFunction: # Should have no exec() calls assert len(exec_calls) == 0, "create_tool_function should not use exec()" + @pytest.mark.asyncio + async def test_request_extra_headers_contextvar_merges_into_request_headers(self): + """Hook-injected headers set by the caller are sent to the upstream API.""" + operation = {} + func = create_tool_function( + path="/protected", + method="get", + operation=operation, + base_url="https://api.example.com", + headers={"X-Static": "static"}, + ) + + token = _request_extra_headers.set( + {"Authorization": "Bearer signed-jwt", "X-Trace": "trace-id"} + ) + try: + with patch(GET_ASYNC_CLIENT_TARGET) as mock_client: + async_client = _create_mock_client("get", "ok") + mock_client.return_value = async_client + + result = await func() + + assert result == "ok" + call_kwargs = async_client.get.call_args.kwargs + assert call_kwargs["headers"] == { + "X-Static": "static", + "Authorization": "Bearer signed-jwt", + "X-Trace": "trace-id", + } + finally: + _request_extra_headers.reset(token) + class TestBuildInputSchema: """Test that build_input_schema preserves original parameter names."""