import json from functools import partial from unittest import mock import pytest from dapd_db_schema.schemas.uow_meta import DBXJobStatusEnum from requests.models import Response from dapd_helpers.dbx_client import DatabricksJob @pytest.mark.parametrize( 'content, is_fail', (({ 'run_id': 16894952, 'number_in_job': 16894952 }, False), ({ 'number_in_job': 16894952 }, True)) ) @mock.patch('requests.request') def test_start_job(requests_post, content, is_fail): logger = mock.Mock() res = Response() res.status_code = 200 res._content = bytes(json.dumps(content), encoding='utf8') # pylint: disable=W0212 requests_post.return_value = res dbx_client = DatabricksJob( dbx_url='http://dbx.com', dbx_token='token123', logger=logger, request_attempts=1, sleep_timeout=1 ) start_job = partial( dbx_client.start_job, job_id=12345, commands={ 'key1': 'value1', 'key2': 'value2' } ) if is_fail: with pytest.raises(KeyError): start_job() else: run_id = start_job() assert run_id == content['run_id'] @pytest.mark.parametrize( 'content, expected_status, is_fail', (({ 'state': { 'life_cycle_state': 'PENDING', 'state_message': 'Waiting for cluster', 'user_cancelled_or_timedout': False } }, DBXJobStatusEnum.IN_PROGRESS.value, False), ({ 'state': { 'life_cycle_state': 'TERMINATED', 'result_state': 'FAILED', 'state_message': '', 'user_cancelled_or_timedout': False } }, DBXJobStatusEnum.FAILED.value, False), ({ 'state': { 'life_cycle_state': 'PENDING', 'result_state': 'SUCCESS', 'state_message': '', 'user_cancelled_or_timedout': False } }, DBXJobStatusEnum.COMPLETE.value, False), ({ 'state': { 'state_message': '', 'user_cancelled_or_timedout': False } }, None, True)) ) @mock.patch('requests.request') def test_get_run_state(requests_post, content, expected_status, is_fail): logger = mock.Mock() res = Response() res.status_code = 200 res._content = bytes(json.dumps(content), encoding='utf8') # pylint: disable=W0212 requests_post.return_value = res dbx_client = DatabricksJob( dbx_url='http://dbx.com', dbx_token='token123', logger=logger, request_attempts=1, sleep_timeout=1 ) get_run_state = partial( dbx_client.get_run_state, run_id=12345, ) if is_fail: with pytest.raises(KeyError): get_run_state() else: run_state = get_run_state() assert run_state == expected_status