"""Tests for account_tax_info model.""" from datetime import timedelta from decimal import Decimal from abacus_account.constants.constants import TAX_EMPLOYMENT_TYPES from abacus_account.models.account_tax_info import AccountTaxInfo from tests.utils.factories import AccountFactory from tests.utils.factories import AccountTaxInfoFactory def test_create_account_tax(faker): """Test creating an account_tax_info.""" account = AccountFactory.create() test_date = faker.future_date() AccountTaxInfo.create( account_id=account.account_id, country_of_tax_residence='GBR', tax_employment_type=TAX_EMPLOYMENT_TYPES[0], certificate_of_residence_expiration_date=test_date, is_wht_applicable=True, is_resident_of_spanish_islands=True, wht_rate_override=Decimal('10.01') ) result = AccountTaxInfo.query.all() assert len(result) == 1 assert result[0].account_tax_info_id assert result[0].account_id == account.account_id assert result[0].is_sba_signed is False assert result[0].is_vat_exempt is True assert result[0].is_tax_treaty_claimed is False assert result[0].tax_employment_type == TAX_EMPLOYMENT_TYPES[0] assert result[0].certificate_of_residence_expiration_date == test_date assert result[0].is_wht_applicable is True assert result[0].is_resident_of_spanish_islands is True assert result[0].wht_rate_override == Decimal('10.01') def test_get_account_tax_info_by_id(): """Test getting an account_tax_info by account_tax_info_id.""" object_id = 123 account_tax_info = AccountTaxInfoFactory.create(account_tax_info_id=object_id) result = AccountTaxInfo.get_by_id(object_id) assert result.account_tax_info_id == account_tax_info.account_tax_info_id assert result.account_id == account_tax_info.account_id assert result.is_vat_exempt == account_tax_info.is_vat_exempt def test_stream_all(account_fixtures): """Test for stream_all method of AccountTaxInfo model.""" account_ids = (1, 2, 3, 4, 5) # created in account_fixtures for account in account_fixtures: AccountTaxInfoFactory.create(account=account) results = AccountTaxInfo.stream_all().all() assert len(list(results)) == len(account_ids) results_filtered = AccountTaxInfo.stream_all((1, 4, 2, 10)) assert len(list(results_filtered)) == 3 def test_get_filtered_items(account_fixtures, faker): """Test get_filtered_items method.""" account_ids = (1, 2, 3, 4, 5) # created in account_fixtures # noqa test_expiration_date = faker.past_date() for account in account_fixtures: AccountTaxInfoFactory.create( account=account, certificate_of_residence_expiration_date=test_expiration_date ) items, total_count = AccountTaxInfo.get_filtered_items(2, 1) assert [item.account_id for item in items] == [2, 3] assert total_count == 5 items, total_count = AccountTaxInfo.get_filtered_items(10, 0, [3, 4]) assert [item.account_id for item in items] == [3, 4] assert total_count == 2 items, total_count = AccountTaxInfo.get_filtered_items( 10, 0, [3, 4], certificate_of_residence_expiration_date_start=test_expiration_date, certificate_of_residence_expiration_date_end=test_expiration_date ) assert [item.account_id for item in items] == [3, 4] assert total_count == 2 items, total_count = AccountTaxInfo.get_filtered_items( 10, 0, [3, 4], certificate_of_residence_expiration_date_start=test_expiration_date - timedelta(days=1), # noqa: E501 certificate_of_residence_expiration_date_end=test_expiration_date + timedelta(days=1) # noqa: E501 ) assert [item.account_id for item in items] == [3, 4] assert total_count == 2 items, total_count = AccountTaxInfo.get_filtered_items( 10, 0, [3, 4], certificate_of_residence_expiration_date_start=test_expiration_date + timedelta(days=1), # noqa: E501 certificate_of_residence_expiration_date_end=test_expiration_date + timedelta(days=2) # noqa: E501 ) assert items == [] assert total_count == 0 items, total_count = AccountTaxInfo.get_filtered_items( 10, 0, [3, 4], certificate_of_residence_expiration_date_start=test_expiration_date - timedelta(days=1), # noqa: E501 certificate_of_residence_expiration_date_end=test_expiration_date - timedelta(days=2) # noqa: E501 ) assert items == [] assert total_count == 0