"""Tests for contract model.""" from datetime import date from decimal import Decimal from unittest.mock import patch import pytest from abacus_contract.constants.constants import ( ADVANCE_STATUSES, CONTRACT_LIFECYCLE_STATUSES, CONTRACT_TYPES, ) from abacus_contract.models.contract import Contract from abacus_contract.models.deleted_contract import DeletedContract from abacus_contract.tests.utils.factories import ( AccountContractFactory, ContractAdvanceFactory, ContractFactory, ContractLifecycleFactory, ContractReserveFactory, LegacyContractFactory, ReferenceSapProfitCenterFactory, ReferenceSigningEntityFactory, SigningEntitySapProfitCenterFactory, ) from royalties.models.run_controller import RunController from royalties.tests.utils.factories import ( AccountingRunFactory, RunControllerContractFactory, RunControllerFactory, ) FF_PATH = 'abacus_contract.models.contract.is_single_supply_chain_company_codes_enabled' def test_get_contract_by_id(): """Test getting a contract by it's contract_id.""" contract = ContractFactory.create() contract_info = Contract.get_by_id(contract.contract_id) assert contract_info.contract_id == contract.contract_id assert contract_info.contract_name == contract.contract_name assert contract_info.summary_note == contract.summary_note assert contract_info.general_note == contract.general_note def test_get_contracts_by_ids(): """Test getting a list of contracts by a list of contract_ids.""" contracts = ContractFactory.create_batch(5) contract_ids = [c.contract_id for c in contracts] res = Contract.get_by_ids(contract_ids) assert len(res) == len(contracts) assert set(contract.contract_id for contract in res) == set(contract_ids) def test_get_contract_with_legacy(): """Create a contract attachment.""" contract = ContractFactory.create() legacy_contract = LegacyContractFactory.create(contract=contract) contract_info = Contract.get_by_id(contract.contract_id) assert contract_info.contract_id == contract.contract_id assert contract_info.contract_name == contract.contract_name assert ( contract_info.legacy_contract.legacy_contract_id == legacy_contract.legacy_contract_id ) def test_get_by_accounts(create_mock_account): """Test getting contracts by accounts.""" contracts = ContractFactory.create_batch(3) for contract in contracts[:2]: account_id = contract.contract_id AccountContractFactory.create(contract=contract, account_id=account_id) res = Contract.get_by_accounts([1, 2]) assert len(res) == 2 assert res == contracts[:2] assert contracts[-1] not in res def test_by_oa_contract_ids(): """Testing getting contracts by orchard admin contract ids.""" oa_contract_ids = [1, 2] contracts = ContractFactory.create_batch(5) for ind, oa_contract_id in enumerate(oa_contract_ids): LegacyContractFactory.create( contract=contracts[ind], oa_contract_id=oa_contract_id ) assert Contract.count() == 5 assert Contract.get_by_legacy_contract_ids(oa_contract_ids) == [ contracts[0], contracts[1], ] def test_find_by_name(): """Test getting contracts by contract_name.""" ContractFactory.create(contract_name='Test Contract') assert Contract.find_by_name("Doesn't exist") is None found = Contract.find_by_name('test contract') assert found.contract_name == 'Test Contract' @pytest.mark.db('mysql') def test_get_filtered_query(mock_contracts): """Get contracts by request parameters.""" contracts = Contract.get_filtered_query(contract_name='test').all() assert len(contracts) == 2 assert ( all( [ contract.contract_name in ['Test Distribution Contract', 'Test NR Contract'] for contract in contracts ] ) is True ) contracts = Contract.get_filtered_query(contract_name='distribution contract').all() assert len(contracts) == 4 assert ( all( [ contract.contract_name in [ 'Distribution Contract 1', 'Distribution Contract 2', 'Distribution Contract 3', 'Test Distribution Contract', ] for contract in contracts ] ) is True ) # A contract name having '\' character contracts = Contract.get_filtered_query(contract_name='A1 LaFlare\\Amigo').all() assert len(contracts) == 1 assert contracts[0].contract_name == 'A1 LaFlare\\Amigo Records, LLC' # A contract name having '%' character contracts = Contract.get_filtered_query(contract_name='%').all() assert len(contracts) == 1 assert contracts[0].contract_name == 'Dylan Bukov 100% t/a Dybbukk' # query by contract id contracts = Contract.get_filtered_query(search_term='11').all() assert len(contracts) == 1 assert contracts[0].contract_name == 'NR Contract 2' assert contracts[0].contract_id == 11 # query by account ids contracts = Contract.get_filtered_query(contract_name=None, account_ids=[2]).all() assert len(contracts) == 7 assert ( all( [ contract.contract_name in [ 'NR Contract 1', 'NR Contract 2', 'NR Contract 3', 'Dylan Bukov 100% t/a Dybbukk', 'Test NR Contract', 'Contract 8675309', 'Contract sort order', ] for contract in contracts ] ) is True ) @pytest.mark.db('mysql') def test_get_filtered_query_by_search_term(mock_contracts): """Get contracts by search term.""" query = Contract.get_filtered_query(search_term='8675309') contracts = query.all() assert len(contracts) == 2 assert contracts[0].contract_id == 8675309 assert contracts[1].contract_name == 'Contract 8675309' def test_get_filtered_query_by_contract_type(mock_contracts): """Get contracts by request parameter contract_type.""" contracts = Contract.get_filtered_query( contract_type=CONTRACT_TYPES.DISTRIBUTION ).all() assert len(contracts) == 7 assert ( all( [ contract.contract_name in [ 'Distribution Contract 1', 'Distribution Contract 2', 'Distribution Contract 3', 'A1 LaFlare\\Amigo Records, LLC', 'Test Distribution Contract', 'Contract 8675309', 'Contract sort order', ] for contract in contracts ] ) is True ) @pytest.mark.skip(reason='This endpoint is deprecated; skip failing tests') def test_get_contract_vat_info_by_contract_ids( create_mock_account, create_mock_account_tax_info ): """Testing getting contracts vat info by list of contract ids.""" contract_ids = [1, 2] for contract_id in contract_ids: contract = ContractFactory.create(contract_id=contract_id) AccountContractFactory.create(contract=contract, account_id=contract_id) result = Contract.get_contract_vat_info_by_contract_ids(contract_ids) assert result == [ (1, 1, 'GBR', True, Decimal('20.00'), Decimal('20.00')), (2, 2, 'GBR', False, Decimal('20.00'), None), ] def test_get_contract_contract_reserve(): """Test getting contract contract reserve returns un-deleted contract.""" contract = ContractFactory.create() reserve = ContractReserveFactory.create( contract=contract, reserve_rate=25, reserve_release_offset_in_months=1, installments_in_months=4, release_schedule=[25, 25, 25, 25], ) ContractReserveFactory.create( contract=contract, reserve_rate=50, reserve_release_offset_in_months=1, installments_in_months=2, release_schedule=[25, 25], deleted_by='test', deleted_at='2022-01-17', ) result = Contract.get_by_id_or_error(contract.contract_id) assert result.contract_reserve == reserve contract2 = ContractFactory.create() result = Contract.get_by_id_or_error(contract2.contract_id) assert not result.contract_reserve def test_stream_all(create_mock_account, create_mock_account_payment_term): """Test stream_all query.""" contracts = ContractFactory.create_batch(3) excluded_contract = ContractFactory.create(is_excluded_from_accounting_run=True) for contract in contracts: AccountContractFactory.create(contract=contract) init_contract_lifecycle = ContractLifecycleFactory.create( contract=contracts[0], lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.INIT ) active_contract_lifecycle = ContractLifecycleFactory.create( contract=contracts[1], lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.ACTIVE, lifecycle_term_end=date(2050, 1, 1), ) terminated_contract_lifecycle = ContractLifecycleFactory.create( contract=contracts[2], lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.TERMINATED, lifecycle_term_start=date(2020, 1, 1), lifecycle_term_end=date(2022, 1, 1), ) res = Contract.stream_all().fetchall() assert len(res) == 2 assert init_contract_lifecycle.contract_id not in ( contract.contract_id for contract in res ) assert excluded_contract.contract_id not in ( contract.contract_id for contract in res ) assert res[0].term_start == active_contract_lifecycle.lifecycle_term_start assert res[0].term_end == active_contract_lifecycle.lifecycle_term_end assert res[1].term_start == terminated_contract_lifecycle.lifecycle_term_start assert res[1].term_end == terminated_contract_lifecycle.lifecycle_term_end assert contracts[1].term_start != active_contract_lifecycle.lifecycle_term_start assert contracts[1].term_end != active_contract_lifecycle.lifecycle_term_end assert contracts[2].term_start != terminated_contract_lifecycle.lifecycle_term_start assert contracts[2].term_end != terminated_contract_lifecycle.lifecycle_term_end def test_get_filtered_query_included_excluded_run(mock_contracts): """Get contracts by request parameter is_excluded_from_accounting_run.""" # get excluded contracts contracts = Contract.get_filtered_query(is_excluded_from_accounting_run=1).all() assert len(contracts) == 6 assert ( all( [ contract.contract_name in [ 'A1 LaFlare\\Amigo Records, LLC', 'Test Distribution Contract', 'NR Contract 1', 'NR Contract 2', 'NR Contract 3', 'Test NR Contract', ] for contract in contracts ] ) is True ) # get included contracts contracts = Contract.get_filtered_query(is_excluded_from_accounting_run=0).all() assert len(contracts) == 6 assert ( all( [ contract.contract_name in [ 'Distribution Contract 1', 'Distribution Contract 2', 'Distribution Contract 3', 'Dylan Bukov 100% t/a Dybbukk', 'Contract 8675309', 'Contract sort order', ] for contract in contracts ] ) is True ) def test_get_filtered_query_contract_statuses(mock_contracts): """Get contracts by request parameter contract_statuses.""" contracts = Contract.get_filtered_query(contract_statuses='init,terminated').all() assert len(contracts) == 6 assert ( all( [ contract.contract_name in [ 'Distribution Contract 1', 'NR Contract 2', 'Test Distribution Contract', 'NR Contract 3', 'Dylan Bukov 100% t/a Dybbukk', 'Test NR Contract', ] for contract in contracts ] ) is True ) contracts = Contract.get_filtered_query( contract_statuses='active,in_collection_period,to_be_terminated' ).all() assert len(contracts) == 5 assert ( all( [ contract.contract_name in [ 'Distribution Contract 3', 'A1 LaFlare\\Amigo Records, LLC', 'NR Contract 1', 'Contract 8675309', 'Contract sort order', ] for contract in contracts ] ) is True ) def test_get_filtered_query_by_run_controller_ids(mock_contracts): """Get contracts by request parameter run_controller_ids.""" run_controllers = RunController.query.all() run_controller_ids = [ str(run_controller.run_controller_id) for run_controller in run_controllers ] contracts = Contract.get_filtered_query( run_controller_ids=','.join(run_controller_ids) ).all() assert len(contracts) == 12 assert ( all( [ contract.contract_name in [ 'Distribution Contract 1', 'Distribution Contract 2', 'Distribution Contract 3', 'A1 LaFlare\\Amigo Records, LLC', 'Test Distribution Contract', 'NR Contract 1', 'NR Contract 2', 'NR Contract 3', 'Dylan Bukov 100% t/a Dybbukk', 'Test NR Contract', 'Contract 8675309', 'Contract sort order', ] for contract in contracts ] ) is True ) contracts = Contract.get_filtered_query(run_controller_ids='2').all() assert len(contracts) == 0 def test_get_sap_profit_center_by_contract_id_returns_expected_data( create_mock_account, ): """Test valid contract returns correct profit center info.""" profit_center = ReferenceSapProfitCenterFactory.create() signing_entity = ReferenceSigningEntityFactory.create( reference_sap_profit_center=profit_center ) contract = ContractFactory.create(reference_signing_entity=signing_entity) account_id = contract.contract_id AccountContractFactory.create(contract=contract, account_id=account_id) result = Contract.get_sap_profit_center_by_contract_id(contract.contract_id) assert result is not None assert result.contract_id == contract.contract_id assert result.Prctr == profit_center.profit_center assert result.Bukrs == profit_center.company_code def test_get_sap_profit_center_by_contract_id_returns_none_for_invalid_id(): """Test that method returns None when contract_id does not exist.""" result = Contract.get_sap_profit_center_by_contract_id(999999) assert result is None def test_get_contract_with_profit_center_missing_relationships(): """Test that method returns None if required relationships are missing.""" contract = ContractFactory.create() result = Contract.get_sap_profit_center_by_contract_id(contract.contract_id) assert result is None @patch(FF_PATH, return_value=True) def test_get_sap_profit_center_by_contract_id_ff_on_uses_contract_pc( _, create_mock_account ): """FF ON: PC is resolved from contract.reference_sap_profit_center_id directly.""" se_pc = ReferenceSapProfitCenterFactory.create() se = ReferenceSigningEntityFactory.create(reference_sap_profit_center=se_pc) # Contract's PC differs from SE's legacy PC — only the new path will resolve correctly. contract_pc = ReferenceSapProfitCenterFactory.create() SigningEntitySapProfitCenterFactory.create( reference_signing_entity=se, reference_sap_profit_center=contract_pc ) contract = ContractFactory.create( reference_signing_entity=se, reference_sap_profit_center_id=contract_pc.reference_sap_profit_center_id, ) AccountContractFactory.create(contract=contract, account_id=contract.contract_id) result = Contract.get_sap_profit_center_by_contract_id(contract.contract_id) assert result is not None assert result.Prctr == contract_pc.profit_center assert result.Bukrs == contract_pc.company_code @patch(FF_PATH, return_value=False) def test_get_sap_profit_center_by_contract_id_ff_off_uses_se_pc(_, create_mock_account): """FF OFF: PC is resolved via the legacy SE→PC join, ignoring contract.reference_sap_profit_center_id.""" se_pc = ReferenceSapProfitCenterFactory.create() se = ReferenceSigningEntityFactory.create(reference_sap_profit_center=se_pc) # Contract.PC is set to a different PC, but FF OFF means we don't read it. other_pc = ReferenceSapProfitCenterFactory.create() SigningEntitySapProfitCenterFactory.create( reference_signing_entity=se, reference_sap_profit_center=other_pc ) contract = ContractFactory.create( reference_signing_entity=se, reference_sap_profit_center_id=other_pc.reference_sap_profit_center_id, ) AccountContractFactory.create(contract=contract, account_id=contract.contract_id) result = Contract.get_sap_profit_center_by_contract_id(contract.contract_id) assert result is not None assert result.Prctr == se_pc.profit_center assert result.Bukrs == se_pc.company_code def test_get_sap_profit_center_by_contract_id_both_ff_branches_identical_for_backfilled_contract( create_mock_account, ): """For a backfilled contract (SE.PC == contract.PC), both FF branches return identical output.""" pc = ReferenceSapProfitCenterFactory.create() se = ReferenceSigningEntityFactory.create(reference_sap_profit_center=pc) contract = ContractFactory.create( reference_signing_entity=se, reference_sap_profit_center_id=pc.reference_sap_profit_center_id, ) AccountContractFactory.create(contract=contract, account_id=contract.contract_id) with patch(FF_PATH, return_value=False): off_result = Contract.get_sap_profit_center_by_contract_id(contract.contract_id) with patch(FF_PATH, return_value=True): on_result = Contract.get_sap_profit_center_by_contract_id(contract.contract_id) assert off_result is not None and on_result is not None assert tuple(off_result) == tuple(on_result) def test_can_be_deleted_true(): """Test checking if a contract can be deleted when it can.""" contract = ContractFactory.create() result = Contract.can_be_deleted(contract.contract_id) assert result is True @pytest.mark.skip(reason='This test case works locally but is failing on Jenkins.') def test_accounting_run_with_revenue_can_not_be_deleted( mock_revenue_associated_with_accounting_run, ): """Test checking if a contract can be deleted if revenue is associated with contract.""" result = Contract.can_be_deleted(9012890) assert result is False def test_accounting_run_with_no_revenue_can_be_be_deleted(create_mock_account): """Test checking if a contract can be deleted if no revenue is associated with contract.""" mock_contract = ContractFactory.create( contract_name='Contract 8675309', ) AccountContractFactory.create(contract=mock_contract, account_id=2) mock_run_controller = RunControllerFactory.create() RunControllerContractFactory.create( run_controller=mock_run_controller, contract=mock_contract ) AccountingRunFactory.create(run_controller=mock_run_controller) result = Contract.can_be_deleted(mock_contract.contract_id) assert result is True def test_can_be_deleted_false(): """Test checking if a contract can be deleted when it cannot.""" contract = ContractFactory.create() ContractAdvanceFactory.create( contract=contract, advance_status=ADVANCE_STATUSES.IN_REVIEW ) result = Contract.can_be_deleted(contract.contract_id) assert result is False @pytest.mark.skip(reason='This endpoint is under development.') def test_delete(): """Test deleting a contract.""" contract_id = ContractFactory.create().contract_id Contract.delete(contract_id) existing_contract = Contract.get_by_id(contract_id) assert existing_contract is None deleted_contract = DeletedContract.get_by_id(1) assert deleted_contract.contract_id == contract_id assert deleted_contract.contract_data is not None def get_primary_contract_without_lifecycle_by_account_id_and_contract_type( create_mock_account, ): """Test get_primary_contract_by_account_id_and_contract_type query. The contract without lifecycle rules is primary. """ account_id = 1 primary_contract = ContractFactory.create(is_primary_contract=True) AccountContractFactory.create(contract=primary_contract, account_id=account_id) contracts = ContractFactory.create_batch(3) for contract in contracts: AccountContractFactory.create(contract=contract, account_id=account_id) ContractLifecycleFactory.create( contract=contracts[0], lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.INIT ) ContractLifecycleFactory.create( contract=contracts[1], lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.ACTIVE, lifecycle_term_end=date(2050, 1, 1), ) ContractLifecycleFactory.create( contract=contracts[2], lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.TERMINATED, lifecycle_term_start=date(2020, 1, 1), lifecycle_term_end=date(2022, 1, 1), ) res = Contract.get_primary_contract_by_account_id_and_contract_type( account_id=account_id, contract_type=CONTRACT_TYPES.DISTRIBUTION ) assert len(res) == 1 assert res.contract_id == primary_contract.contract_id def get_primary_contract_with_lifecycle_by_account_id_and_contract_type( create_mock_account, ): """Test get_primary_contract_by_account_id_and_contract_type query. The contract having lifecycle rules is primary. """ account_id = 1 contract = ContractFactory.create() AccountContractFactory.create(contract=contract, account_id=account_id) contracts = ContractFactory.create_batch(3) primary_contract = ContractFactory.create(is_primary_contract=True) AccountContractFactory.create(contract=primary_contract, account_id=2) for contract in contracts: AccountContractFactory.create(contract=contract, account_id=account_id) ContractLifecycleFactory.create( contract=contracts[0], lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.TERMINATED ) ContractLifecycleFactory.create( contract=contracts[1], lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.TERMINATED ) ContractLifecycleFactory.create( contract=contracts[2], lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.TERMINATED ) res = Contract.get_primary_contract_by_account_id_and_contract_type( account_id=2, contract_type=CONTRACT_TYPES.DISTRIBUTION ) assert len(res) == 1 assert res.contract_id == primary_contract.contract_id def get_primary_terminated_contract_by_account_id_and_contract_type( create_mock_account, ): """Test get_primary_contract_by_account_id_and_contract_type query. All contracts are terminated for account 1. """ account_id = 1 contract = ContractFactory.create() AccountContractFactory.create(contract=contract, account_id=account_id) contracts = ContractFactory.create_batch(3) primary_contract = ContractFactory.create(is_primary_contract=True) AccountContractFactory.create(contract=primary_contract, account_id=2) for contract in contracts: AccountContractFactory.create(contract=contract, account_id=account_id) ContractLifecycleFactory.create( contract=contracts[0], lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.ACTIVE ) ContractLifecycleFactory.create( contract=contracts[0], lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.TERMINATED ) ContractLifecycleFactory.create( contract=contracts[1], lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.TERMINATED ) ContractLifecycleFactory.create( contract=contracts[2], lifecycle_status=CONTRACT_LIFECYCLE_STATUSES.TERMINATED ) res = Contract.get_primary_contract_by_account_id_and_contract_type( account_id=1, contract_type=CONTRACT_TYPES.DISTRIBUTION ) assert len(res) == 0