"""Tests for contract model.""" from datetime import date from decimal import Decimal import pytest from abacus_contract.constants.constants import ADVANCE_STATUSES from abacus_contract.constants.constants import CONTRACT_LIFECYCLE_STATUSES from abacus_contract.constants.constants import CONTRACT_TYPES from abacus_contract.models.contract import Contract from abacus_contract.models.deleted_contract import DeletedContract from tests.utils.factories import AccountContractFactory from tests.utils.factories import ContractAdvanceFactory from tests.utils.factories import ContractFactory from tests.utils.factories import ContractLifecycleFactory from tests.utils.factories import ContractReserveFactory from tests.utils.factories import LegacyContractFactory from tests.utils.factories import ReferenceSapProfitCenterFactory from tests.utils.factories import ReferenceSigningEntityFactory 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 not contracts[-1] 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' 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 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.""" contracts = Contract.get_filtered_query( run_controller_ids='1,2' ).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_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 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='Skip failing test') 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 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