import unittest.mock import pytest from db_schema import apps from db_schema.factories.apps import UnitOfWorkFactory from apps_etl_databricks import exceptions from apps_etl_databricks.objects.job import DBXJobJar, DBXJobSpark, PersistentDBXJob from apps_etl_databricks.objects.run import RunState @pytest.mark.parametrize( 'mode,klass', ( ('jar_params', DBXJobJar), ('spark_submit_params', DBXJobSpark), ) ) @unittest.mock.patch('time.sleep', return_value=None) @unittest.mock.patch('requests.get') @unittest.mock.patch('requests.post') def test_dbx_job_jar(mocked_post, mocked_get, _, mode, klass): run_id, job_id = 10, 12 base_url, token = 'base_url', 'token' command, contexts = ['--dsp', 'apple'], ['sme:US'] statuses = unittest.mock.Mock() statuses.licensors_and_contexts.return_value = contexts statuses.is_empty.return_value = False statuses.is_backfill_context.return_value = False statuses.is_backfill_uow.return_value = False licensors = unittest.mock.Mock() licensors.get_names.return_value = [] licensors.is_empty.return_value = True mocked_post.return_value.json.return_value = {'run_id': run_id} mocked_get.return_value.json.return_value = {'state': RunState(life_cycle_state='IN_PROGRESS')} job = klass(unittest.mock.Mock(), base_url, token, job_id) actual_run_id = job.run(licensors, statuses, command[:]) mocked_post.assert_called_once_with( url=f'https://{base_url}/api/2.0/jobs/run-now', headers={'Authorization': f'Bearer {token}'}, json={ 'job_id': job_id, mode: [*command, '--context', *contexts] }, ) assert actual_run_id == run_id @pytest.mark.parametrize( 'mode,klass', ( ('jar_params', DBXJobJar), ('spark_submit_params', DBXJobSpark), ) ) @unittest.mock.patch('time.sleep', return_value=None) @unittest.mock.patch('requests.get') @unittest.mock.patch('requests.post') def test_dbx_job_jar_backfill(mocked_post, mocked_get, _, mode, klass): run_id, job_id = 10, 12 base_url, token = 'base_url', 'token' command, contexts = ['--dsp', 'apple'], ['*'] statuses = unittest.mock.Mock() statuses.licensors_and_contexts.return_value = contexts statuses.is_empty.return_value = False statuses.is_backfill_context.return_value = True statuses.is_backfill_uow.return_value = True licensors = unittest.mock.Mock() licensors.get_names.return_value = ['sme'] mocked_post.return_value.json.return_value = {'run_id': run_id} mocked_get.return_value.json.return_value = {'state': RunState(life_cycle_state='IN_PROGRESS')} job = klass(unittest.mock.Mock(), base_url, token, job_id) actual_run_id = job.run(licensors, statuses, command[:]) mocked_post.assert_called_once_with( url=f'https://{base_url}/api/2.0/jobs/run-now', headers={'Authorization': f'Bearer {token}'}, json={ 'job_id': job_id, mode: [*command, '--context', 'sme:*'] }, ) assert actual_run_id == run_id @pytest.mark.integration @unittest.mock.patch('time.sleep', return_value=None) @unittest.mock.patch('requests.get') def test_persistent_job(mocked_get, _, session): uow = UnitOfWorkFactory() dbx_job_mocked = unittest.mock.Mock() dbx_job_mocked.run.return_value = run_id = 999 dbx_job_mocked.job.return_value = job_id = 111 statuses = unittest.mock.Mock() statuses.licensors_and_contexts.return_value = (('sme', 'US'), ) statuses.is_empty.return_values = False licensors = unittest.mock.Mock() licensors.get_names.return_value = [] licensors.is_empty.return_value = True mocked_get.return_value.json.return_value = {'state': RunState(life_cycle_state='IN_PROGRESS')} persistent_job = PersistentDBXJob( logger=unittest.mock.Mock(), job=dbx_job_mocked, session=session, uow_id=uow.unit_of_work_id, sf_name='test-sf-name', dbx_job_max_lifetime=3600, ) persistent_job.run(licensors, statuses, command=[]) dbx_execution = session.query(apps.DatabricksExecution).one_or_none() assert dbx_execution assert dbx_execution.unit_of_work_id == uow.unit_of_work_id assert dbx_execution.dbx_run_id == run_id assert dbx_execution.dbx_job_id == job_id @pytest.mark.integration @unittest.mock.patch('requests.get') def test_persistent_job_was_not_skipped(mocked_get, session): uow = UnitOfWorkFactory() dbx_job_mocked = unittest.mock.Mock() mocked_get.return_value.json.return_value = {'state': RunState(life_cycle_state='IN_PROGRESS')} persistent_job = PersistentDBXJob( logger=unittest.mock.Mock(), job=dbx_job_mocked, session=session, uow_id=uow.unit_of_work_id, sf_name='test-sf-name', dbx_job_max_lifetime=3600, ) is_skipped = persistent_job.check_job_was_skipped(run_id=1) assert not is_skipped @pytest.mark.integration @unittest.mock.patch('time.sleep', return_value=None) @unittest.mock.patch('requests.get') def test_check_job_was_skipped_raised_exception(mocked_get, _, session): uow = UnitOfWorkFactory() dbx_job_mocked = unittest.mock.Mock() dbx_job_mocked.run.return_value = 7 dbx_job_mocked.job.return_value = 14 statuses = unittest.mock.Mock() statuses.licensors_and_contexts.return_value = (('sme', '2382::80032614'), ) statuses.is_empty.return_values = False licensors = unittest.mock.Mock() licensors.get_names.return_value = [] licensors.is_empty.return_value = True mocked_get.return_value.json.return_value = {'state': RunState(life_cycle_state='SKIPPED')} persistent_job = PersistentDBXJob( logger=unittest.mock.Mock(), job=dbx_job_mocked, session=session, uow_id=uow.unit_of_work_id, sf_name='test-sf-name', dbx_job_max_lifetime=3600, ) with pytest.raises(exceptions.DBXExecutionWasSkipped): persistent_job.run(licensors, statuses, command=[]) # exception has been raised, but data dbx_execution = session.query(apps.DatabricksExecution).one_or_none() assert dbx_execution