"""TaxWithholdingOverride model tests.""" import datetime from decimal import Decimal import pytest from payee.models.tax_withholding_override import TaxWithholdingOverride from tests.utils.factories import AccountPayeeFactory class TestTaxWithholdingOverride: """Test cases for TaxWithholdingOverride model.""" def test_create_tax_withholding_override_with_all_fields(self): """Test creating a TaxWithholdingOverride with all fields populated.""" account_payee = AccountPayeeFactory.create() rate_override = Decimal('15.25') certificate_expiration_date = datetime.date(2025, 9, 30) message = 'message' created_at = datetime.datetime(2025, 9, 18, 16, 0, 0) last_modified = datetime.datetime(2025, 9, 18, 16, 0, 0) override = TaxWithholdingOverride.create( account_payee_id=account_payee.account_payee_id, rate_override=rate_override, certificate_expiration_date=certificate_expiration_date, message=message, created_by='test_user', created_at=created_at, last_modified_by='test_user', last_modified=last_modified, ) assert override.tax_withholding_override_id is not None assert override.account_payee_id == account_payee.account_payee_id assert override.rate_override == rate_override assert override.certificate_expiration_date == certificate_expiration_date assert override.message == message assert override.created_by == 'test_user' assert override.created_at == created_at assert override.last_modified_by == 'test_user' assert override.last_modified == last_modified def test_unique_constraint_on_account_payee_id(self): """Test that account_payee_id must be unique.""" account_payee = AccountPayeeFactory.create() TaxWithholdingOverride.create( account_payee_id=account_payee.account_payee_id, created_by='test_user1', created_at=datetime.datetime.now(), last_modified_by='test_user1', last_modified=datetime.datetime.now(), ) with pytest.raises(Exception): TaxWithholdingOverride.create( account_payee_id=account_payee.account_payee_id, # duplicate account_payee_id created_by='test_user2', created_at=datetime.datetime.now(), last_modified_by='test_user2', last_modified=datetime.datetime.now(), ) def test_get_by_account_payee_id(self): """Test retrieving override by account_payee_id.""" account_payee = AccountPayeeFactory.create() override = TaxWithholdingOverride.create( account_payee_id=account_payee.account_payee_id, rate_override=Decimal('20.00'), message='Test override', created_by='test_user', created_at=datetime.datetime.now(), last_modified_by='test_user', last_modified=datetime.datetime.now(), ) found_override = TaxWithholdingOverride.get_by_account_payee_id( account_payee.account_payee_id ) assert found_override is not None assert ( found_override.tax_withholding_override_id == override.tax_withholding_override_id ) assert found_override.rate_override == Decimal('20.00') assert found_override.message == 'Test override' def test_get_by_account_payee_ids(self): """Test retrieving overrides by a list of account_payee_ids.""" num_of_overrides = 3 account_payees = [ AccountPayeeFactory.create(account_id=i + 1) for i in range(num_of_overrides) ] account_payees_ids = [ap.account_payee_id for ap in account_payees] created_overrides = [] for i, account_payee in enumerate(account_payees): override = TaxWithholdingOverride.create( account_payee_id=account_payee.account_payee_id, rate_override=Decimal('20.00'), message=f'test_message {i}', created_by='test_user', created_at=datetime.datetime.now(), last_modified_by='test_user', last_modified=datetime.datetime.now(), ) created_overrides.append(override) found_overrides = TaxWithholdingOverride.get_by_account_payee_ids( account_payees_ids ) assert len(found_overrides) == num_of_overrides found_by_account_payee_id = {o.account_payee_id: o for o in found_overrides} for i, account_payee_id in enumerate(account_payees_ids): assert account_payee_id in found_by_account_payee_id override = found_by_account_payee_id[account_payee_id] assert override.rate_override == Decimal('20.00') assert override.message == f'test_message {i}'