"""Test cases for ows-manual-adjustment connector.""" from unittest import mock import pytest from accounting.flows.reserve_payouts.connectors import ows_manual_adjustment APP_NAME = 'swf-accounting' OWS_MANUAL_ADJUSTMENT = 'ows-manual-adjustment' @pytest.mark.parametrize('status_code, expected_result', [ (400, False), (500, False), (200, True), ]) @mock.patch( 'accounting.flows.reserve_payouts.connectors.ows_manual_adjustment' '.requests') def test_health_check(mocked_request, status_code, expected_result): """Test health_check function.""" response = mock.MagicMock(status_code=status_code) mocked_request.process.return_value = response result = ows_manual_adjustment.health_check() assert result.bool == expected_result mocked_request.process.assert_called_with( application=APP_NAME, environment=None, method='GET', service_name=OWS_MANUAL_ADJUSTMENT, path='/health') @mock.patch( 'accounting.flows.reserve_payouts.connectors.ows_manual_adjustment' '.request_with_retry') @mock.patch( 'accounting.flows.reserve_payouts.connectors.ows_manual_adjustment' '.functools') @mock.patch( 'accounting.flows.reserve_payouts.connectors.ows_manual_adjustment' '.ows_manual_adjustment') @mock.patch( 'accounting.flows.reserve_payouts.connectors.ows_manual_adjustment' '.setting') @mock.patch( 'accounting.flows.reserve_payouts.connectors.ows_manual_adjustment' '.requests') def test_post_manual_adjustment( mocked_request, setting, ows_manual_adjustment_constant, functools, request_with_retry): """Test post_manual_adjustment connector function.""" app_name = 'test_app_name' env = 'test_env' service_name = 'test_ows_manual_adjustment' path = '/test-manual-adjustment' num_retries = 42 setting.APP_NAME = app_name setting.ENVIRONMENT = env setting.OWS_MANUAL_ADJUSTMENT = service_name setting.OWS_MANUAL_ADJUSTMENT_NUM_RETRIES = num_retries ows_manual_adjustment_constant.MANUAL_ADJUSTMENT_ENDPOINT = path request_dict = mock.MagicMock() response = mock.MagicMock() request_with_retry.return_value = response mock_partial_request = mock.MagicMock() functools.partial.return_value = mock_partial_request result = ows_manual_adjustment.post_manual_adjustment(request_dict) assert result is response functools.partial.assert_called_once_with( mocked_request.process, application=app_name, environment=env, method='POST', service_name=service_name, path=path, json=request_dict) request_with_retry.assert_called_once_with( mock_partial_request, num_retries) @mock.patch( 'accounting.flows.reserve_payouts.connectors.ows_manual_adjustment' '.ows_manual_adjustment') @mock.patch( 'accounting.flows.reserve_payouts.connectors.ows_manual_adjustment' '.setting') @mock.patch( 'accounting.flows.reserve_payouts.connectors.ows_manual_adjustment' '.request_with_retry') def test_post_manual_adjustment_raises( request_with_retry, setting, ows_manual_adjustment_constant): """Test post_manual_adjustment raises the exception.""" request_with_retry.side_effect = Exception() with pytest.raises(Exception): ows_manual_adjustment.post_manual_adjustment(mock.MagicMock()) @pytest.mark.parametrize('return_codes, expected_call_num', [ ([200], 1), ([1, 200], 2), ([1, 2, 3, 200], 4), ([1] * 10, 4), ]) @mock.patch( 'accounting.flows.reserve_payouts.connectors.ows_manual_adjustment' '.time') @mock.patch( 'accounting.flows.reserve_payouts.connectors.ows_manual_adjustment' '.setting') def test_request_with_retry(setting, time, return_codes, expected_call_num): """Test request_with_retry utility function.""" retry_codes = (1, 2, 3) setting.HTTP_CODES_TO_RETRY = retry_codes time_to_sleep = 42 setting.OWS_MANUAL_ADJUSTMENT_TIME_TO_SLEEP = time_to_sleep num_retries = 4 responses = [mock.MagicMock(status_code=code) for code in return_codes] mock_request = mock.MagicMock(side_effect=responses) expected_request_calls = [mock.call() for i in range(expected_call_num)] expected_time_sleep_calls = [ mock.call(time_to_sleep * i) for i, code in enumerate(return_codes[:num_retries]) if code in retry_codes] ows_manual_adjustment.request_with_retry(mock_request, num_retries) mock_request.assert_has_calls(expected_request_calls) time.sleep.assert_has_calls(expected_time_sleep_calls)