"""Run controller model tests.""" from royalties.constants.constants import CONTRACT_TYPES from royalties.models.run_controller import RunController from royalties.tests.utils.factories import ( RunControllerContractFactory, RunControllerFactory, ) def test_create(test_app_request): """Create a run controller.""" RunController.create( contract_type=CONTRACT_TYPES.DISTRIBUTION, run_controller_name='My run controller', ) controllers = RunController.query.all() assert len(controllers) == 1 assert controllers[0].run_controller_name == 'My run controller' assert controllers[0].contract_type == CONTRACT_TYPES.DISTRIBUTION def test_find_by_name(run_controller_fixtures): """Find a run controller by name.""" assert RunController.find_by_name("Doesn't exist") is None found = RunController.find_by_name('High Priority 1') assert found.run_controller_name == 'High Priority 1' def tests_get_all_associated_to_contracts(account_contract_fixtures): """Get a list of all run controllers that are associated to contracts.""" run_controllers = RunControllerFactory.create_batch(2) rc_contract = RunControllerContractFactory.create( run_controller=run_controllers[0], contract_id=1 ) results = RunController.get_all_associated_to_contracts( contract_type=run_controllers[0].contract_type ) assert len(results) < len(run_controllers) assert len(results) == 1 assert results[0].run_controller_id == rc_contract.run_controller_id def test_base_list_query(): """Get all run controllers, filtered (or not) by contract type.""" run_controllers = [ RunControllerFactory.create(contract_type=CONTRACT_TYPES.DISTRIBUTION), RunControllerFactory.create(contract_type=CONTRACT_TYPES.LEGACY_DISTRIBUTION), ] all_results = RunController.base_list_query().all() assert len(all_results) == len(run_controllers) distribution_results = RunController.base_list_query( contract_type=CONTRACT_TYPES.DISTRIBUTION ).all() assert len(distribution_results) == 1 assert distribution_results[0].contract_type == CONTRACT_TYPES.DISTRIBUTION legacy_distribution_results = RunController.base_list_query( 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_list_query_active_only(): """Get only active run controllers.""" run_controllers = RunControllerFactory.create_batch(2) run_controllers[0].deleted_at = '2023-09-18' run_controllers[0].deleted_by = 'someone' res = RunController.base_list_query(active_only=True).all() assert len(res) < len(run_controllers) def test_count_contracts(): """Test counting contracts associated with run controller.""" run_controller = RunControllerFactory.create() RunControllerContractFactory.create_batch(3, run_controller=run_controller) count = run_controller.count_contracts(run_controller.run_controller_id) assert count == 3