"""Unit tests for the `adjustment_file_generate` helper functions.""" from unittest.mock import MagicMock, patch import pytest from lib.constants import ( DAG_ADJUSTMENT_FILE_GENERATE_FLOWTHROUGH_EVENT_NAME, DAG_ADJUSTMENT_FILE_GENERATE_TARGET_TYPE, ) from tasks.adjustment_file_generate import helpers @patch('tasks.adjustment_file_generate.helpers.event') def test_get_event_from_params(mock_event, mock_adjustment_file_generate_dag_run): """Test getting the Abacus event from the DAG run config.""" mock_abacus_event = MagicMock() mock_event.get_abacus_event.return_value = mock_abacus_event mock_event.validate_event_for_handler.return_value = True res = helpers.get_event_from_params(mock_adjustment_file_generate_dag_run) assert res mock_event.get_abacus_event.assert_called_once_with( mock_adjustment_file_generate_dag_run ) mock_event.validate_event_for_handler.assert_called_once_with( mock_abacus_event, target_type=DAG_ADJUSTMENT_FILE_GENERATE_TARGET_TYPE, event_names=[DAG_ADJUSTMENT_FILE_GENERATE_FLOWTHROUGH_EVENT_NAME], ) @patch('tasks.adjustment_file_generate.helpers.ows') def test_get_abacus_state(mock_ows, mock_adjustment_file_generate_states): """Test getting an Abacus state.""" action_name = 'upload_file' statement_period_adjustment_file_id = 1 mock_ows.get_abacus_states.return_value = ( mock_adjustment_file_generate_states ) res = helpers.get_abacus_state( action_name, statement_period_adjustment_file_id ) assert res.get('action_name') == action_name mock_ows.get_abacus_states.assert_called_once_with( DAG_ADJUSTMENT_FILE_GENERATE_TARGET_TYPE, statement_period_adjustment_file_id ) # Test when the state is not found action_name = 'unknown' mock_ows.reset_mock() mock_ows.get_abacus_states.return_value = ( mock_adjustment_file_generate_states ) with pytest.raises(ValueError): helpers.get_abacus_state( action_name, statement_period_adjustment_file_id ) mock_ows.get_abacus_states.assert_called_once_with( DAG_ADJUSTMENT_FILE_GENERATE_TARGET_TYPE, statement_period_adjustment_file_id )