"""Unit tests for adjustment_file_upload helpers.""" from unittest.mock import patch import pytest from lib import constants from tasks.adjustment_file_upload import helpers @patch('tasks.adjustment_file_upload.helpers.event') def test_get_event_from_params(mock_event, mock_adjustment_file_upload_dag_run): """Test getting abacus_event from dag run config.""" mock_event.get_abacus_event.return_value = 'an event' mock_event.validate_event_for_handler.return_value = True res = helpers.get_event_from_params(mock_adjustment_file_upload_dag_run) assert res mock_event.get_abacus_event.assert_called_once_with( mock_adjustment_file_upload_dag_run ) mock_event.validate_event_for_handler.assert_called_once_with( 'an event', event_name=constants.DAG_ADJUSTMENT_FILE_UPLOAD_EVENT_NAME, target_type=constants.DAG_ADJUSTMENT_FILE_UPLOAD_TARGET_TYPE ) @patch('tasks.adjustment_file_upload.helpers.ows') def test_get_abacus_state_success(mock_ows, mock_adjustment_file_upload_states): """Test getting a sales_file's 'upload_file' abacus_state.""" action_name = constants.STATEMENT_PERIOD_ADJUSTMENT_FILE_ACTIONS.UPLOAD_FILE mock_ows.get_abacus_states.return_value = mock_adjustment_file_upload_states res = helpers.get_abacus_state(action_name, 123) assert res.get('action_name') == action_name mock_ows.get_abacus_states.assert_called_once_with(helpers.PARENT_TABLE_NAME, 123) @patch('tasks.adjustment_file_upload.helpers.ows') def test_get_abacus_states_failure(mock_ows, mock_adjustment_file_upload_states): """Test that an error is raised when abacus_state record does not exist.""" action_name = 'not_an_action' mock_ows.get_abacus_states.return_value = [ mock_adjustment_file_upload_states[0] ] with pytest.raises(ValueError): helpers.get_abacus_state(action_name, 123) mock_ows.get_abacus_states.assert_called_once_with(helpers.PARENT_TABLE_NAME, 123)