From 7483a39c1d1ef21a4d6f05e40f74c262a2e404b5 Mon Sep 17 00:00:00 2001 From: Harshit28j Date: Sat, 21 Mar 2026 00:13:36 +0530 Subject: [PATCH] fix: mcp headers logging --- litellm/litellm_core_utils/litellm_logging.py | 19 ++ .../proxy/_experimental/mcp_server/server.py | 10 +- tests/mcp_tests/test_mcp_logging.py | 301 ++++++++++++++++++ 3 files changed, 329 insertions(+), 1 deletion(-) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index fea139a64b4..5799ff6bfb0 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -4763,6 +4763,25 @@ class StandardLoggingPayloadSetup: ): clean_metadata["requester_metadata"] = _potential_requester_metadata + # Populate requester_custom_headers from proxy_server_request + # headers so custom x-* headers are available in logging callbacks. + if proxy_server_request is not None: + _request_headers = proxy_server_request.get("headers", {}) + if _request_headers and isinstance(_request_headers, dict): + custom_headers = { + k: v + for k, v in _request_headers.items() + if k.lower().startswith("x-") + and v is not None + and isinstance(v, str) + and not re.match( + r"^x-mcp-.+-(authorization|x-api-key)$", + k.lower(), + ) + } + if custom_headers: + clean_metadata["requester_custom_headers"] = custom_headers + if ( EnterpriseStandardLoggingPayloadSetupVAR and proxy_server_request is not None diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index cd06de2a2df..9c91d364d59 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -378,12 +378,20 @@ if MCP_AVAILABLE: body_data["litellm_trace_id"] = chain_id body_data["litellm_session_id"] = chain_id + _scope_headers: list[tuple[bytes, bytes]] = [ + (b"content-type", b"application/json"), + ] + if raw_headers: + for k, v in raw_headers.items(): + _scope_headers.append( + (k.lower().encode("latin-1"), v.encode("latin-1")) + ) request = Request( scope={ "type": "http", "method": "POST", "path": "/mcp/tools/call", - "headers": [(b"content-type", b"application/json")], + "headers": _scope_headers, } ) if user_api_key_auth is not None: diff --git a/tests/mcp_tests/test_mcp_logging.py b/tests/mcp_tests/test_mcp_logging.py index d9ecb594b7b..3432bbb5c73 100644 --- a/tests/mcp_tests/test_mcp_logging.py +++ b/tests/mcp_tests/test_mcp_logging.py @@ -368,3 +368,304 @@ async def test_mcp_tool_call_hook(): print("logged_standard_logging_payload", logged_standard_logging_payload) assert logged_standard_logging_payload is not None, "Standard logging payload should not be None" assert logged_standard_logging_payload["response_cost"] == 1.42 + + +@pytest.mark.asyncio +async def test_mcp_custom_headers_in_logging_callback(): + """Test that custom x-* headers from MCP requests appear in requester_custom_headers.""" + litellm.logging_callback_manager._reset_all_callbacks() + mock_result = CallToolResult( + content=[TextContent(type="text", text="Test response")], + isError=False, + ) + + mock_client = AsyncMock() + mock_client.call_tool = AsyncMock(return_value=mock_result) + mock_client.list_tools = AsyncMock( + return_value=[ + MCPTool( + name="add_tools", + description="Test tool", + inputSchema={ + "type": "object", + "properties": {"test": {"type": "string"}}, + }, + ) + ] + ) + + def mock_client_constructor(*args, **kwargs): + return mock_client + + local_mcp_server_manager = MCPServerManager() + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient", + mock_client_constructor, + ): + await local_mcp_server_manager.load_servers_from_config( + mcp_servers_config={ + "test_server": { + "url": "https://example.com/mcp", + } + } + ) + + test_logger = TestMCPLogger() + litellm.callbacks = [test_logger] + + await local_mcp_server_manager._initialize_tool_name_to_mcp_server_name_mapping() + local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[ + "add_tools" + ] = "test_server" + local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[ + "test_server-add_tools" + ] = "test_server" + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + local_mcp_server_manager, + ), patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + local_mcp_server_manager, + ): + # Set auth context WITH raw_headers containing custom headers + raw_headers = { + "x-custom-header-foo": "bar", + "x-trace-id": "abc123", + "content-type": "application/json", + "user-agent": "test-agent/1.0", + } + set_auth_context( + user_api_key_auth=UserAPIKeyAuth( + api_key="test", + user_id="test_user", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="mcp-test-permissions", + mcp_servers=list( + local_mcp_server_manager.get_all_mcp_server_ids() + ), + ), + ), + mcp_servers=list( + local_mcp_server_manager.get_all_mcp_server_ids() + ), + raw_headers=raw_headers, + ) + + await mcp_server_tool_call( + name="test_server-add_tools", + arguments={"test": "test"}, + ) + + await asyncio.sleep(2) + + payload = test_logger.standard_logging_payload + assert payload is not None, "Standard logging payload should not be None" + + metadata = payload["metadata"] + custom_headers = metadata.get("requester_custom_headers") + assert custom_headers is not None, ( + "requester_custom_headers should be populated" + ) + assert custom_headers["x-custom-header-foo"] == "bar" + assert custom_headers["x-trace-id"] == "abc123" + # Non x-* headers should NOT be in requester_custom_headers + assert "content-type" not in custom_headers + assert "user-agent" not in custom_headers + + +@pytest.mark.asyncio +async def test_mcp_sensitive_headers_not_in_logging_callback(): + """Test that sensitive headers are excluded from requester_custom_headers. + + Disables the enterprise logging override so we test the open-source + filtering in isolation (the enterprise code has its own filter). + """ + litellm.logging_callback_manager._reset_all_callbacks() + mock_result = CallToolResult( + content=[TextContent(type="text", text="Test response")], + isError=False, + ) + + mock_client = AsyncMock() + mock_client.call_tool = AsyncMock(return_value=mock_result) + mock_client.list_tools = AsyncMock( + return_value=[ + MCPTool( + name="add_tools", + description="Test tool", + inputSchema={ + "type": "object", + "properties": {"test": {"type": "string"}}, + }, + ) + ] + ) + + def mock_client_constructor(*args, **kwargs): + return mock_client + + local_mcp_server_manager = MCPServerManager() + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient", + mock_client_constructor, + ): + await local_mcp_server_manager.load_servers_from_config( + mcp_servers_config={ + "test_server": { + "url": "https://example.com/mcp", + } + } + ) + + test_logger = TestMCPLogger() + litellm.callbacks = [test_logger] + + await local_mcp_server_manager._initialize_tool_name_to_mcp_server_name_mapping() + local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[ + "add_tools" + ] = "test_server" + local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[ + "test_server-add_tools" + ] = "test_server" + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + local_mcp_server_manager, + ), patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + local_mcp_server_manager, + ), patch( + "litellm.litellm_core_utils.litellm_logging.EnterpriseStandardLoggingPayloadSetupVAR", + None, + ): + # Include sensitive headers that should be stripped + raw_headers = { + "authorization": "Bearer secret-token", + "x-litellm-api-key": "sk-secret-key", + "x-mcp-github-authorization": "Bearer gh-token", + "x-mcp-zapier-x-api-key": "zapier-secret", + "x-custom-safe-header": "safe-value", + } + set_auth_context( + user_api_key_auth=UserAPIKeyAuth( + api_key="test", + user_id="test_user", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="mcp-test-permissions", + mcp_servers=list( + local_mcp_server_manager.get_all_mcp_server_ids() + ), + ), + ), + mcp_servers=list( + local_mcp_server_manager.get_all_mcp_server_ids() + ), + raw_headers=raw_headers, + ) + + await mcp_server_tool_call( + name="test_server-add_tools", + arguments={"test": "test"}, + ) + + await asyncio.sleep(2) + + payload = test_logger.standard_logging_payload + assert payload is not None, "Standard logging payload should not be None" + + metadata = payload["metadata"] + custom_headers = metadata.get("requester_custom_headers") + + # The safe custom header should be present + assert custom_headers is not None + assert custom_headers.get("x-custom-safe-header") == "safe-value" + + # Sensitive headers must NOT appear + assert "authorization" not in custom_headers + # x-litellm-api-key is stripped by clean_headers (SpecialHeaders) + assert "x-litellm-api-key" not in custom_headers + # Server-specific MCP auth headers are filtered by regex + assert "x-mcp-github-authorization" not in custom_headers + assert "x-mcp-zapier-x-api-key" not in custom_headers + + +@pytest.mark.asyncio +async def test_mcp_logging_without_raw_headers(): + """Test no regression when raw_headers is None (existing behavior).""" + litellm.logging_callback_manager._reset_all_callbacks() + mock_result = CallToolResult( + content=[TextContent(type="text", text="Test response")], + isError=False, + ) + + mock_client = AsyncMock() + mock_client.call_tool = AsyncMock(return_value=mock_result) + mock_client.list_tools = AsyncMock( + return_value=[ + MCPTool( + name="add_tools", + description="Test tool", + inputSchema={ + "type": "object", + "properties": {"test": {"type": "string"}}, + }, + ) + ] + ) + + def mock_client_constructor(*args, **kwargs): + return mock_client + + local_mcp_server_manager = MCPServerManager() + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient", + mock_client_constructor, + ): + await local_mcp_server_manager.load_servers_from_config( + mcp_servers_config={ + "test_server": { + "url": "https://example.com/mcp", + } + } + ) + + test_logger = TestMCPLogger() + litellm.callbacks = [test_logger] + + await local_mcp_server_manager._initialize_tool_name_to_mcp_server_name_mapping() + local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[ + "add_tools" + ] = "test_server" + local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[ + "test_server-add_tools" + ] = "test_server" + + with patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + local_mcp_server_manager, + ), patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + local_mcp_server_manager, + ): + # No raw_headers — existing behavior + _set_authorized_user( + local_mcp_server_manager.get_all_mcp_server_ids() + ) + + await mcp_server_tool_call( + name="test_server-add_tools", + arguments={"test": "test"}, + ) + + await asyncio.sleep(2) + + payload = test_logger.standard_logging_payload + assert payload is not None, "Standard logging payload should not be None" + # Should not crash; requester_custom_headers should be None or empty + metadata = payload["metadata"] + custom_headers = metadata.get("requester_custom_headers") + assert custom_headers is None or custom_headers == {}