"""Test run controller contract logic.""" from operator import itemgetter from unittest.mock import patch import pytest from abacus_account.tests.utils.factories import AccountFactory from abacus_contract.tests.utils.factories import ( AccountContractFactory, ContractFactory, ContractLifecycleFactory, ) from royalties.constants.constants import CONTRACT_LIFECYCLE_STATUSES, CONTRACT_TYPES from royalties.constants.error import ( ERROR_DIFFERENT_RUN_CONTROLLER, ERROR_RUN_CONTROLLER_DOES_NOT_MATCH_CONTRACT_TYPE, ) from royalties.logic import run_controller_contract as logic from royalties.logic.run_controller_contract import RunControllerValidationError from royalties.schemas.run_controller_contract import RunControllerContractSchema from royalties.tests.utils.factories import ( RunControllerContractFactory, RunControllerFactory, ) @patch('royalties.logic.run_controller_contract._base_create_run_controller_contract') @patch('royalties.logic.run_controller_contract.models') def test_create_run_controller_contract(mock_models, mock_create): """Test successfully creating a run_controller_contract.""" account_id = 1 contract_type = CONTRACT_TYPES.DISTRIBUTION mock_response = { 'run_controller_contract_id': 1, 'contract_id': 1, 'run_controller_id': 1, } mock_create.return_value = mock_response mock_models.RunControllerContract.commit_changes.return_value = True params = { 'account_id': account_id, 'contract_type': contract_type, 'contract_id': mock_response.get('contract_id'), 'run_controller_id': mock_response.get('run_controller_id'), } res = logic.create_run_controller_contract(**params) assert res.status == 201 assert res.message == mock_response mock_create.assert_called_once_with( account_id, mock_response.get('contract_id'), mock_response.get('run_controller_id'), contract_type, ) mock_models.RunControllerContract.commit_changes.assert_called_once() @patch('royalties.logic.run_controller_contract.models') def test_get_run_controller_contracts_by_account_success(mock_models, remove_fks): """Test getting run_controller_contracts by account.""" account_id = 1 contract_type = CONTRACT_TYPES.DISTRIBUTION run_controller_contracts = RunControllerContractFactory.create_batch(2) mock_models.RunControllerContract.get_by_account_id.return_value = ( RunControllerContractSchema().dump(run_controller_contracts, many=True) ) res = logic.get_run_controller_contracts_by_account(account_id, contract_type) assert res.status == 200 assert len(res.message) == len(run_controller_contracts) mock_models.RunControllerContract.get_by_account_id.assert_called_once_with( account_id, contract_type ) @patch('royalties.logic.run_controller_contract.models') def test_get_run_controller_contracts_by_account_error(mock_models): """Test an error is raised if account_id is not specified.""" account_id = None contract_type = CONTRACT_TYPES.DISTRIBUTION res = logic.get_run_controller_contracts_by_account(account_id, contract_type) assert res.status == 400 assert res.errors mock_models.RunControllerContract.get_by_account_id.assert_not_called() @patch('royalties.logic.run_controller_contract.models') def test_get_contracts_run_controllers_with_dataload_format(mock_models): """Test getting run_controller_contracts with dataload format.""" run_controller_contracts = RunControllerContractFactory.create_batch(2) contract_ids = [rcc.contract_id for rcc in run_controller_contracts] contract_ids.append(9999) mock_models.RunControllerContract.get_by_contract_ids.return_value = ( run_controller_contracts ) mock_models.RunControllerContractSchema.dump.return_value = run_controller_contracts res = logic.get_contracts_run_controllers_with_dataload_format(contract_ids) assert res.status == 200 assert res.message == [ { 'data': { 'run_controller_contract_id': run_controller_contracts[ 0 ].run_controller_contract_id, 'contract_id': run_controller_contracts[0].contract_id, 'run_controller_id': run_controller_contracts[0].run_controller_id, 'run_controller_name': run_controller_contracts[ 0 ].run_controller.run_controller_name, } }, { 'data': { 'run_controller_contract_id': run_controller_contracts[ 1 ].run_controller_contract_id, 'contract_id': run_controller_contracts[1].contract_id, 'run_controller_id': run_controller_contracts[1].run_controller_id, 'run_controller_name': run_controller_contracts[ 1 ].run_controller.run_controller_name, } }, {'data': None}, ] mock_models.RunControllerContract.get_by_contract_ids.assert_called_once_with( contract_ids ) @patch('royalties.logic.run_controller_contract._validate_run_controller_contract') @patch('royalties.logic.run_controller_contract.models') @patch('royalties.logic.run_controller_contract.add_run_controller_to_acct_period') def test_base_create_run_controller_contract( mock_add_to_period, mock_models, mock_validate, account_contract_fixtures ): """Test base create run_controller_contract method.""" account_id = 1 contract_id = 123 contract_type = CONTRACT_TYPES.LEGACY_DISTRIBUTION run_controller = RunControllerFactory.create() RunControllerContractFactory.create(contract_id=1, run_controller=run_controller) mock_models.RunController.get_by_id_or_error.return_value = run_controller mock_models.RunController.count_contracts.return_value = 10 mock_add_to_period.return_value = True mock_validate.return_value = None mock_models.RunControllerContract.build.return_value = { 'contract_id': contract_id, 'contract_type': contract_type, 'run_controller_id': run_controller.run_controller_id, } res = logic._base_create_run_controller_contract( account_id, contract_id, run_controller.run_controller_id, contract_type ) assert res mock_validate.assert_called_once_with(account_id, run_controller, contract_type) mock_add_to_period.assert_not_called() mock_models.RunControllerContract.build.assert_called_once_with( contract_id=contract_id, run_controller_id=run_controller.run_controller_id ) @patch('royalties.logic.run_controller_contract._validate_run_controller_contract') @patch('royalties.logic.run_controller_contract.models') @patch('royalties.logic.run_controller_contract.add_run_controller_to_acct_period') def test_base_create_run_controller_contract_no_contracts( mock_add_to_period, mock_models, mock_validate ): """Test run controller is added to acct period when first contract is attached.""" account_id = 1 contract_id = 123 contract_type = CONTRACT_TYPES.LEGACY_DISTRIBUTION run_controller = RunControllerFactory.create() mock_models.RunController.get_by_id_or_error.return_value = run_controller mock_models.RunController.count_contracts.return_value = 0 mock_add_to_period.return_value = True mock_validate.return_value = None mock_models.RunControllerContract.build.return_value = { 'contract_id': contract_id, 'run_controller_id': run_controller.run_controller_id, } res = logic._base_create_run_controller_contract( account_id, contract_id, contract_type, run_controller.run_controller_id ) assert res mock_add_to_period.assert_called_once_with(run_controller) @patch('royalties.logic.run_controller_contract.models') def test_validate_run_controller_contract_is_valid( mock_models, account_contract_fixtures ): """Test run controller contract params pass validation.""" account_id = 1 contract_type = CONTRACT_TYPES.LEGACY_DISTRIBUTION run_controller = RunControllerFactory.create() run_controller_contract = RunControllerContractFactory.create( contract_id=1, run_controller=run_controller ) mock_models.RunControllerContract.get_by_account_id.return_value = [ run_controller_contract ] res = logic._validate_run_controller_contract( account_id, run_controller, contract_type ) assert res is None mock_models.RunControllerContract.get_by_account_id.assert_called_once_with( account_id, contract_type ) @patch('royalties.logic.run_controller_contract.models') def test_validate_run_controller_contract_is_invalid( mock_models, account_contract_fixtures ): """Test run controller contract params fail validation.""" account_id = 1 contract_type = CONTRACT_TYPES.LEGACY_DISTRIBUTION run_controller = RunControllerFactory.create() run_controller_contract = RunControllerContractFactory.create( contract_id=1, run_controller=run_controller ) run_controller_contract.run_controller_name = run_controller.run_controller_name mock_models.RunControllerContract.get_by_account_id.return_value = [ run_controller_contract ] different_run_controller = RunControllerFactory.create() error_msg = ERROR_DIFFERENT_RUN_CONTROLLER.format( contract_type=contract_type.capitalize(), run_controller_name=run_controller.run_controller_name, ) with pytest.raises(RunControllerValidationError, match=error_msg): logic._validate_run_controller_contract( account_id, different_run_controller, contract_type ) @patch( 'royalties.logic.run_controller_contract.is_abacus_contract_page_update_run_controller_ff_enabled' ) def test_update_run_controller_contract_feature_flag_disabled(mock_is_feature_enabled): """Test update run controller contract fails if feature flag is disabled.""" mock_is_feature_enabled.return_value = False res = logic.update_run_controller_contract(contract_id=1, run_controller_id=1) assert res.status == 403 assert res.errors['message'] == 'Unauthorized to update run controller contract.' assert res.errors['code'] == 'Unauthorized' @patch( 'royalties.logic.run_controller_contract.is_abacus_contract_page_update_run_controller_ff_enabled' ) @patch('royalties.logic.run_controller_contract.response.create_fatal_response') def test_update_run_controller_contract_fetch_contract_lifecycle_not_found( mock_create_fatal_response, mock_is_feature_enabled ): """Test update fails if contract lifecycle is not found. Feature flag is on. """ mock_is_feature_enabled.return_value = True contract_id = 301 account_id = 100 AccountFactory.create(account_id=account_id) # contract_lifecycle is None contract = ContractFactory.create( contract_id=contract_id, contract_type=CONTRACT_TYPES.DISTRIBUTION ) contract.account_contract = AccountContractFactory.create( account_id=account_id, contract_id=contract.contract_id ) logic.update_run_controller_contract(contract_id=contract_id, run_controller_id=1) mock_create_fatal_response.assert_called_once_with( message=f'Contract lifecycle not found for contract {contract_id}.' ) @patch( 'royalties.logic.run_controller_contract.is_abacus_contract_page_update_run_controller_ff_enabled' ) @patch('royalties.logic.run_controller_contract.validation_error') def test_update_run_controller_contract_fetch_contract_lifecycle_terminated( mock_validation_error, mock_is_feature_enabled ): """Test update fails if contract lifecycle is terminated. Feature flag is on. """ mock_is_feature_enabled.return_value = True contract_id = 302 account_id = 100 AccountFactory.create(account_id=account_id) contract = ContractFactory.create( contract_id=contract_id, contract_type=CONTRACT_TYPES.DISTRIBUTION ) # Lifecycle status is terminated ContractLifecycleFactory.create( contract=contract, lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.TERMINATED ) contract.account_contract = AccountContractFactory.create( account_id=account_id, contract_id=contract.contract_id ) logic.update_run_controller_contract(contract_id=contract_id, run_controller_id=1) mock_validation_error.assert_called_once_with( { 'contract_id': f'Contract {contract_id} is already terminated. Run controller cannot be updated.' } ) @patch( 'royalties.logic.run_controller_contract.is_abacus_contract_page_update_run_controller_ff_enabled' ) @patch('royalties.logic.run_controller_contract.response.create_fatal_response') def test_update_run_controller_contract_contract_not_found( mock_create_fatal_response, mock_is_feature_enabled ): """Test update fails if the initial contract is not found and feature flag is on.""" mock_is_feature_enabled.return_value = True # No contract exists for this ID contract_id = 9999 logic.update_run_controller_contract(contract_id=contract_id, run_controller_id=1) # Results into no contract lifecycle error. mock_create_fatal_response.assert_called_once_with( message=f'Contract lifecycle not found for contract {contract_id}.' ) @patch( 'royalties.logic.run_controller_contract.is_abacus_contract_page_update_run_controller_ff_enabled' ) @patch('royalties.logic.run_controller_contract.response.create_fatal_response') def test_update_run_controller_contract_account_not_found( mock_create_fatal_response, mock_is_feature_enabled ): """Test update fails if the account for the contract is not found. Feature flag is on. """ mock_is_feature_enabled.return_value = True contract_id = 302 contract = ContractFactory.create( contract_id=contract_id, contract_type=CONTRACT_TYPES.DISTRIBUTION ) ContractLifecycleFactory.create( contract=contract, lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.ACTIVE ) # No account_contract object logic.update_run_controller_contract(contract_id=contract_id, run_controller_id=1) mock_create_fatal_response.assert_called_once_with( message=f'Account ID not found in contract {contract_id}.' ) @patch( 'royalties.logic.run_controller_contract.is_abacus_contract_page_update_run_controller_ff_enabled' ) @patch('royalties.logic.run_controller_contract.validation_error') def test_update_run_controller_contract_type_mismatch( mock_validation_error, mock_is_feature_enabled ): """Test update fails if run controller's type doesn't match contract's type. Feature flag is on. """ mock_is_feature_enabled.return_value = True contract_id = 101 account_id = 100 account = AccountFactory.create(account_id=account_id) contract = ContractFactory.create( contract_id=contract_id, contract_type=CONTRACT_TYPES.DISTRIBUTION ) ContractLifecycleFactory.create( contract=contract, lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.ACTIVE ) contract.account_contract = AccountContractFactory.create( account_id=account.account_id, contract_id=contract.contract_id ) # The proposed run controller has a different type proposed_run_controller = RunControllerFactory.create( contract_type=CONTRACT_TYPES.NEIGHBOURING_RIGHTS ) logic.update_run_controller_contract( contract_id=contract_id, run_controller_id=proposed_run_controller.run_controller_id, ) mock_validation_error.assert_called_once_with( { 'run_controller_id': ERROR_RUN_CONTROLLER_DOES_NOT_MATCH_CONTRACT_TYPE.format( run_controller_name=proposed_run_controller.run_controller_name, run_controller_contract_type=proposed_run_controller.contract_type, contract_id=contract_id, contract_type=CONTRACT_TYPES.DISTRIBUTION, ) } ) @patch( 'royalties.logic.run_controller_contract.is_abacus_contract_page_update_run_controller_ff_enabled' ) def test_update_run_controller_contract_with_multiple_contract_types_and_lifecycle_statuses_success_ff_on( mock_is_feature_enabled, ): """Test update run controller contract succeeds and feature flag is on. Account contains multiple contracts of mixed types and different lifecycle statuses. """ mock_is_feature_enabled.return_value = True contract_id_1 = 101 contract_id_2 = 102 contract_id_3 = 103 contract_id_4 = 104 account_id = 1 contract_type = CONTRACT_TYPES.DISTRIBUTION AccountFactory.create(account_id=account_id) contract_to_update = ContractFactory.create( contract_id=contract_id_1, contract_type=contract_type ) contract_sibling = ContractFactory.create( contract_id=contract_id_2, contract_type=contract_type ) # Wrong contract type contract_wrong_type = ContractFactory.create( contract_id=contract_id_3, contract_type=CONTRACT_TYPES.NEIGHBOURING_RIGHTS ) contract_terminated = ContractFactory.create( contract_id=contract_id_4, contract_type=contract_type ) for c in [ contract_to_update, contract_sibling, contract_wrong_type, contract_terminated, ]: c.account_contract = AccountContractFactory.create( account_id=account_id, contract_id=c.contract_id ) ContractLifecycleFactory.create( contract=contract_to_update, lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.ACTIVE ) ContractLifecycleFactory.create( contract=contract_sibling, lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.ACTIVE ) ContractLifecycleFactory.create( contract=contract_wrong_type, lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.ACTIVE, ) # Terminated lifecycle ContractLifecycleFactory.create( contract=contract_terminated, lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.TERMINATED, ) proposed_run_controller = RunControllerFactory.create(contract_type=contract_type) current_run_controller = RunControllerFactory.create( contract_type=CONTRACT_TYPES.DISTRIBUTION ) rcc_to_update = RunControllerContractFactory.create( contract=contract_to_update, run_controller=current_run_controller ) rcc_sibling = RunControllerContractFactory.create( contract=contract_sibling, run_controller=current_run_controller ) RunControllerContractFactory.create( contract=contract_wrong_type, run_controller=current_run_controller ) RunControllerContractFactory.create( contract=contract_terminated, run_controller=current_run_controller ) expect_message = [ { 'run_controller_contract_id': rcc_sibling.run_controller_contract_id, 'contract_id': contract_sibling.contract_id, 'run_controller_id': proposed_run_controller.run_controller_id, 'run_controller_name': proposed_run_controller.run_controller_name, }, { 'run_controller_contract_id': rcc_to_update.run_controller_contract_id, 'contract_id': contract_to_update.contract_id, 'run_controller_id': proposed_run_controller.run_controller_id, 'run_controller_name': proposed_run_controller.run_controller_name, }, ] res = logic.update_run_controller_contract( contract_id=contract_id_1, run_controller_id=proposed_run_controller.run_controller_id, ) assert res.status == 200 assert sorted(res.message, key=itemgetter('contract_id')) == sorted( expect_message, key=itemgetter('contract_id') )