"""Unit tests for WorksheetPaymentContractAdvance model.""" from datetime import datetime, timedelta from decimal import Decimal from abacus_common_logic.connectors.database import db import pytest from sqlalchemy import select from payment.models.worksheet_payment_contract_advance import ( WorksheetPaymentContractAdvance, ) from tests.utils.factories import WorksheetPaymentContractAdvanceFactory def test_create(mock_statement_periods, mock_exchange_rates, mock_contract_advance): """Test create WorksheetPaymentContractAdvance instance.""" WorksheetPaymentContractAdvance.create( contract_advance_id=1, payment_name='test payment', amount=Decimal('100.00'), currency_code='USD', amount_payee_currency=Decimal('100.00'), payee_currency_code='USD', exchange_rate=Decimal('0.79'), statement_period_id=1, exchange_rate_statement_period_id=1, withholding_tax_amount=Decimal('-10.00'), vat_amount=Decimal('20.00'), amount_after_withholding_and_vat=Decimal('110.00'), withholding_tax_amount_payee_currency=Decimal('110.00'), vat_amount_payee_currency=Decimal('20.00'), amount_after_withholding_and_vat_payee_currency=Decimal('110.00'), us_source_income_rate=Decimal('50.00'), salesforce_id='test_id', is_internal=True, ) worksheet_payment_contract_advances = ( db.session.execute(select(WorksheetPaymentContractAdvance)).scalars().all() ) assert len(worksheet_payment_contract_advances) == 1 worksheet_payment_contract_advance = worksheet_payment_contract_advances[0] assert worksheet_payment_contract_advance.payment_name == 'test payment' assert worksheet_payment_contract_advance.amount == Decimal('100.00') assert worksheet_payment_contract_advance.currency_code == 'USD' assert worksheet_payment_contract_advance.amount_payee_currency == Decimal('100.00') assert worksheet_payment_contract_advance.payee_currency_code == 'USD' assert worksheet_payment_contract_advance.exchange_rate == Decimal('0.79') assert worksheet_payment_contract_advance.statement_period_id == 1 assert worksheet_payment_contract_advance.exchange_rate_statement_period_id == 1 assert worksheet_payment_contract_advance.withholding_tax_amount == Decimal( '-10.00' ) assert worksheet_payment_contract_advance.vat_amount == Decimal('20.00') assert ( worksheet_payment_contract_advance.amount_after_withholding_and_vat == Decimal('110.00') ) assert ( worksheet_payment_contract_advance.withholding_tax_amount_payee_currency == Decimal('110.00') ) assert worksheet_payment_contract_advance.vat_amount_payee_currency == Decimal( '20.00' ) assert ( worksheet_payment_contract_advance.amount_after_withholding_and_vat_payee_currency # noqa == Decimal('110.00') ) assert worksheet_payment_contract_advance.us_source_income_rate == Decimal('50.00') assert worksheet_payment_contract_advance.salesforce_id == 'test_id' assert worksheet_payment_contract_advance.is_internal assert not worksheet_payment_contract_advance.deleted_at assert not worksheet_payment_contract_advance.deleted_by def test_get_by_id(mock_statement_periods, mock_exchange_rates, mock_contract_advance): """Test get_by_id method.""" contract_advance_id = 1 created_item = WorksheetPaymentContractAdvanceFactory.create( contract_advance_id=contract_advance_id ) found_item = WorksheetPaymentContractAdvance.get_by_id( created_item.worksheet_payment_contract_advance_id ) assert found_item == created_item created_item.update_attributes(deleted_at=datetime.now(), deleted_by='default_user') created_item.commit_changes() found_item = WorksheetPaymentContractAdvance.get_by_id( created_item.worksheet_payment_contract_advance_id ) assert found_item is None def test_default_order( mock_statement_periods, mock_exchange_rates, mock_contract_advance ): """Test default order.""" contract_advance_id = 1 created_item1 = WorksheetPaymentContractAdvanceFactory.create( contract_advance_id=contract_advance_id, created_at=datetime.now() - timedelta(days=2), ) created_item2 = WorksheetPaymentContractAdvanceFactory.create( contract_advance_id=contract_advance_id, created_at=datetime.now() - timedelta(days=1), ) items = ( db.session.execute( select(WorksheetPaymentContractAdvance).order_by( WorksheetPaymentContractAdvance.default_order() ) ) .scalars() .all() ) assert items == [created_item2, created_item1] def test_filter_active_internal( mock_statement_periods, mock_exchange_rates, mock_contract_advance ): """Test filter_active_internal.""" contract_advance_id = 1 WorksheetPaymentContractAdvanceFactory.create( contract_advance_id=contract_advance_id, deleted_at=datetime.now(), is_internal=True, ) WorksheetPaymentContractAdvanceFactory.create( contract_advance_id=contract_advance_id, is_internal=False ) expected_item = WorksheetPaymentContractAdvanceFactory.create( contract_advance_id=contract_advance_id, is_internal=True ) items = ( db.session.execute(WorksheetPaymentContractAdvance.filter_active_internal()) .scalars() .all() ) assert items == [expected_item] def test_filter_for( mock_statement_periods, mock_exchange_rates, mock_contract_advance, create_mock_payment_state, ): """Test filter_for method.""" test_salesforce_id = 'test salesforce id for search' test_contract_advance_id = 1 expected_item = WorksheetPaymentContractAdvanceFactory.create( salesforce_id=test_salesforce_id ) WorksheetPaymentContractAdvanceFactory.create() query = WorksheetPaymentContractAdvance.query query = WorksheetPaymentContractAdvance.join_states(query) items = query.filter( *WorksheetPaymentContractAdvance.filter_for( { 'salesforce_id': test_salesforce_id, 'contract_advance_id': test_contract_advance_id, 'payment_states': ['init', 'complete'], } ) ).all() assert items == [expected_item] expected_deleted_item = WorksheetPaymentContractAdvanceFactory.create( salesforce_id=test_salesforce_id, deleted_at=datetime.now() ) items = query.filter( *WorksheetPaymentContractAdvance.filter_for( { 'salesforce_id': test_salesforce_id, 'contract_advance_id': test_contract_advance_id, 'payment_states': ['init', 'complete', None], } ) ).all() assert items == [expected_item, expected_deleted_item] items = query.filter( *WorksheetPaymentContractAdvance.filter_for( { 'salesforce_id': test_salesforce_id, 'contract_advance_id': test_contract_advance_id, 'payment_states': ['init', 'complete', None], 'active': True, } ) ).all() assert items == [expected_item] def test_get_filtered_active_internal_records( mock_statement_periods, mock_exchange_rates, mock_contract_advance ): """Test get_filtered_active_internal_records method.""" contract_advance_id = 1 WorksheetPaymentContractAdvanceFactory.create( contract_advance_id=contract_advance_id, deleted_at=datetime.now(), is_internal=True, ) WorksheetPaymentContractAdvanceFactory.create( contract_advance_id=contract_advance_id, is_internal=False ) expected_item = WorksheetPaymentContractAdvanceFactory.create( contract_advance_id=contract_advance_id, is_internal=True ) items, total_count = ( WorksheetPaymentContractAdvance.get_filtered_active_internal_records( offset=0, limit=10, salesforce_id=None, contract_advance_id=None, payment_statuses=[], ) ) assert total_count == 1 assert items == [expected_item] def test_get_filtered_active_internal_records_limit_zero( mock_statement_periods, mock_exchange_rates, mock_contract_advance ): """Test get_filtered_active_internal_records returns empty list when limit is 0.""" WorksheetPaymentContractAdvanceFactory.create( contract_advance_id=1, is_internal=True ) items, total_count = ( WorksheetPaymentContractAdvance.get_filtered_active_internal_records( offset=0, limit=0, salesforce_id=None, contract_advance_id=None, payment_statuses=[], ) ) assert total_count == 1 assert items == [] def test_delete_by_id_or_error( mock_statement_periods, mock_exchange_rates, mock_contract_advance, create_mock_payment_state, ): """Test delete_by_id_or_error method.""" WorksheetPaymentContractAdvanceFactory() WorksheetPaymentContractAdvanceFactory() WorksheetPaymentContractAdvance.delete_by_id_or_error(1) with pytest.raises(Exception) as excinfo: WorksheetPaymentContractAdvance.delete_by_id_or_error(2) assert excinfo.value.code == 400 assert ( excinfo.value.description == 'The worksheet action "send_payments" should have the "init" status.' ) assert WorksheetPaymentContractAdvance.get_by_id(1) is None assert WorksheetPaymentContractAdvance.get_by_id(2)