diff --git a/airflow/providers/sftp/hooks/sftp.py b/airflow/providers/sftp/hooks/sftp.py index 61b36cbeb1e66..ec1f860bf2371 100644 --- a/airflow/providers/sftp/hooks/sftp.py +++ b/airflow/providers/sftp/hooks/sftp.py @@ -320,3 +320,12 @@ def append_matching_path_callback(list_): ) return files, dirs, unknowns + + def test_connection(self) -> Tuple[bool, str]: + """Test the SFTP connection by checking if remote entity '/some/path' exists""" + try: + conn = self.get_conn() + conn.pwd + return True, "Connection successfully tested" + except Exception as e: + return False, str(e) diff --git a/tests/providers/sftp/hooks/test_sftp.py b/tests/providers/sftp/hooks/test_sftp.py index 445d50e98dfa4..35a52562a63a7 100644 --- a/tests/providers/sftp/hooks/test_sftp.py +++ b/tests/providers/sftp/hooks/test_sftp.py @@ -304,6 +304,37 @@ def test_get_tree_map(self): assert dirs == [os.path.join(TMP_PATH, TMP_DIR_FOR_TESTS, SUB_DIR)] assert unknowns == [] + @mock.patch('airflow.providers.sftp.hooks.sftp.SFTPHook.get_connection') + @mock.patch( + 'airflow.providers.sftp.hooks.sftp.SFTPHook.get_conn.pwd', side_effect=Exception('Connection Error') + ) + def test_connection_failure(self, mock_get_connection, mock_pwd): + connection = Connection( + login='login', + host='host', + ) + mock_get_connection.return_value = connection + + hook = SFTPHook() + status, msg = hook.test_connection() + assert status is False + assert msg == 'Connection Error' + + @mock.patch('airflow.providers.sftp.hooks.sftp.SFTPHook.get_connection') + @mock.patch('airflow.providers.sftp.hooks.sftp.SFTPHook.get_conn.pwd') + def test_connection_success(self, mock_get_connection, mock_pwd): + connection = Connection( + login='login', + host='host', + ) + mock_get_connection.return_value = connection + mock_pwd.return_value = '/home/some_user' + + hook = SFTPHook() + status, msg = hook.test_connection() + assert status is True + assert msg == 'Connection successfully tested' + def tearDown(self): shutil.rmtree(os.path.join(TMP_PATH, TMP_DIR_FOR_TESTS)) os.remove(os.path.join(TMP_PATH, TMP_FILE_FOR_TESTS))