"""Tax withholding override logic tests.""" import datetime from decimal import Decimal from unittest.mock import Mock, patch import pytest from payee.logic import tax_withholding_override as logic from payee.schemas.tax_withholding_override import TaxWithholdingOverrideSchema from payee.utils.exception import AccountPayeeNotFoundException from tests.utils.factories import TaxWithholdingOverrideFactory as factory @patch('payee.logic.tax_withholding_override.TaxWithholdingOverride') def test_get_tax_withholding_override_by_account_payee_ids_with_overrides( tax_withholding_override_mock, ): """Test getting tax withholding overrides.""" data = [ factory.build(account_payee_id=1), factory.build(account_payee_id=3), ] tax_withholding_override_mock.get_by_account_payee_ids.return_value = data res = logic.get_tax_withholding_override_by_account_payee_ids([1, 2, 3]) tax_withholding_override_mock.get_by_account_payee_ids.assert_called_once_with( [1, 2, 3] ) assert res == { 'items': [ {'data': data[0]}, {'data': None}, {'data': data[1]}, ] } @patch('payee.logic.tax_withholding_override.AccountPayee') @patch('payee.logic.tax_withholding_override.TaxWithholdingOverride') def test_create_tax_withholding_override_success(mock_model, mock_account_payee): """Test successful creation of new tax withholding override.""" account_payee_id = 123 input_data = { 'rate_override': Decimal('15.25'), 'certificate_expiration_date': datetime.date(2025, 12, 31), 'message': 'Override for tax treaty', } mock_instance = factory.build( account_payee_id=account_payee_id, rate_override=Decimal('15.25'), certificate_expiration_date=datetime.date(2025, 12, 31), message='Override for tax treaty', ) mock_account_payee.get_payee_by_id.return_value = Mock() mock_model.get_by_account_payee_id.return_value = None mock_model.build.return_value = mock_instance result = logic.create_or_update_tax_withholding_override( account_payee_id, **input_data ) schema = TaxWithholdingOverrideSchema() dumped_result = schema.dump(result) assert dumped_result['account_payee_id'] == account_payee_id assert dumped_result['rate_override'] == '15.25' assert dumped_result['message'] == 'Override for tax treaty' @patch('payee.logic.tax_withholding_override.AccountPayee') @patch('payee.logic.tax_withholding_override.TaxWithholdingOverride') def test_update_tax_withholding_override_success(mock_model, mock_account_payee): """Test successful update of existing tax withholding override.""" account_payee_id = 456 input_data = {'rate_override': Decimal('25.00'), 'message': 'new message'} existing_instance = factory.build( account_payee_id=account_payee_id, rate_override=Decimal('10.00'), certificate_expiration_date=datetime.date(2025, 6, 30), message='old message', ) mock_account_payee.get_payee_by_id.return_value = Mock() mock_update_attributes = Mock(return_value=existing_instance) existing_instance.update_attributes = mock_update_attributes existing_instance.rate_override = input_data['rate_override'] existing_instance.message = input_data['message'] mock_model.get_by_account_payee_id.return_value = existing_instance result = logic.create_or_update_tax_withholding_override( account_payee_id, **input_data ) schema = TaxWithholdingOverrideSchema() dumped_result = schema.dump(result) mock_update_attributes.assert_called_once_with(**input_data) assert dumped_result['account_payee_id'] == account_payee_id assert dumped_result['rate_override'] == '25.00' assert dumped_result['message'] == 'new message' @patch('payee.logic.tax_withholding_override.AccountPayee') @patch('payee.logic.tax_withholding_override.TaxWithholdingOverride') def test_create_or_update_tax_withholding_override_account_payee_not_found( mock_model, mock_account_payee ): """Test that AccountPayeeNotFoundException is raised when account payee doesn't exist.""" account_payee_id = 1000 input_data = {'rate_override': Decimal('15.25'), 'message': 'dummy message'} mock_account_payee.get_payee_by_id.return_value = None with pytest.raises(AccountPayeeNotFoundException) as exc_info: logic.create_or_update_tax_withholding_override(account_payee_id, **input_data) assert str(account_payee_id) in str(exc_info.value) mock_model.get_by_account_payee_id.assert_not_called() mock_model.build.assert_not_called()