"""Tests for vendor contract logic.""" from unittest.mock import patch import pytest from marshmallow.exceptions import ValidationError from abacus_legacy_sync.constants import constants, error from abacus_legacy_sync.logic import vendor_contract as logic from abacus_legacy_sync.schemas.vendor_contract import VendorContractDetailSchema from tests.utils.factories import ( CountryFactory, CurrenciesFactory, CustomerMasterMasterFactory, VendorContractFactory, ) @patch('abacus_legacy_sync.logic.vendor_contract.models') @patch('abacus_legacy_sync.logic.vendor_contract._validate_store_exclusion') @patch('abacus_legacy_sync.models.ows.carveouts_python.save_account_carveouts') def test_create_vendor_contract_success( mock_save_carveouts, mock_validate_store_exclusion, mock_models ): """Test success response of create_vendor_contract method.""" currency = CurrenciesFactory.create() country = CountryFactory.create() vendor_contract = VendorContractFactory.create() mock_models.Currencies.get_by_iso_code.return_value = currency mock_models.Country.get_by_iso_codes.return_value = [country] mock_models.VendorContract.create.return_value = vendor_contract mock_validate_store_exclusion.return_value = ['1', '2'] post_data = { 'vendor_id': 77, 'cont_start': '2022-01-01', 'cont_end': '2122-12-31', 'contract_type': constants.VENDOR_CONTRACT_CONTRACT_TYPES.VENDOR_TERM, 'release_term': 0, 'opt_out': 'N', 'is_automatic_rollover': 'Y', 'payment_interval': constants.VENDOR_CONTRACT_PAYMENT_INTERVALS.MONTH, 'pay_after': constants.VENDOR_CONTRACT_PAY_AFTER_INTERVALS.AFTER_30, 'digital_split': 0.75, 'currency_code': 'USD', 'country_exclusion': ['USA'], 'store_exclusion': ['1', '2'], 'distribution_type_id': 1, } response = logic.create_vendor_contract(**post_data) expected_model_params = dict(post_data) del expected_model_params['currency_code'] del expected_model_params['country_exclusion'] del expected_model_params['distribution_type_id'] del expected_model_params['store_exclusion'] expected_model_params['currency_id'] = currency.currency_id assert response.status == 201 assert response.message == VendorContractDetailSchema().dump(vendor_contract) mock_models.VendorContract.create.assert_called_once_with(**expected_model_params) mock_save_carveouts.assert_called_once_with( 1, { 'service': [ {'service_id': '1', 'distribution_types': [1]}, {'service_id': '2', 'distribution_types': [1]}, ], 'country': ['US'], }, ) @patch('abacus_legacy_sync.logic.vendor_contract._validate_store_exclusion') @patch('abacus_legacy_sync.logic.vendor_contract._validate_country_exclusion') @patch('abacus_legacy_sync.models.ows.carveouts_python.save_account_carveouts') def test_update_vendor_contract_store_exclusion( mock_save_carveouts, mock_validate_country_exclusion, mock_validate_store_exclusion ): """Test update_vendor_contract function for store_exclusion.""" vendor_contract = VendorContractFactory.create() mock_validate_country_exclusion.return_value = ['US'] mock_validate_store_exclusion.return_value = ['1', '2'] mock_put_data = { 'country_exclusion': ['USA'], 'store_exclusion': ['1', '2'], 'distribution_type_id': 1, } response = logic.update_vendor_contract(vendor_contract, **mock_put_data) assert response.status == 200 mock_validate_store_exclusion.assert_called_once_with( mock_put_data['store_exclusion'] ) mock_save_carveouts.assert_called_once_with( 1, { 'service': [ {'service_id': '1', 'distribution_types': [1]}, {'service_id': '2', 'distribution_types': [1]}, ], 'country': ['US'], }, ) @patch('abacus_legacy_sync.logic.vendor_contract.models') def test_validate_store_exclusion_success(mock_models): """Test _validate_store_exclusion function.""" customer_master_master = CustomerMasterMasterFactory.create() customer_master_master_id = str(customer_master_master.customer_master_master_id) mock_models.CustomerMasterMaster.get_by_ids.return_value = [customer_master_master] mock_store_exclusion = [customer_master_master_id] result = logic._validate_store_exclusion(mock_store_exclusion) assert result == mock_store_exclusion @patch('abacus_legacy_sync.logic.vendor_contract.models') def test_validate_store_exclusion_error(mock_models): """Test _validate_store_exclusion function when store_ids not found.""" mock_models.CustomerMasterMaster.get_by_ids.return_value = [] mock_store_exclusion = ['1234', '143'] with pytest.raises(ValidationError) as err: logic._validate_store_exclusion(mock_store_exclusion) assert err.value.args[0] == error.ERROR_INVALID_STORE_ID.format( ','.join(mock_store_exclusion) )