"""Tests for dataloader handlers.""" from unittest.mock import MagicMock, patch import pytest from abacus_common_logic.test_utils.helpers import get_message from flask import testing as flask_testing from owsresponse import response from royalties.constants.error import ERROR_INVALID_STATEMENT_PERIOD_IDS from royalties.tests.utils.factories import StatementPeriodFactory @patch('royalties.blueprints.dataloader.logic') def test_get_statement_periods_by_ids(mock_logic, fixture_client): """Get statement periods by a list of IDs.""" statement_periods = StatementPeriodFactory.create_batch(4) statement_period_ids = [sp.statement_period_id for sp in statement_periods] mock_logic.get_statement_periods_by_ids.return_value = response.Response( status=200, message='OK' ) res = fixture_client.post( '/dataloader/statement-periods/', json=statement_period_ids ) assert res.status_code == 200 mock_logic.get_statement_periods_by_ids.assert_called_once_with( statement_period_ids ) @pytest.mark.parametrize( ['standalone_check_result', 'pdp_check_result', 'expected_status_code'], [ pytest.param(True, None, 200, id='standalone check passes'), pytest.param(False, True, 200, id='pdp check passes'), pytest.param(False, False, 403, id='pdp check fails'), ], ) @patch('royalties.blueprints.dataloader.authorization') @patch('royalties.blueprints.dataloader.flask_request') def test_get_statement_periods_by_ids_authorization( mock_flask_request: MagicMock, mock_authorization: MagicMock, standalone_check_result: bool, pdp_check_result: bool | None, expected_status_code: int, fixture_client: flask_testing.FlaskClient, ) -> None: """Test authorization for getting statement periods by IDs.""" mock_flask_request.verify_rules_access_standalone.return_value = ( standalone_check_result ) mock_authorization.pdp_authorize_resource.return_value = pdp_check_result res = fixture_client.post( '/dataloader/statement-periods/', json=[1, 2], ) assert res.status_code == expected_status_code mock_flask_request.verify_rules_access_standalone.assert_called_once() if not standalone_check_result: mock_authorization.pdp_authorize_resource.assert_called_once_with( resource_id=0, resource_type='statement_period', ) @patch('royalties.blueprints.dataloader.logic') def test_get_statement_periods_by_ids_error(mock_logic, fixture_client): """Get statement periods by a list of IDs when the data is incorrect.""" res = fixture_client.post('/dataloader/statement-periods/', json=['test']) assert res.status_code == 400 assert get_message(res) == ERROR_INVALID_STATEMENT_PERIOD_IDS mock_logic.get_statement_periods_by_ids.assert_not_called()