"""Tests for contract term model.""" from abacus_contract.constants import constants from abacus_contract.models.contract_term import ContractTerm from tests.utils.factories import AccountContractFactory from tests.utils.factories import ContractFactory from tests.utils.factories import ContractTermFactory def test_create_contract_term(): """Test to create contract_term record.""" mock_contract = ContractFactory.create() mock_contract_term = ContractTerm.create( contract_id=mock_contract.contract_id, contract_term_name='Test contract term', term_type='label', attachments=['123'] ) result = ContractTerm.query.all() assert len(result) == 1 assert result[0] == mock_contract_term def test_get_contract_base_term(): """Test getting base term for specified contract.""" contract = ContractFactory.create() contract_term = ContractTermFactory.create( contract=contract, is_base_term=True ) base_term = ContractTerm.get_contract_base_term(contract_term.contract_id) assert base_term.contract_id == contract.contract_id assert base_term.is_base_term is True def test_get_contract_base_term_soft_deleted(): """Test getting a base term filters out soft-deleted contract_terms.""" contract_term = ContractTermFactory.create(is_base_term=True) contract_term.deleted_at = '2022-01-01' base_term = ContractTerm.get_contract_base_term(contract_term.contract_id) assert not base_term def test_get_contract_terms_by_account_and_term_type(create_mock_account): """Test getting contract terms for a specified account and term_type.""" account_id = 1 tracks_attachments_1 = ['01234', '56781'] tracks_attachments_2 = ['128776', '129751'] term_type = constants.CONTRACT_TERM_TYPES.TRACK attachments = ['128776'] contract_1 = ContractFactory.create() AccountContractFactory.create( account_id=account_id, contract=contract_1 ) ContractTermFactory.create( contract=contract_1, term_type=constants.CONTRACT_TERM_TYPES.TRACK, attachments=tracks_attachments_1 ) contract_2 = ContractFactory.create() AccountContractFactory.create( account_id=account_id, contract=contract_2 ) contract_term_2 = ContractTermFactory.create( contract=contract_2, term_type=constants.CONTRACT_TERM_TYPES.TRACK, attachments=tracks_attachments_2 ) result = ContractTerm.get_contract_terms_by_account_and_term_type( account_id, attachments, term_type) assert result[0].contract_term_id == contract_term_2.contract_term_id assert len(result) == 1 def test_get_contract_terms_by_account_and_term_type_soft_deleted(create_mock_account): """Test getting contract terms filters out soft-deleted records.""" account_id = 1 tracks_attachments = ['128776', '129751'] term_type = constants.CONTRACT_TERM_TYPES.TRACK attachments = ['128776'] contract = ContractFactory.create() AccountContractFactory.create( account_id=account_id, contract=contract ) contract_term = ContractTermFactory.create( contract=contract, term_type=constants.CONTRACT_TERM_TYPES.TRACK, attachments=tracks_attachments ) contract_term.deleted_at = '2022-01-01' result = ContractTerm.get_contract_terms_by_account_and_term_type( account_id, attachments, term_type) assert not result