import unittest from unittest import mock from db_schema.schemas.exploration import DisassembleContentStatus from exp_dbx_status_checker_lambda import Input from exp_dbx_status_checker_lambda.dbx_service import (DatabricksJob, DatabricksService, RunState) from exp_dbx_status_checker_lambda.exceptions import ( DBXStateError, DisassembleCSRecordNotFound) class DatabricksServiceTestCase(unittest.TestCase): def setUp(self) -> None: self.payload = Input(disassemble_content_status_id=123, ) self.pg_repo = mock.Mock() self.dbx_service = DatabricksService(mock.Mock(), mock.Mock(), self.payload, self.pg_repo) def test_update_db_ok(self): self.pg_repo.get_disassemble_cs_by_id.side_effect = lambda *args: DisassembleContentStatus( dbx_run_id=111) result = self.dbx_service._get_run_id(123) self.assertEqual(result, 111) def test_update_db_correct_raises_exc(self): self.pg_repo.get_disassemble_cs_by_id.side_effect = lambda *args: DisassembleContentStatus( ) self.assertRaises(DisassembleCSRecordNotFound, self.dbx_service._get_run_id, 123) class DatabricksJobTestCase(unittest.TestCase): def setUp(self) -> None: self.dbx_service = DatabricksJob(dbx_url="dbx_url", dbx_token="dbx_token") @unittest.mock.patch('requests.get') def test_get_run_state(self, mocked_get): mocked_get.return_value.json.return_value = { 'state': { "life_cycle_state": 'IN_PROGRESS' } } run_id = 10 _state = self.dbx_service.get_run_state(run_id) self.assertEqual(_state, RunState(life_cycle_state='IN_PROGRESS')) mocked_get.assert_called_once_with( url='https://dbx_url/api/2.0/jobs/runs/get', headers={'Authorization': 'Bearer dbx_token'}, params={'run_id': run_id}, ) class RunStateTestCase(unittest.TestCase): def test_get_state_running(self): run_state = RunState(life_cycle_state='RUNNING') state = run_state.get_state() self.assertEqual(state, 'RUNNING') def test_get_state_success(self): run_state = RunState( life_cycle_state='TERMINATED', result_state='SUCCESS', ) state = run_state.get_state() self.assertEqual(state, 'SUCCESS') def test_get_state_correct_raises_exc(self): run_state = RunState(state_message='test', ) self.assertRaises(DBXStateError, run_state.get_state)