"""Test validate_format task.""" from unittest.mock import patch import pytest from tasks.adjustment_file_upload.validate_format import\ validate_format_task @patch('tasks.adjustment_file_upload.validate_format.aws') @patch('tasks.adjustment_file_upload.validate_format.ows') @patch('tasks.adjustment_file_upload.validate_format.helpers') def test_validate_format_task( mock_helpers, mock_ows, mock_aws, mock_statement_period_adjustment_file_details, mock_adjustment_file_upload_dag_run, mock_valid_adjustment_file ): """Test validate headers and content.""" statement_period_adjustment_file_id = 1 mock_response = mock_statement_period_adjustment_file_details mock_ows.get_statement_period_adjustment_file_details.return_value = mock_response mock_aws.split_path.return_value.bucket = 'qa-abacus-adjustments' mock_aws.split_path.return_value.key = 'valid/location/test_file.csv' mock_aws.get_file.return_value = mock_valid_adjustment_file mock_helpers.get_event_from_params.return_value.target_id = \ statement_period_adjustment_file_id validate_format_task(mock_adjustment_file_upload_dag_run) mock_ows.get_statement_period_adjustment_file_details.assert_called_once_with( statement_period_adjustment_file_id) mock_aws.split_path.assert_called_once_with( 's3://qa-abacus-adjustments/valid/location/test_file.csv') mock_aws.get_file.assert_called_once_with( '437795906767', 'qa-abacus-adjustments', 'valid/location/test_file.csv') mock_helpers.get_event_from_params.assert_called_once_with( mock_adjustment_file_upload_dag_run) mock_ows.update_adjustment_file.assert_not_called() @patch('tasks.adjustment_file_upload.validate_format.aws') @patch('tasks.adjustment_file_upload.validate_format.ows') @patch('tasks.adjustment_file_upload.validate_format.helpers') def test_validate_format_task_no_valid_file( mock_helpers, mock_ows, mock_aws, mock_adjustment_file_upload_dag_run, ): """Test validate_format task raises an exception when the task fails.""" statement_period_adjustment_file_id = 1 mock_response = {} mock_ows.get_statement_period_adjustment_file_details.return_value = mock_response mock_helpers.get_event_from_params.return_value.target_id = \ statement_period_adjustment_file_id with pytest.raises(AssertionError) as missing_upload_error: validate_format_task(mock_adjustment_file_upload_dag_run) assert 'Missing adjustment file' == str(missing_upload_error.value) mock_ows.get_statement_period_adjustment_file_details.assert_called_once_with( statement_period_adjustment_file_id) assert not mock_aws.split_path.called assert not mock_ows.update_adjustment_file.called @patch('tasks.adjustment_file_upload.validate_format.aws') @patch('tasks.adjustment_file_upload.validate_format.ows') @patch('tasks.adjustment_file_upload.validate_format.helpers') def test_validate_format_task_invalid_header_count( mock_helpers, mock_ows, mock_aws, mock_statement_period_adjustment_file_details, mock_adjustment_file_upload_dag_run, mock_invalid_header_count_file ): """Test validate_format task with invalid header count.""" statement_period_adjustment_file_id = 1 mock_response = mock_statement_period_adjustment_file_details mock_ows.get_statement_period_adjustment_file_details.return_value = mock_response mock_aws.split_path.return_value.bucket = 'qa-abacus-adjustments' mock_aws.split_path.return_value.key = 'valid/location/test_file.csv' mock_aws.get_file.return_value = mock_invalid_header_count_file mock_helpers.get_event_from_params.return_value.target_id = \ statement_period_adjustment_file_id with pytest.raises(Exception) as excinfo: validate_format_task(mock_adjustment_file_upload_dag_run) mock_ows.get_statement_period_adjustment_file_details.assert_called_once_with( statement_period_adjustment_file_id) mock_aws.split_path.assert_called_once_with( 's3://qa-abacus-adjustments/valid/location/test_file.csv') mock_ows.update_adjustment_file.assert_called_once_with( statement_period_adjustment_file_id, body=dict(error_type='format_error')) assert str(excinfo.value) == 'Missing columns' @patch('tasks.adjustment_file_upload.validate_format.aws') @patch('tasks.adjustment_file_upload.validate_format.ows') @patch('tasks.adjustment_file_upload.validate_format.helpers') def test_validate_format_task_invalid_headers( mock_helpers, mock_ows, mock_aws, mock_statement_period_adjustment_file_details, mock_adjustment_file_upload_dag_run, mock_invalid_adjustment_file ): """Test validate_format task with invalid data.""" statement_period_adjustment_file_id = 1 mock_response = mock_statement_period_adjustment_file_details mock_ows.get_statement_period_adjustment_file_details.return_value = mock_response mock_aws.split_path.return_value.bucket = 'qa-abacus-adjustments' mock_aws.split_path.return_value.key = 'valid/location/test_file.csv' mock_aws.get_file.return_value = mock_invalid_adjustment_file mock_helpers.get_event_from_params.return_value.target_id = \ statement_period_adjustment_file_id with pytest.raises(Exception) as excinfo: validate_format_task(mock_adjustment_file_upload_dag_run) mock_ows.get_statement_period_adjustment_file_details.assert_called_once_with( statement_period_adjustment_file_id) mock_aws.split_path.assert_called_once_with( 's3://qa-abacus-adjustments/valid/location/test_file.csv') mock_ows.update_adjustment_file.assert_called_once_with( statement_period_adjustment_file_id, body=dict(error_type='format_error')) assert str(excinfo.value) == 'Mismatched or missing column names' @patch('tasks.adjustment_file_upload.validate_format.aws') @patch('tasks.adjustment_file_upload.validate_format.ows') @patch('tasks.adjustment_file_upload.validate_format.helpers') def test_validate_format_task_no_adjustment_records( mock_helpers, mock_ows, mock_aws, mock_statement_period_adjustment_file_details, mock_adjustment_file_upload_dag_run, mock_empty_adjustment_file ): """Test validate_format task for an empty adjustment file.""" statement_period_adjustment_file_id = 1 mock_response = mock_statement_period_adjustment_file_details mock_ows.get_statement_period_adjustment_file_details.return_value = mock_response mock_aws.split_path.return_value.bucket = 'qa-abacus-adjustments' mock_aws.split_path.return_value.key = 'valid/location/test_file.csv' mock_aws.get_file.return_value = mock_empty_adjustment_file mock_helpers.get_event_from_params.return_value.target_id = \ statement_period_adjustment_file_id with pytest.raises(Exception) as excinfo: validate_format_task(mock_adjustment_file_upload_dag_run) mock_ows.get_statement_period_adjustment_file_details.assert_called_once_with( statement_period_adjustment_file_id) mock_aws.split_path.assert_called_once_with( 's3://qa-abacus-adjustments/valid/location/test_file.csv') mock_ows.update_adjustment_file.assert_called_once_with( statement_period_adjustment_file_id, body=dict(error_type='format_error')) assert str(excinfo.value) == 'There are no adjustment records in the file' @patch('tasks.adjustment_file_upload.validate_format.aws') @patch('tasks.adjustment_file_upload.validate_format.ows') @patch('tasks.adjustment_file_upload.validate_format.helpers') def test_validate_format_task_too_many_rows( mock_helpers, mock_ows, mock_aws, mock_statement_period_adjustment_file_details, mock_adjustment_file_upload_dag_run, mock_too_many_rows_adjustment_file ): """Test validate_format task for a file with too many rows.""" statement_period_adjustment_file_id = 1 mock_response = mock_statement_period_adjustment_file_details mock_ows.get_statement_period_adjustment_file_details.return_value = mock_response mock_aws.split_path.return_value.bucket = 'qa-abacus-adjustments' mock_aws.split_path.return_value.key = 'valid/location/test_file.csv' mock_aws.get_file.return_value = mock_too_many_rows_adjustment_file mock_helpers.get_event_from_params.return_value.target_id = \ statement_period_adjustment_file_id with pytest.raises(Exception) as excinfo: validate_format_task(mock_adjustment_file_upload_dag_run) mock_ows.get_statement_period_adjustment_file_details.assert_called_once_with( statement_period_adjustment_file_id) mock_aws.split_path.assert_called_once_with( 's3://qa-abacus-adjustments/valid/location/test_file.csv') mock_ows.update_adjustment_file.assert_called_once_with( statement_period_adjustment_file_id, body=dict(error_type='row_count_error')) assert str(excinfo.value) == 'The file exceeds the maximum number of rows'