"""Unit tests for WorksheetPaymentContractAdvance logic.""" from decimal import Decimal from unittest.mock import call, MagicMock, patch import pytest from payment.logic.exceptions import EntityDoesNotExist, LogicError from payment.logic.worksheet_payment_contract_advance import ( _round_currency_conversion, _validate_params, create_worksheet_payment_contract_advance, dataload_worksheets_by_ids, delete_worksheet_payment_contract_advance, get_worksheet_payment_contract_advance, list_worksheet_payment_contract_advances, update_worksheet_payment_contract_advance, ) from payment.models.generic import Items from tests.utils.factories import ( ExchangeRateFactory, WorksheetPaymentContractAdvanceFactory, ) def test_validate_params_success(): """Test success execution for _validate_params.""" exchange_rate = ExchangeRateFactory.create() try: _validate_params(1, 'USD', 'GBP', exchange_rate) except Exception as e: pytest.fail(f'Unexpected error: {e}') try: _validate_params(1, 'USD', 'USD') except Exception as e: pytest.fail(f'Unexpected error: {e}') @patch( 'payment.logic.worksheet_payment_contract_advance.WorksheetPaymentContractAdvance' ) def test_validate_params_failure_invalid_currency( mock_worksheet_model, mock_statement_periods, mock_exchange_rates, mock_contract_advance, ): """Test failure execution for _validate_params.""" exchange_rate = ExchangeRateFactory.create( from_currency_code='GBP', to_currency_code='USD' ) with pytest.raises(LogicError): _validate_params(1, '___', 'USD') with pytest.raises(LogicError): _validate_params(1, 'USD', '___') with pytest.raises(LogicError): _validate_params(1, 'USD', 'USD', exchange_rate) with pytest.raises(LogicError): _validate_params(1, 'GBP', 'GBP', exchange_rate) @patch( 'payment.logic.worksheet_payment_contract_advance.WorksheetPaymentContractAdvance' ) def test_validate_params_failure_existing_worksheet( mock_worksheet_model, mock_statement_periods, mock_exchange_rates, mock_contract_advance, ): """Test failure when an active worksheet already exists for the contract advance.""" exchange_rate = ExchangeRateFactory.create() worksheet = WorksheetPaymentContractAdvanceFactory.create() mock_worksheet_model.get_filtered_query.return_value.first.return_value = worksheet with pytest.raises(LogicError): _validate_params(worksheet.contract_advance_id, 'GBP', 'USD', exchange_rate) def test_round_currency_conversion(): """Test _round_currency_conversion logic.""" assert _round_currency_conversion(Decimal('100.0001')) == Decimal('100.01') @patch( 'payment.logic.worksheet_payment_contract_advance.WorksheetPaymentContractAdvance' ) @patch('payment.logic.worksheet_payment_contract_advance.ExchangeRate') @patch('payment.logic.worksheet_payment_contract_advance._validate_params') def test_create_worksheet_payment_contract_advance_success( mock_validate_params, mock_exchange_rate_model, mock_worksheet_model, mock_statement_periods, mock_exchange_rates, mock_contract_advance, ): """Test success execution for create_worksheet_payment_contract_advance.""" exchange_rate = ExchangeRateFactory.create() worksheet = WorksheetPaymentContractAdvanceFactory.create() mock_exchange_rate_model.get_by_id_or_error.return_value = exchange_rate mock_worksheet_model.create.return_value = worksheet result = create_worksheet_payment_contract_advance( contract_advance_id=1, exchange_rate_id=exchange_rate.exchange_rate_id, statement_period_id=1, payment_name='test_name', amount=Decimal('100.00'), currency_code=exchange_rate.from_currency_code, payee_currency_code=exchange_rate.to_currency_code, withholding_tax_amount=Decimal('-10.00'), vat_amount=Decimal('20.00'), amount_after_withholding_and_vat=Decimal('110.00'), us_source_income_rate=Decimal('90.00'), is_internal=False, ) rate = exchange_rate.rate amount_payee_currency = _round_currency_conversion(Decimal('100.00') * rate) withholding_tax_amount_payee_currency = _round_currency_conversion( Decimal('-10.00') * rate ) vat_amount_payee_currency = _round_currency_conversion(Decimal('20.00') * rate) amount_after_withholding_and_vat_payee_currency = ( amount_payee_currency + withholding_tax_amount_payee_currency + vat_amount_payee_currency ) assert mock_validate_params.call_args_list == [ call( 1, exchange_rate.from_currency_code, exchange_rate.to_currency_code, exchange_rate, ) ] assert mock_worksheet_model.create.call_args_list == [ call( contract_advance_id=1, statement_period_id=1, exchange_rate_statement_period_id=exchange_rate.statement_period_id, payment_name='test_name', amount=Decimal('100.00'), currency_code=exchange_rate.from_currency_code, amount_payee_currency=amount_payee_currency, payee_currency_code=exchange_rate.to_currency_code, exchange_rate=exchange_rate.rate, withholding_tax_amount=Decimal('-10.00'), vat_amount=Decimal('20.00'), amount_after_withholding_and_vat=Decimal('110.00'), withholding_tax_amount_payee_currency=withholding_tax_amount_payee_currency, # noqa vat_amount_payee_currency=vat_amount_payee_currency, amount_after_withholding_and_vat_payee_currency=amount_after_withholding_and_vat_payee_currency, # noqa us_source_income_rate=Decimal('90.00'), is_internal=False, ) ] assert result == worksheet @patch('payment.logic.worksheet_payment_contract_advance.ExchangeRate') def test_create_worksheet_payment_contract_advance_failure_no_exchange_rate( mock_exchange_rate_model, mock_statement_periods, mock_exchange_rates, mock_contract_advance, ): """Test no exchange rate error for create_worksheet_payment_contract_advance.""" mock_exchange_rate_model.get_by_id_or_error.side_effect = Exception with pytest.raises(LogicError, match='Existing exchange rate is required'): create_worksheet_payment_contract_advance( contract_advance_id=1, exchange_rate_id=1, statement_period_id=1, payment_name='test_name', amount=Decimal('100.00'), currency_code='USD', payee_currency_code='GBP', withholding_tax_amount=Decimal('-10.00'), vat_amount=Decimal('20.00'), amount_after_withholding_and_vat=Decimal('110.00'), us_source_income_rate=Decimal('90.00'), is_internal=False, ) @patch('payment.logic.worksheet_payment_contract_advance.ExchangeRate') @patch('payment.logic.worksheet_payment_contract_advance._validate_params') def test_create_worksheet_payment_contract_advance_failure_invalid_params( mock_validate_params, mock_exchange_rate_model, mock_statement_periods, mock_exchange_rates, mock_contract_advance, ): """Test invalid params for create_worksheet_payment_contract_advance.""" exchange_rate = ExchangeRateFactory.create() mock_exchange_rate_model.get_by_id_or_error.return_value = exchange_rate mock_validate_params.side_effect = LogicError('mocked error') with pytest.raises(LogicError, match='mocked error'): create_worksheet_payment_contract_advance( contract_advance_id=1, exchange_rate_id=1, statement_period_id=1, payment_name='test_name', amount=Decimal('100.00'), currency_code='USD', payee_currency_code='GBP', withholding_tax_amount=Decimal('-10.00'), vat_amount=Decimal('20.00'), amount_after_withholding_and_vat=Decimal('110.00'), us_source_income_rate=Decimal('90.00'), is_internal=False, ) @patch( 'payment.logic.worksheet_payment_contract_advance.WorksheetPaymentContractAdvance' ) def test_get_worksheet_payment_contract_advance_success(mock_model): """Test get_worksheet_payment_contract_advance logic.""" worksheet = WorksheetPaymentContractAdvanceFactory.build() mock_model.get_by_id.return_value = worksheet result = get_worksheet_payment_contract_advance(1) mock_model.get_by_id.assert_called_once_with(1) assert result == worksheet @patch( 'payment.logic.worksheet_payment_contract_advance.WorksheetPaymentContractAdvance' ) def test_get_worksheet_payment_contract_advance_not_found(mock_model): """Test get_worksheet_payment_contract_advance raises EntityDoesNotExist when not found.""" mock_model.get_by_id.return_value = None with pytest.raises(EntityDoesNotExist): get_worksheet_payment_contract_advance(999) @patch( 'payment.logic.worksheet_payment_contract_advance.WorksheetPaymentContractAdvance' ) def test_update_worksheet_payment_contract_advance_success( mock_model, mock_statement_periods, mock_exchange_rates, mock_contract_advance, ): """Test update_worksheet_payment_contract_advance logic.""" test_salesforce_id = 'new test salesforce id' mock_obj = MagicMock() mock_model.get_by_id.return_value = mock_obj result = update_worksheet_payment_contract_advance( 1, salesforce_id=test_salesforce_id ) mock_model.get_by_id.assert_called_once_with(1) assert mock_obj.update_attributes.call_args_list == [ call(salesforce_id=test_salesforce_id) ] assert mock_model.commit_changes.called assert result == mock_obj @patch( 'payment.logic.worksheet_payment_contract_advance.WorksheetPaymentContractAdvance' ) def test_update_worksheet_payment_contract_advance_not_found(mock_model): """Test update_worksheet_payment_contract_advance raises EntityDoesNotExist when not found.""" mock_model.get_by_id.return_value = None with pytest.raises(EntityDoesNotExist): update_worksheet_payment_contract_advance(999, salesforce_id='some_id') @patch( 'payment.logic.worksheet_payment_contract_advance.WorksheetPaymentContractAdvance' ) def test_delete_worksheet_payment_contract_advance_success(mock_model): """Test delete_worksheet_payment_contract_advance logic.""" delete_worksheet_payment_contract_advance(1) mock_model.delete_by_id_or_error.assert_called_once_with(1) @patch( 'payment.logic.worksheet_payment_contract_advance.WorksheetPaymentContractAdvance' ) def test_list_worksheet_payment_contract_advances_success(mock_model): """Test list_worksheet_payment_contract_advances logic.""" worksheet = WorksheetPaymentContractAdvanceFactory.build() mock_model.get_filtered_active_internal_records.return_value = ([worksheet], 1) result = list_worksheet_payment_contract_advances( offset=0, limit=10, salesforce_id=None, contract_advance_id=None, payment_statuses=[], ) assert isinstance(result, Items) assert result.items == [worksheet] assert result.total_count == 1 mock_model.get_filtered_active_internal_records.assert_called_once_with( offset=0, limit=10, salesforce_id=None, contract_advance_id=None, payment_statuses=[], ) @patch( 'payment.logic.worksheet_payment_contract_advance.WorksheetPaymentContractAdvance' ) def test_list_worksheet_payment_contract_advances_with_payment_statuses(mock_model): """Test list_worksheet_payment_contract_advances passes payment_statuses to model.""" mock_model.get_filtered_active_internal_records.return_value = ([], 0) list_worksheet_payment_contract_advances( offset=0, limit=10, salesforce_id=None, contract_advance_id=None, payment_statuses=['init'], ) mock_model.get_filtered_active_internal_records.assert_called_once_with( offset=0, limit=10, salesforce_id=None, contract_advance_id=None, payment_statuses=['init'], ) @patch( 'payment.logic.worksheet_payment_contract_advance.WorksheetPaymentContractAdvance' ) def test_dataload_worksheets_by_ids( mock_model, mock_statement_periods, mock_exchange_rates, mock_contract_advance ): """Test dataload_worksheets_by_ids function.""" worksheet = WorksheetPaymentContractAdvanceFactory.build() worksheet_ids = [worksheet.worksheet_payment_contract_advance_id, 999] mock_model.get_filtered_all.return_value = [worksheet] result = dataload_worksheets_by_ids(worksheet_ids) assert result == {'items': [{'data': worksheet}, {'data': None}]} mock_model.get_filtered_all.assert_called_once_with( worksheet_ids=worksheet_ids, active=True )