"""Tests for AbacusContractInfo class module.""" import re from datetime import date from decimal import Decimal from unittest.mock import MagicMock, patch import httpx import pytest from owsclient.test import OwsClientMock from sync_contract.abacus_contract_info import LEGACY_SCHEDULE_REGEX, AbacusContractInfo from sync_contract.ows_royalties import SERVICE @patch('sync_contract.abacus_contract_info.get_contract_states') def test_get_legacy_sync_state_success(mock_get_contract_states, mock_contract_states): """Test successfully getting contract states and filtering for legacy_sync state.""" abacus_contract_id = 123 mock_get_contract_states.return_value = mock_contract_states abacus_contract_info = AbacusContractInfo(abacus_contract_id) res = abacus_contract_info.get_legacy_sync_state() assert res == mock_contract_states[0] mock_get_contract_states.assert_called_once_with(abacus_contract_id) @patch('sync_contract.abacus_contract_info.get_contract_states') def test_get_legacy_sync_state_error(mock_get_contract_states): """Test that an error is raised if a legacy_sync state does not exist.""" abacus_contract_id = 123 mock_contract_states = [{'action_name': 'nope'}] mock_get_contract_states.return_value = mock_contract_states abacus_contract_info = AbacusContractInfo(abacus_contract_id) with pytest.raises(ValueError) as e: abacus_contract_info.get_legacy_sync_state() assert 'abacus_state not found' in str(e.value) mock_get_contract_states.assert_called_once_with(abacus_contract_id) def test_abacus_contract_info_properties( mock_contract_details, mock_contract_exclusions, mock_contract_lifecycle, mock_contract_terms, mock_contract_term_conditions, mock_account_contracts, mock_account_payment_terms, mock_contract_states, ows_client_mock: OwsClientMock, ): """Test for AbacusContractInfo instance's props.""" abacus_contract_id = 1 abacus_contract_term_id = 1 account_id = mock_contract_details['account_id'] AbacusContractInfo.get_legacy_sync_state = MagicMock( return_value=mock_contract_states[0] ) abacus_contract_info = AbacusContractInfo(abacus_contract_id) ows_client_mock.get(SERVICE, f'/contract/{abacus_contract_id}/exclusions/').mock( return_value=httpx.Response(200, json=mock_contract_exclusions) ) ows_client_mock.get( SERVICE, f'/contract/{abacus_contract_id}/contract-lifecycle/' ).mock(return_value=httpx.Response(200, json=mock_contract_lifecycle)) ows_client_mock.get(SERVICE, f'/contract/{abacus_contract_id}/').mock( return_value=httpx.Response(200, json=mock_contract_details) ) ows_client_mock.get( SERVICE, f'/contracts/{abacus_contract_id}/contract-terms/' ).mock(return_value=httpx.Response(200, json=mock_contract_terms)) ows_client_mock.get( SERVICE, f'/contract-term/{abacus_contract_term_id}/conditions/' ).mock(return_value=httpx.Response(200, json=mock_contract_term_conditions)) ows_client_mock.get( 'ows-abacus-account', f'/account/{account_id}/account-payment-term/' ).mock(return_value=httpx.Response(200, json=mock_account_payment_terms)) ows_client_mock.get(SERVICE, '/contracts/').mock( return_value=httpx.Response(200, json=mock_account_contracts) ) abacus_contract_info.load_data_from_ows() assert abacus_contract_info.legacy_sync_state == mock_contract_states[0] assert abacus_contract_info.account_id == mock_contract_details['account_id'] assert abacus_contract_info.contract_type == mock_contract_details['contract_type'] assert ( abacus_contract_info.country_exclusions == mock_contract_exclusions['exclusions']['countries'] ) assert ( abacus_contract_info.currency_code == mock_account_payment_terms['currency_code'] ) assert ( abacus_contract_info.oa_contract_id == mock_contract_details['oa_contract_id'] ) assert abacus_contract_info.pay_after == '30' assert abacus_contract_info.payment_interval == 'month' assert abacus_contract_info.term_end == mock_contract_details['term_end'] assert abacus_contract_info.term_rate == Decimal('0.6') assert abacus_contract_info.term_start != mock_contract_details['term_start'] assert ( abacus_contract_info.term_start == mock_contract_lifecycle['lifecycle_term_start'] ) AbacusContractInfo.get_legacy_sync_state.assert_called_once() def test_abacus_contract_info_properties_without_contract_lifecycle( mock_contract_details, mock_contract_exclusions, mock_contract_lifecycle, mock_contract_terms, mock_contract_term_conditions, mock_account_contracts, mock_account_payment_terms, mock_contract_states, ows_client_mock: OwsClientMock, ): """Test for AbacusContractInfo term_start when there is no contract_lifecycle.""" abacus_contract_id = 1 abacus_contract_term_id = 1 account_id = mock_contract_details['account_id'] AbacusContractInfo.get_legacy_sync_state = MagicMock( return_value=mock_contract_states[0] ) abacus_contract_info = AbacusContractInfo(abacus_contract_id) ows_client_mock.get(SERVICE, f'/contract/{abacus_contract_id}/exclusions/').mock( return_value=httpx.Response(200, json=mock_contract_exclusions) ) ows_client_mock.get( SERVICE, f'/contract/{abacus_contract_id}/contract-lifecycle/' ).mock(return_value=httpx.Response(200, json={})) ows_client_mock.get(SERVICE, f'/contract/{abacus_contract_id}/').mock( return_value=httpx.Response(200, json=mock_contract_details) ) ows_client_mock.get( SERVICE, f'/contracts/{abacus_contract_id}/contract-terms/' ).mock(return_value=httpx.Response(200, json=mock_contract_terms)) ows_client_mock.get( SERVICE, f'/contract-term/{abacus_contract_term_id}/conditions/' ).mock(return_value=httpx.Response(200, json=mock_contract_term_conditions)) ows_client_mock.get( 'ows-abacus-account', f'/account/{account_id}/account-payment-term/' ).mock(return_value=httpx.Response(200, json=mock_account_payment_terms)) ows_client_mock.get(SERVICE, '/contracts/').mock( return_value=httpx.Response(200, json=mock_account_contracts) ) abacus_contract_info.load_data_from_ows() assert abacus_contract_info.term_start == str(date.today()) assert ( abacus_contract_info.term_start != mock_contract_lifecycle['lifecycle_term_start'] ) def test_abacus_contract_info_property_term_rate_without_base_term( mock_contract_details, mock_contract_exclusions, mock_contract_terms_without_base_term, mock_contract_term_conditions, mock_account_contracts, mock_account_payment_terms, mock_contract_states, mock_contract_lifecycle, ows_client_mock: OwsClientMock, ): """Test for AbacusContractInfo instance's props.""" abacus_contract_id = 1 abacus_contract_term_id = 1 account_id = mock_contract_details['account_id'] AbacusContractInfo.get_legacy_sync_state = MagicMock( return_value=mock_contract_states[0] ) abacus_contract_info = AbacusContractInfo(abacus_contract_id) ows_client_mock.get(SERVICE, f'/contract/{abacus_contract_id}/exclusions/').mock( return_value=httpx.Response(200, json=mock_contract_exclusions) ) ows_client_mock.get(SERVICE, f'/contract/{abacus_contract_id}/').mock( return_value=httpx.Response(200, json=mock_contract_details) ) ows_client_mock.get( SERVICE, f'/contracts/{abacus_contract_id}/contract-terms/' ).mock(return_value=httpx.Response(200, json=mock_contract_terms_without_base_term)) ows_client_mock.get( SERVICE, f'/contract-term/{abacus_contract_term_id}/conditions/' ).mock(return_value=httpx.Response(200, json=mock_contract_term_conditions)) ows_client_mock.get( 'ows-abacus-account', f'/account/{account_id}/account-payment-term/' ).mock(return_value=httpx.Response(200, json=mock_account_payment_terms)) ows_client_mock.get(SERVICE, '/contracts/').mock( return_value=httpx.Response(200, json=mock_account_contracts) ) ows_client_mock.get( SERVICE, f'/contract/{abacus_contract_id}/contract-lifecycle/' ).mock(return_value=httpx.Response(200, json=mock_contract_lifecycle)) abacus_contract_info.load_data_from_ows() assert abacus_contract_info.term_rate == Decimal('0.8') @patch('sync_contract.abacus_contract_info.get_contract_details') @patch('sync_contract.abacus_contract_info.get_contract_exclusions') @patch('sync_contract.abacus_contract_info.get_contract_lifecycle') @patch('sync_contract.abacus_contract_info.get_contract_terms') @patch('sync_contract.abacus_contract_info.get_contract_term_conditions') @patch('sync_contract.abacus_contract_info.get_account_payment_terms') def test_load_data_from_ows_failure( mock_get_account_payment_terms, mock_get_contract_term_conditions, mock_get_contract_terms, mock_get_contract_lifecycle, mock_get_contract_exclusions, mock_get_contract_details, ): """Test load data failure.""" abacus_contract_id = 123 AbacusContractInfo.get_legacy_sync_state.side_effect = ValueError('nope') abacus_contract_info = AbacusContractInfo(abacus_contract_id) with pytest.raises(ValueError): abacus_contract_info.load_data_from_ows() AbacusContractInfo.get_legacy_sync_state.assert_called() mock_get_contract_details.assert_not_called() mock_get_contract_exclusions.assert_not_called() mock_get_contract_lifecycle.assert_not_called() mock_get_contract_terms.assert_not_called() mock_get_contract_term_conditions.assert_not_called() mock_get_account_payment_terms.assert_not_called() def test_legacy_schedule_regex(): """Tests for LEGACY_SCHEDULE_REGEX constant.""" payment_schedule = '30_days_after_quarter_end' result = re.search(LEGACY_SCHEDULE_REGEX, payment_schedule) assert result.group('pay_after') == '30' assert result.group('payment_interval') == 'quarter' payment_schedule = '45_days_after_month_end' result = re.search(LEGACY_SCHEDULE_REGEX, payment_schedule) assert result.group('pay_after') == '45' assert result.group('payment_interval') == 'month' payment_schedule = 'some_bad_string' result = re.search(LEGACY_SCHEDULE_REGEX, payment_schedule) assert result is None def test_update_country_codes_for_serbia(): """Test _update_country_codes method. When Serbia country code is included in the contract_exclusions list. """ abacus_contract_id = 123 contract_exclusions = ['RUS', 'SRB'] abacus_contract_info = AbacusContractInfo(abacus_contract_id) res = abacus_contract_info._update_country_codes(contract_exclusions) assert res == ['RUS', 'SCG'] def test_update_country_codes_for_montenegro(): """Test _update_country_codes method. When Montenegro country code is included in the contract_exclusions list. """ abacus_contract_id = 123 contract_exclusions = ['RUS', 'MNE'] abacus_contract_info = AbacusContractInfo(abacus_contract_id) res = abacus_contract_info._update_country_codes(contract_exclusions) assert res == ['RUS', 'SCG'] def test_update_country_codes_for_both_serbia_and_montenegro(): """Test _update_country_codes method. When both Serbia and Montenegro are included in the contract_exclusions list. """ abacus_contract_id = 123 contract_exclusions = ['RUS', 'SRB', 'MNE'] abacus_contract_info = AbacusContractInfo(abacus_contract_id) res = abacus_contract_info._update_country_codes(contract_exclusions) assert res == ['RUS', 'SCG'] contract_exclusions = ['SRB', 'MNE'] abacus_contract_info = AbacusContractInfo(abacus_contract_id) res = abacus_contract_info._update_country_codes(contract_exclusions) assert res == ['SCG'] def test_update_country_codes(): """Test _update_country_codes method. When Serbia/Montenegro is not included in the contract_exclusions list. """ abacus_contract_id = 123 contract_exclusions = ['RUS', 'ASM', 'AUS', 'IMN'] abacus_contract_info = AbacusContractInfo(abacus_contract_id) res = abacus_contract_info._update_country_codes(contract_exclusions) assert res == contract_exclusions