"""Contract Exclusion logic tests.""" from unittest.mock import patch from abacus_contract.constants.error import ( ERROR_CONTRACT_EXCLUSION_DOES_NOT_EXIST, ERROR_OA_CONTRACT_EXISTS, ERROR_UNKNOWN_COUNTRY, ) from abacus_contract.logic import contract_exclusion as logic from abacus_contract.tests.utils.factories import ( ContractExclusionFactory, ContractFactory, LegacyContractFactory, ) @patch('abacus_contract.logic.contract_exclusion.models') def test_create_contract_exclusions(mock_models): """Test create contract exclusions logic.""" contract = ContractFactory.create() mock_models.Contract.get_by_id_or_error.return_value = contract mock_models.ContractExclusion.create.return_value = None params = {'countries': ['AFG'], 'stores': []} result = logic.create_contract_exclusions(contract.contract_id, params) assert result.status == 201 mock_models.ContractExclusion.create.assert_called_once_with( contract_id=contract.contract_id, exclusions=params, ) @patch('abacus_contract.logic.contract_exclusion.models') def test_create_contract_exclusions_error(mock_models): """Test creating exclusions for invalid country_code.""" contract = ContractFactory.create() mock_models.Contract.get_by_id_or_error.return_value = contract params = {'countries': ['test'], 'stores': []} res = logic.create_contract_exclusions(contract.contract_id, params) assert res.status == 400 assert res.errors['message'] == ERROR_UNKNOWN_COUNTRY.format(code='test') mock_models.ContractExclusion.build.assert_not_called() mock_models.ContractExclusion.commit_changes.assert_not_called() @patch('abacus_contract.logic.contract_exclusion.models') def test_create_contract_exclusions_error_oa_contract_exists(mock_models): """Test creating exclusions if oa contract exists.""" contract = ContractFactory.create() LegacyContractFactory.create(contract=contract) mock_models.Contract.get_by_id_or_error.return_value = contract params = {'countries': ['USA', 'UKR']} res = logic.create_contract_exclusions(contract.contract_id, params) assert res.status == 400 assert res.errors['message'] == ERROR_OA_CONTRACT_EXISTS params_2 = {'stores': ['111', '115']} res_2 = logic.create_contract_exclusions(contract.contract_id, params_2) assert res_2.status == 400 assert res_2.errors['message'] == ERROR_OA_CONTRACT_EXISTS params_empty = {'countries': [], 'stores': []} res_3 = logic.create_contract_exclusions(contract.contract_id, params_empty) assert res_3.status == 400 assert res_3.errors['message'] == ERROR_OA_CONTRACT_EXISTS mock_models.ContractExclusion.build.assert_not_called() mock_models.ContractExclusion.commit_changes.assert_not_called() @patch('abacus_contract.logic.contract_exclusion.models') def test_get_exclusions_by_contract(mock_models): """Test to get contract exclusions by contract_id.""" contract = ContractFactory.create() exclusion = ContractExclusionFactory.create(contract=contract) mock_models.Contract.get_by_id_or_error.return_value = contract res = logic.get_exclusions_by_contract(contract.contract_id) assert res.status == 200 assert res.message['contract_exclusion_id'] == exclusion.contract_exclusion_id def test_get_exclusion_records_by_contract_ids() -> None: """Serialize exclusions for authorized contract ids as flat records. Each record is a flat dict carrying contract_id (not the dataload wrapping). """ contract = ContractFactory.create() other_contract = ContractFactory.create() contract_without_exclusion = ContractFactory.create() exclusion = ContractExclusionFactory.create(contract=contract) other_exclusion = ContractExclusionFactory.create(contract=other_contract) res = logic.get_exclusion_records_by_contract_ids( [ contract.contract_id, other_contract.contract_id, contract_without_exclusion.contract_id, ] ) assert len(res) == 2 contract_ids = {record['contract_id'] for record in res} # The id with no exclusion produces no record; the shaper nulls it downstream. assert contract_ids == {contract.contract_id, other_contract.contract_id} assert contract_without_exclusion.contract_id not in contract_ids exclusion_ids = {record['contract_exclusion_id'] for record in res} assert exclusion_ids == { exclusion.contract_exclusion_id, other_exclusion.contract_exclusion_id, } @patch('abacus_contract.logic.contract_exclusion.models') def test_get_exclusion_records_by_contract_ids_empty_list_skips_query( mock_models, ) -> None: """No authorized ids: returns an empty list without querying. The dataloader helper only ever passes authorized ids, so an id that never resolved to an account must never reach this query. """ res = logic.get_exclusion_records_by_contract_ids([]) assert res == [] mock_models.ContractExclusion.get_by_contract_ids.assert_not_called() @patch('abacus_contract.logic.contract_exclusion.models') def test_update_contract_exclusions(mock_models): """Test to update existing contract exclusions.""" contract = ContractFactory.create() exclusion = ContractExclusionFactory.create(contract=contract) params = {'countries': ['USA', 'UKR'], 'stores': ['12']} mock_models.Contract.get_by_id_or_error.return_value = contract res = logic.update_contract_exclusions(contract.contract_id, params) assert res.status == 200 assert res.message['contract_id'] == contract.contract_id assert res.message['contract_exclusion_id'] == exclusion.contract_exclusion_id assert res.message['exclusions'] == params @patch('abacus_contract.logic.contract_exclusion.models') def test_update_contract_exclusions_no_record(mock_models): """Test to update contract exclusion that does not exist.""" contract = ContractFactory.create() params = {'countries': ['USA', 'UKR'], 'stores': ['12']} mock_models.Contract.get_by_id_or_error.return_value = contract res = logic.update_contract_exclusions(contract.contract_id, params) assert res.status == 404 assert res.errors['message'] == ERROR_CONTRACT_EXCLUSION_DOES_NOT_EXIST.format( contract_id=contract.contract_id ) mock_models.ContractExclusion.commit_changes.assert_not_called() @patch('abacus_contract.logic.contract_exclusion.models') def test_update_contract_exclusions_country_code_error(mock_models): """Test updating exclusions for invalid country_code.""" contract = ContractFactory.create() ContractExclusionFactory.create(contract=contract) mock_models.Contract.get_by_id_or_error.return_value = contract params = {'countries': ['test'], 'stores': []} res = logic.update_contract_exclusions(contract.contract_id, params) assert res.status == 400 assert res.errors['message'] == ERROR_UNKNOWN_COUNTRY.format(code='test') mock_models.ContractExclusion.commit_changes.assert_not_called()