diff --git a/airflow/providers/apache/hdfs/hooks/webhdfs.py b/airflow/providers/apache/hdfs/hooks/webhdfs.py index fa4e219eeddd2..04ce0d4f0cefb 100644 --- a/airflow/providers/apache/hdfs/hooks/webhdfs.py +++ b/airflow/providers/apache/hdfs/hooks/webhdfs.py @@ -17,6 +17,7 @@ # under the License. """Hook for Web HDFS""" import logging +import socket from hdfs import HdfsError, InsecureClient @@ -56,27 +57,35 @@ def __init__(self, webhdfs_conn_id='webhdfs_default', proxy_user=None): def get_conn(self): """ Establishes a connection depending on the security mode set via config or environment variable. - :return: a hdfscli InsecureClient or KerberosClient object. :rtype: hdfs.InsecureClient or hdfs.ext.kerberos.KerberosClient """ - connections = self.get_connections(self.webhdfs_conn_id) + connection = self._find_valid_server() + if connection is None: + raise AirflowWebHDFSHookException("Failed to locate the valid server.") + return connection + def _find_valid_server(self): + connections = self.get_connections(self.webhdfs_conn_id) for connection in connections: + host_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + self.log.info("Trying to connect to %s:%s", connection.host, connection.port) try: - self.log.debug('Trying namenode %s', connection.host) - client = self._get_client(connection) - client.status('/') - self.log.debug('Using namenode %s for hook', connection.host) - return client + conn_check = host_socket.connect_ex((connection.host, connection.port)) + if conn_check == 0: + self.log.info('Trying namenode %s', connection.host) + client = self._get_client(connection) + client.status('/') + self.log.info('Using namenode %s for hook', connection.host) + host_socket.close() + return client + else: + self.log.info("Could not connect to %s:%s", connection.host, connection.port) + host_socket.close() except HdfsError as hdfs_error: - self.log.debug('Read operation on namenode %s failed with error: %s', - connection.host, hdfs_error) - - hosts = [connection.host for connection in connections] - error_message = 'Read operations failed on the namenodes below:\n{hosts}'.format( - hosts='\n'.join(hosts)) - raise AirflowWebHDFSHookException(error_message) + self.log.info('Read operation on namenode %s failed with error: %s', + connection.host, hdfs_error) + return None def _get_client(self, connection): connection_str = 'http://{host}:{port}'.format(host=connection.host, port=connection.port) diff --git a/tests/providers/apache/hdfs/hooks/test_webhdfs.py b/tests/providers/apache/hdfs/hooks/test_webhdfs.py index ad034b411396c..4c70bd99dd56c 100644 --- a/tests/providers/apache/hdfs/hooks/test_webhdfs.py +++ b/tests/providers/apache/hdfs/hooks/test_webhdfs.py @@ -35,8 +35,10 @@ def setUp(self): Connection(host='host_1', port=123), Connection(host='host_2', port=321, login='user') ]) - def test_get_conn(self, mock_get_connections, mock_insecure_client): + @patch("airflow.providers.apache.hdfs.hooks.webhdfs.socket") + def test_get_conn(self, socket_mock, mock_get_connections, mock_insecure_client): mock_insecure_client.side_effect = [HdfsError('Error'), mock_insecure_client.return_value] + socket_mock.socket.return_value.connect_ex.return_value = 0 conn = self.webhdfs_hook.get_conn() mock_insecure_client.assert_has_calls([ @@ -52,10 +54,13 @@ def test_get_conn(self, mock_get_connections, mock_insecure_client): Connection(host='host_1', port=123) ]) @patch('airflow.providers.apache.hdfs.hooks.webhdfs._kerberos_security_mode', return_value=True) + @patch("airflow.providers.apache.hdfs.hooks.webhdfs.socket") def test_get_conn_kerberos_security_mode(self, + socket_mock, mock_kerberos_security_mode, mock_get_connections, mock_kerberos_client): + socket_mock.socket.return_value.connect_ex.return_value = 0 conn = self.webhdfs_hook.get_conn() connection = mock_get_connections.return_value[0] @@ -63,7 +68,7 @@ def test_get_conn_kerberos_security_mode(self, 'http://{host}:{port}'.format(host=connection.host, port=connection.port)) self.assertEqual(conn, mock_kerberos_client.return_value) - @patch('airflow.providers.apache.hdfs.hooks.webhdfs.WebHDFSHook.get_connections', return_value=[]) + @patch('airflow.providers.apache.hdfs.hooks.webhdfs.WebHDFSHook._find_valid_server', return_value=None) def test_get_conn_no_connection_found(self, mock_get_connection): with self.assertRaises(AirflowWebHDFSHookException): self.webhdfs_hook.get_conn()