"""WorksheetAccountContractPayableDetails repository tests.""" from decimal import Decimal from abacus_common_logic.connectors.database import db from sqlalchemy import select from payment.models import WorksheetAccountContractPayableDetails from payment.repository import worksheet_account_contract_payable_details as repository from tests.utils.factories import ( PaymentGroupPaymentAccountDetailFactory, PaymentGroupPaymentAccountFactory, PaymentGroupPaymentFactory, WorksheetAccountContractClosingBalanceFactory, WorksheetAccountContractPayableDetailsFactory, WorksheetPayableBalanceAfterTaxFactory, ) def test_get_by_statement_period_id( mock_contracts, mock_accounts, mock_abacus_event, mock_statement_periods, mock_ledger_account_contracts, ): """Get contracts by statement period id.""" statement_period_id = 1 created_item = WorksheetAccountContractPayableDetails.create( statement_period_id=statement_period_id, contract_id=1, account_id=1, target_table='test_table', target_id=1, payable_detail_type_id=1, amount_payable=Decimal('10.00'), currency='USD', ) created_item2 = WorksheetAccountContractPayableDetails.create( statement_period_id=statement_period_id, contract_id=1, account_id=1, target_table='test_table', target_id=1, payable_detail_type_id=1, amount_payable=Decimal('10.00'), currency='USD', ) created_item3 = WorksheetAccountContractPayableDetails.create( statement_period_id=3, contract_id=1, account_id=1, target_table='test_table', target_id=1, payable_detail_type_id=1, amount_payable=Decimal('10.00'), currency='USD', ) found_items, total_count = repository.get_filtered_active_records( statement_period_id=1 ) assert len(found_items) == 2 assert total_count == 2 assert created_item in found_items assert created_item2 in found_items assert created_item3 not in found_items assert found_items[0].statement_period_id == statement_period_id assert found_items[1].statement_period_id == statement_period_id found_items, total_count = repository.get_filtered_active_records( limit=1, offset=1, statement_period_id=1 ) assert len(found_items) == 1 assert total_count == 2 assert created_item2 in found_items assert created_item not in found_items assert created_item3 not in found_items assert found_items[0].statement_period_id == statement_period_id found_items, total_count = repository.get_filtered_active_records( limit=1, offset=2, statement_period_id=1, ) assert len(found_items) == 0 assert total_count == 2 found_items, total_count = repository.get_filtered_active_records( limit=10, statement_period_id=3, ) assert len(found_items) == 1 assert total_count == 1 assert created_item3 in found_items assert created_item not in found_items assert created_item2 not in found_items assert found_items[0].statement_period_id == 3 def test_get_by_worksheet_after_tax_ids( mock_contracts, mock_accounts, mock_abacus_event, mock_statement_periods, mock_ledger_account_contracts, mock_worksheet_account_contract_payable_after_tax, ): """Get contracts by worksheet account contract payable after tax ids.""" created_item = WorksheetAccountContractPayableDetails.create( statement_period_id=1, worksheet_account_contract_payable_after_tax_id=1, contract_id=1, account_id=1, target_table='test_table', target_id=1, payable_detail_type_id=1, amount_payable=Decimal('10.00'), currency='USD', ) created_item2 = WorksheetAccountContractPayableDetails.create( statement_period_id=1, worksheet_account_contract_payable_after_tax_id=2, contract_id=1, account_id=1, target_table='test_table', target_id=1, payable_detail_type_id=1, amount_payable=Decimal('10.00'), currency='USD', ) created_item3 = WorksheetAccountContractPayableDetails.create( statement_period_id=3, contract_id=1, account_id=1, target_table='test_table', target_id=1, payable_detail_type_id=1, amount_payable=Decimal('10.00'), currency='USD', ) found_items, total_count = repository.get_filtered_active_records( # noqa statement_period_id=1, worksheet_after_tax_ids=[1] ) assert len(found_items) == 1 assert total_count == 1 assert created_item in found_items assert created_item2 not in found_items assert found_items[0].worksheet_account_contract_payable_after_tax_id == 1 found_items, total_count = repository.get_filtered_active_records( statement_period_id=1, worksheet_after_tax_ids=[1, 2] ) assert len(found_items) == 2 assert total_count == 2 assert created_item2 in found_items assert created_item in found_items assert created_item3 not in found_items assert found_items[0].worksheet_account_contract_payable_after_tax_id == 1 assert found_items[1].worksheet_account_contract_payable_after_tax_id == 2 def test_get_by_detail_group( mock_contracts, mock_accounts, mock_abacus_event, mock_statement_periods, mock_ledger_account_contracts, mock_worksheet_account_contract_payable_after_tax, ): """Get contracts by detail group names.""" created_item = WorksheetAccountContractPayableDetails.create( statement_period_id=1, worksheet_account_contract_payable_after_tax_id=1, contract_id=1, account_id=1, target_table='test_table', target_id=1, payable_detail_type_id=3, amount_payable=Decimal('10.00'), currency='USD', ) created_item2 = WorksheetAccountContractPayableDetails.create( statement_period_id=1, worksheet_account_contract_payable_after_tax_id=2, contract_id=1, account_id=1, target_table='test_table', target_id=1, payable_detail_type_id=1, amount_payable=Decimal('10.00'), currency='USD', ) created_item3 = WorksheetAccountContractPayableDetails.create( statement_period_id=3, contract_id=1, account_id=1, target_table='test_table', target_id=1, payable_detail_type_id=2, amount_payable=Decimal('10.00'), currency='USD', ) found_items, total_count = repository.get_filtered_active_records( statement_period_id=1, detail_groups=['vat_amount'] ) assert len(found_items) == 1 assert total_count == 1 assert created_item in found_items assert created_item2 not in found_items assert found_items[0].worksheet_account_contract_payable_after_tax_id == 1 found_items, total_count = repository.get_filtered_active_records( statement_period_id=1, detail_groups=['vat_amount', 'pre_tax_amount'] ) assert len(found_items) == 2 assert total_count == 2 assert created_item2 in found_items assert created_item in found_items assert created_item3 not in found_items found_items.sort(key=lambda x: x.payable_detail_type_id) assert found_items[0].payable_detail_type_id == 1 assert found_items[1].payable_detail_type_id == 3 def test_soft_delete_by_payment_group_payment_account( mock_contracts, mock_accounts, mock_abacus_event, mock_statement_periods, mock_ledger_account_contracts, mock_worksheet_account_contract_payable_after_tax, ): """Test soft_delete_by_payment_group_payment_account.""" payment_group_payment = PaymentGroupPaymentFactory.create( payment_group_payment_id=1 ) payment_group_payment_account = PaymentGroupPaymentAccountFactory.create( payment_group_payment=payment_group_payment ) worksheet_closing_balance = WorksheetAccountContractClosingBalanceFactory.create() worksheet_after_tax1 = WorksheetPayableBalanceAfterTaxFactory.create( worksheet_account_contract_closing_balance_id=worksheet_closing_balance.worksheet_account_contract_closing_balance_id, abacus_event_id=1, ) worksheet_after_tax2 = WorksheetPayableBalanceAfterTaxFactory.create( worksheet_account_contract_closing_balance_id=worksheet_closing_balance.worksheet_account_contract_closing_balance_id, abacus_event_id=2, ) PaymentGroupPaymentAccountDetailFactory.create( worksheet_account_contract_payable_after_tax_id=worksheet_after_tax1.worksheet_account_contract_payable_after_tax_id, payment_group_payment_account=payment_group_payment_account, ) PaymentGroupPaymentAccountDetailFactory.create( worksheet_account_contract_payable_after_tax_id=worksheet_after_tax2.worksheet_account_contract_payable_after_tax_id ) payable_details1 = WorksheetAccountContractPayableDetailsFactory.create( worksheet_account_contract_payable_after_tax=worksheet_after_tax1 ) WorksheetAccountContractPayableDetailsFactory.create( worksheet_account_contract_payable_after_tax=worksheet_after_tax2 ) repository.soft_delete_by_payment_group_payment_account( payment_group_payment_account.payment_group_payment_account_id ) assert db.session.execute( select(WorksheetAccountContractPayableDetails).where( WorksheetAccountContractPayableDetails.deleted_at != None, # noqa ) ).scalars().all() == [payable_details1] def test_soft_delete_by_payment_group_payment( mock_contracts, mock_accounts, mock_abacus_event, mock_statement_periods, mock_ledger_account_contracts, mock_worksheet_account_contract_payable_after_tax, ): """Test soft_delete_by_payment_group_payment.""" payment_group_payment = PaymentGroupPaymentFactory.create( payment_group_payment_id=1 ) payment_group_payment_account = PaymentGroupPaymentAccountFactory.create( payment_group_payment=payment_group_payment ) worksheet_closing_balance = WorksheetAccountContractClosingBalanceFactory.create() worksheet_after_tax1 = WorksheetPayableBalanceAfterTaxFactory.create( worksheet_account_contract_closing_balance_id=worksheet_closing_balance.worksheet_account_contract_closing_balance_id, abacus_event_id=1, ) worksheet_after_tax2 = WorksheetPayableBalanceAfterTaxFactory.create( worksheet_account_contract_closing_balance_id=worksheet_closing_balance.worksheet_account_contract_closing_balance_id, abacus_event_id=2, ) PaymentGroupPaymentAccountDetailFactory.create( worksheet_account_contract_payable_after_tax_id=worksheet_after_tax1.worksheet_account_contract_payable_after_tax_id, payment_group_payment_account=payment_group_payment_account, ) PaymentGroupPaymentAccountDetailFactory.create( worksheet_account_contract_payable_after_tax_id=worksheet_after_tax2.worksheet_account_contract_payable_after_tax_id ) payable_details1 = WorksheetAccountContractPayableDetailsFactory.create( worksheet_account_contract_payable_after_tax=worksheet_after_tax1 ) WorksheetAccountContractPayableDetailsFactory.create( worksheet_account_contract_payable_after_tax=worksheet_after_tax2 ) repository.soft_delete_by_payment_group_payment( payment_group_payment.payment_group_payment_id ) assert db.session.execute( select(WorksheetAccountContractPayableDetails).where( WorksheetAccountContractPayableDetails.deleted_at != None, # noqa ) ).scalars().all() == [payable_details1]