"""Run controller contract model tests.""" from abacus_contract.tests.utils.factories import ( AccountContractFactory, ContractFactory, ) from royalties.constants.constants import CONTRACT_TYPES from royalties.models.run_controller_contract import RunControllerContract from royalties.tests.utils.factories import RunControllerContractFactory def test_create(run_controller_fixtures, remove_fks): """Create a run controller contract.""" for run_controller in run_controller_fixtures: RunControllerContractFactory.create( run_controller_id=run_controller.run_controller_id ) assert len(RunControllerContract.query.all()) == len(run_controller_fixtures) def test_get_by_account_id(account_contract_fixtures): """Test getting a run controller contract by account_id.""" account_id = 1 contract = ContractFactory.create() AccountContractFactory.create(account_id=account_id, contract=contract) run_controller_contract = RunControllerContractFactory.create(contract=contract) res = RunControllerContract.get_by_account_id(account_id) assert res assert res[0].run_controller_id == run_controller_contract.run_controller_id assert res[0].contract_type == run_controller_contract.run_controller.contract_type def get_by_account_id(): """Get all run controllers, filtered (or not) by contract type.""" contract = ContractFactory.create() account_id = 1 run_controllers = [ RunControllerContractFactory.create( contract_type=CONTRACT_TYPES.DISTRIBUTION, contract=contract ), RunControllerContractFactory.create( contract_type=CONTRACT_TYPES.LEGACY_DISTRIBUTION, contract=contract ), ] all_results = RunControllerContract.get_by_account_id().all() assert len(all_results) == len(run_controllers) distribution_results = RunControllerContract.get_by_account_id( account_id, contract_type=CONTRACT_TYPES.DISTRIBUTION ).all() assert len(distribution_results) == 1 assert distribution_results[0].contract_type == CONTRACT_TYPES.DISTRIBUTION legacy_distribution_results = RunControllerContract.get_by_account_id( account_id, contract_type=CONTRACT_TYPES.LEGACY_DISTRIBUTION ).all() assert len(legacy_distribution_results) == 1 assert ( legacy_distribution_results[0].contract_type == CONTRACT_TYPES.LEGACY_DISTRIBUTION ) def test_get_by_contract_id(): """Get run_controller_contract by contract_id.""" run_controller_contract = RunControllerContractFactory.create() res = RunControllerContract.get_by_contract_id( run_controller_contract.contract_id ).first() assert ( res.run_controller_contract_id == run_controller_contract.run_controller_contract_id ) def test_get_by_contract_ids(): """Get run_controller_contracts by contract_ids.""" run_controller_contracts = RunControllerContractFactory.create_batch(2) contract_ids = [ run_controller_contract.contract_id for run_controller_contract in run_controller_contracts ] res = RunControllerContract.get_by_contract_ids(contract_ids) assert len(res) == len(run_controller_contracts) assert ( res[0].run_controller_contract_id == run_controller_contracts[0].run_controller_contract_id ) assert ( res[1].run_controller_contract_id == run_controller_contracts[1].run_controller_contract_id )