"""Post tax forms processor tests.""" import logging from unittest.mock import _Call, call, Mock, patch from faker import Faker import pytest from src.constants import AbacusActions, ActionStatuses, DEFAULT_BATCH_SIZE from src.models import ( AccountTaxInfo, AccountTaxInfoBulk, PostTaxFormsInput, TaxFormInfoBulk, TaxFormInfoItem, UpdateAccountTaxInfo, ) from src.processors.base import LogEntry from src.processors.exceptions import ProcessingError from src.processors.tax_forms.post_tax_forms import PostTaxFormsProcessor from tests.unit.factories import ( AccountTaxInfoFactory, PostTaxFormsInputFactory, TaxFormInfoItemFactory, USTaxFormW8BENEFactory, USTaxFormW8BENFactory, USTaxFormW8ECIFactory, USTaxFormW8IMYFactory, USTaxFormW9Factory, ) from tests.utils import create_sample_csv_buffer class TestPostTaxFormsProcessor: """PostTaxFormsProcessor test suite.""" w9_headers = ( 'vendor_id', 'tax_name', 'tax_classification', 'tax_residence_country', 'tin_type', 'tin', 'tax_id_country', 'tax_form_type', 'test_column', ) @patch.object(PostTaxFormsProcessor, '_load_data') @patch.object(PostTaxFormsProcessor, '_load_payee_ids') @patch.object(PostTaxFormsProcessor, '_load_account_tax_info') @patch.object(PostTaxFormsProcessor, '_load_tax_forms') @patch.object(PostTaxFormsProcessor, '_validate_data') @patch.object(PostTaxFormsProcessor, '_post_tax_forms') def test_process( self, mock_post_tax_forms: Mock, mock_validate_data: Mock, mock_load_tax_forms: Mock, mock_load_account_tax_info: Mock, mock_load_payee_ids: Mock, mock_load_data: Mock, faker: Faker, ) -> None: """Test process method.""" processor = PostTaxFormsProcessor(faker.pystr(), faker.pystr()) processor.process() mock_load_data.assert_called_once() mock_load_payee_ids.assert_called_once() mock_load_account_tax_info.assert_called_once() mock_load_tax_forms.assert_called_once() mock_validate_data.assert_called_once() mock_post_tax_forms.assert_called_once() def test_load_data_success(self, faker: Faker) -> None: """Test loading data from CSV file success.""" file_content = [ tuple(list(self.w9_headers) + ['override']), ( '783163', 'Grupo Deer LLC', 'Individual / Sole Proprietor', 'AUS', 'US TIN', '82-0727941', 'USA', 'W9', 'a1', 'TRUE', ), ( '783895', 'C4 Trio LLC', 'LLC-P', 'CAN', 'US TIN', '82-3831537', 'USA', 'W9', 'b2', '', ), ] processor = PostTaxFormsProcessor(faker.pystr(), faker.pystr()) processor._input_file_buffer = create_sample_csv_buffer(file_content) processor._load_data() assert processor._data == { int(row[0]): PostTaxFormsInput.model_validate( {k: v for k, v in zip(file_content[0], row)} ) for row in file_content[1:] } def test_load_data_failure_no_file(self, faker: Faker) -> None: """Test loading data from CSV file failure.""" processor = PostTaxFormsProcessor(faker.pystr(), faker.pystr()) with pytest.raises(ProcessingError, match='Unable to read the file'): processor._load_data() def test_load_data_failure_no_vendor_id_and_country(self, faker: Faker) -> None: """Test loading data from CSV file failure for some rows.""" processor = PostTaxFormsProcessor(faker.pystr(), faker.pystr()) processor._input_file_buffer = create_sample_csv_buffer( ( self.w9_headers, ( '', 'Grupo Deer LLC', 'Individual / Sole Proprietor', 'AUS', 'US TIN', '82-0727941', 'USA', 'W9', 'a1', ), ( '783895', 'C4 Trio LLC', 'LLC-P', '', 'US TIN', '82-3831537', 'USA', 'W9', 'b2', ), ) ) processor._logs = {} with pytest.raises(ProcessingError, match='No items in file to process'): processor._load_data() assert processor._data == {} assert len(processor._logs) == 1 assert logging.ERROR in processor._logs assert len(processor._logs[logging.ERROR]) == 2 assert 'vendor_id' in processor._logs[logging.ERROR][0].message assert 'tax_residence_country' in processor._logs[logging.ERROR][1].message def test_load_data_failure_no_type(self, faker: Faker) -> None: """Test loading data from CSV file failure.""" processor = PostTaxFormsProcessor(faker.pystr(), faker.pystr()) processor._input_file_buffer = create_sample_csv_buffer( ( self.w9_headers, ( '783895', 'C4 Trio LLC', 'LLC-P', 'CAN', 'US TIN', '82-3831537', 'USA', '', 'b2', ), ) ) processor._logs = {} with pytest.raises(ProcessingError, match='No items in file to process'): processor._load_data() assert processor._data == {} assert len(processor._logs) == 1 assert logging.ERROR in processor._logs assert len(processor._logs[logging.ERROR]) == 1 assert 'tax_form_type' in processor._logs[logging.ERROR][0].message assert 'tax_form' in processor._logs[logging.ERROR][0].message def test_load_data_failure_no_items(self, faker: Faker) -> None: """Test loading data from CSV file failure.""" processor = PostTaxFormsProcessor(faker.pystr(), faker.pystr()) processor._input_file_buffer = create_sample_csv_buffer((self.w9_headers,)) with pytest.raises(ProcessingError, match='No items in file to process'): processor._load_data() @patch('src.processors.tax_forms.post_tax_forms.ows_abacus_account') def test_load_payee_ids( self, mock_ows_abacus_account: Mock, faker: Faker, ) -> None: """Test _load_payee_ids method.""" def generate_ids_map(start: int, end: int) -> dict[int, int]: return {10 + i: 200 + i for i in range(start, end)} def generate_ids_tax_form_map( start: int, end: int ) -> dict[int, PostTaxFormsInput]: post_input = PostTaxFormsInputFactory.build( tax_form=USTaxFormW8BENEFactory.build(lob='lob1') ) return {10 + i: post_input for i in range(start, end)} items_count = round(DEFAULT_BATCH_SIZE * 1.3) start = round(DEFAULT_BATCH_SIZE / 2.5) end = items_count - 3 mock_ows_abacus_account.get_payees_by_accounts.side_effect = ( generate_ids_map(start, DEFAULT_BATCH_SIZE), generate_ids_map(DEFAULT_BATCH_SIZE, end), ) processor = PostTaxFormsProcessor(faker.pystr(), faker.pystr()) processor._data = generate_ids_tax_form_map(0, items_count) processor._load_payee_ids() assert mock_ows_abacus_account.get_payees_by_accounts.call_args_list == [ call(list(generate_ids_map(0, DEFAULT_BATCH_SIZE).keys())), call(list(generate_ids_map(DEFAULT_BATCH_SIZE, items_count).keys())), ] assert processor._account_id_account_payee_id == generate_ids_map(start, end) @patch('src.processors.tax_forms.post_tax_forms.ows_abacus_account') def test_load_account_tax_info( self, mock_ows_abacus_account: Mock, faker: Faker, ) -> None: """Test _load_account_tax_info method.""" def generate_tax_form_map(start: int, end: int) -> dict[int, PostTaxFormsInput]: post_input = PostTaxFormsInputFactory.build( tax_form=USTaxFormW8BENEFactory.build(lob='abc') ) return {30 + i: post_input for i in range(start, end)} def generate_map( start: int, end: int, id_only: bool = False ) -> dict[int, AccountTaxInfo | int]: return { 30 + i: ( 201 + i if id_only else AccountTaxInfoFactory.build( account_id=30 + i, account_tax_info_id=201 + i ) ) for i in range(start, end) } items_count = round(DEFAULT_BATCH_SIZE * 1.2) start = round(DEFAULT_BATCH_SIZE / 2.6) end = items_count - 6 mock_ows_abacus_account.get_account_tax_info_bulk.side_effect = [ AccountTaxInfoBulk( items=list(generate_map(start, DEFAULT_BATCH_SIZE, False).values()), total_count=DEFAULT_BATCH_SIZE - start + 1, ), AccountTaxInfoBulk( items=list(generate_map(DEFAULT_BATCH_SIZE, end, False).values()), total_count=end - DEFAULT_BATCH_SIZE + 1, ), ] processor = PostTaxFormsProcessor(faker.pystr(), faker.pystr()) processor._data = generate_tax_form_map(0, items_count) processor._load_account_tax_info() assert mock_ows_abacus_account.get_account_tax_info_bulk.call_args_list == [ call( list(generate_map(0, DEFAULT_BATCH_SIZE, True).keys()), limit=DEFAULT_BATCH_SIZE, ), call( list(generate_map(DEFAULT_BATCH_SIZE, items_count, True).keys()), limit=DEFAULT_BATCH_SIZE, ), ] assert processor._account_id_account_tax_info_id == generate_map( start, end, True ) @patch('src.processors.tax_forms.post_tax_forms.ows_payee') def test_load_tax_forms( self, mock_ows_payee: Mock, faker: Faker, ) -> None: """Test _load_tax_forms method.""" def generate_ids_map(start: int, end: int) -> dict[int, int]: return {10 + i: 120 + i for i in range(start, end)} def generate_tax_info_map( start: int, end: int, id_only: bool = False ) -> dict[int, TaxFormInfoItem | int]: return { 10 + i: ( 20 + i if id_only else TaxFormInfoItemFactory.build( account_payee_id=120 + i, account_payee_tax_form_info_id=20 + i, ) ) for i in range(start, end) } items_count = round(DEFAULT_BATCH_SIZE * 1.1) start = round(DEFAULT_BATCH_SIZE / 2.4) end = items_count - 12 mock_ows_payee.get_tax_form_info_bulk.side_effect = [ TaxFormInfoBulk( items=list( generate_tax_info_map(start, DEFAULT_BATCH_SIZE, False).values() ), total_count=DEFAULT_BATCH_SIZE - start + 1, ), TaxFormInfoBulk( items=list( generate_tax_info_map(DEFAULT_BATCH_SIZE, end, False).values() ), total_count=end - DEFAULT_BATCH_SIZE + 1, ), ] processor = PostTaxFormsProcessor(faker.pystr(), faker.pystr()) processor._account_id_account_payee_id = generate_ids_map(0, items_count) processor._load_tax_forms() assert mock_ows_payee.get_tax_form_info_bulk.call_args_list == [ call( list(generate_ids_map(0, DEFAULT_BATCH_SIZE).values()), limit=DEFAULT_BATCH_SIZE, ), call( list(generate_ids_map(DEFAULT_BATCH_SIZE, items_count).values()), limit=DEFAULT_BATCH_SIZE, ), ] assert ( processor._account_id_account_payee_tax_form_info_id == generate_tax_info_map(start, end, True) ) def test_validate_data(self, faker: Faker) -> None: """Test _validate_data method.""" processor = PostTaxFormsProcessor(faker.pystr(), faker.pystr()) processor._data = { 11: PostTaxFormsInputFactory.build( tax_form=USTaxFormW8BENEFactory.build(lob='abc'), override=None ), 14: PostTaxFormsInputFactory.build(tax_form=USTaxFormW9Factory.build()), 18: PostTaxFormsInputFactory.build(tax_form=USTaxFormW9Factory.build()), 20: PostTaxFormsInputFactory.build(tax_form=USTaxFormW8BENFactory.build()), 21: PostTaxFormsInputFactory.build(tax_form=USTaxFormW8BENFactory.build()), 23: PostTaxFormsInputFactory.build( tax_form=USTaxFormW8ECIFactory.build(lob='abc'), override=True ), } processor._account_id_account_payee_id = { 11: 101, 14: 104, 18: 108, 20: 110, 21: 111, 23: 113, } processor._account_id_account_tax_info_id = {11: 301, 14: 304, 18: 308, 23: 313} processor._account_id_account_payee_tax_form_info_id = {11: 501, 23: 513} expected_data = { i: processor._data[i].model_copy( update=dict( account_payee_id=processor._account_id_account_payee_id[i], account_tax_info_id=processor._account_id_account_tax_info_id[i], account_payee_tax_form_info_id=processor._account_id_account_payee_tax_form_info_id.get( i ), ) ) for i in (14, 18, 23) } processor._validate_data() assert processor._data == expected_data assert processor._logs == { logging.INFO: [LogEntry('Tax forms exist for account 23. Override')], logging.WARNING: [LogEntry('Tax forms exist for account 11. Skip')], logging.ERROR: [ LogEntry('Missing account_tax_info for account 20'), LogEntry('Missing account_tax_info for account 21'), ], } @patch('src.processors.tax_forms.post_tax_forms.ows_abacus_account') @patch('src.processors.tax_forms.post_tax_forms.ows_abacus_state') @patch('src.processors.tax_forms.post_tax_forms.ows_payee') def test_post_tax_forms( self, mock_ows_payee: Mock, mock_ows_abacus_state: Mock, mock_ows_abacus_account: Mock, faker: Faker, ) -> None: """Test _post_tax_forms method.""" items_count = 21 half_count = round(items_count / 2) factories = ( USTaxFormW8BENEFactory, USTaxFormW8BENFactory, USTaxFormW8ECIFactory, USTaxFormW8IMYFactory, USTaxFormW9Factory, ) class LogicItem: def __init__(self, function_name: str) -> None: self.function_name: str = function_name self.results: list[ValueError | None] = [] self.calls: list[_Call] = [] class CheckLogic: def __init__(self) -> None: self.__items: dict[str, LogicItem] = {} self.__call_next: bool = True self.__previous_item: LogicItem | None = None self.__errors: list[LogEntry] = [] self.__error_prefix: str | None = None def __get_item(self, function_name: str) -> LogicItem: if function_name not in self.__items: self.__items[function_name] = LogicItem(function_name) return self.__items[function_name] def set_error_prefix(self, error_prefix: str) -> None: self.__error_prefix = error_prefix def next_iteration(self) -> None: if not self.__call_next and self.__previous_item: self.__errors.append( LogEntry( f'{self.__error_prefix} {self.__previous_item.function_name}' ) ) self.__call_next = True self.__previous_item = None def add( self, function_name: str, call_args: _Call, no_errors: bool, is_called: bool = True, ) -> None: if not self.__call_next: if self.__previous_item: self.__errors.append( LogEntry( f'{self.__error_prefix} {self.__previous_item.function_name}' ) ) self.__previous_item = None return item = self.__get_item(function_name) if is_called: item.results.append( None if no_errors else ValueError(function_name) ) item.calls.append(call_args) self.__call_next = no_errors self.__previous_item = item def get_calls(self, function_name: str) -> list[_Call]: return self.__get_item(function_name).calls def get_results(self, function_name: str) -> list[ValueError | None]: return self.__get_item(function_name).results def get_errors(self) -> list[LogEntry]: return self.__errors data: dict[int, PostTaxFormsInput] = {} check_info = CheckLogic() for index in range(items_count): account_id = 25 + index account_payee_id = 205 + index account_tax_info_id = 600 + index account_payee_tax_form_info_id = 1010 + index country_code = f'c{index}' tax_form = factories[index % len(factories)].build(lob=f'lob {index}') save_tax_form_info_ok = bool(index % half_count) set_state_ok = bool((index + 2) % half_count) update_account_tax_info_ok = bool((index + 3) % half_count) tax_form_exists = bool(index % 2) delete_tax_form_info_ok = bool(tax_form_exists and index % 3) post_input = PostTaxFormsInputFactory.build( tax_form=tax_form, country_of_tax_residence=country_code, account_payee_id=account_payee_id, account_tax_info_id=account_tax_info_id, account_payee_tax_form_info_id=account_payee_tax_form_info_id if tax_form_exists else None, override=tax_form_exists, ) data[account_id] = post_input check_info.set_error_prefix(f'{account_id}/{account_payee_id}') check_info.add( 'delete_tax_form_info', call(account_payee_tax_form_info_id), delete_tax_form_info_ok, tax_form_exists, ) check_info.add( 'save_tax_form_info', call(account_payee_id, tax_form), save_tax_form_info_ok, ) check_info.add( 'set_state', call( account_payee_id, AbacusActions.tax_eligibility, ActionStatuses.complete, 'Upload tax form from S3 file', ), set_state_ok, ) check_info.add( 'update_account_tax_info', call( account_tax_info_id, UpdateAccountTaxInfo( country_of_tax_residence=country_code, is_tax_treaty_claimed=getattr( tax_form, 'tax_treaty_claim', None ), ), ), update_account_tax_info_ok, ) check_info.next_iteration() mock_ows_payee.delete_tax_form_info.side_effect = check_info.get_results( 'delete_tax_form_info' ) mock_ows_payee.save_tax_form_info.side_effect = check_info.get_results( 'save_tax_form_info' ) mock_ows_abacus_state.set_state.side_effect = check_info.get_results( 'set_state' ) mock_ows_abacus_account.update_account_tax_info.side_effect = ( check_info.get_results('update_account_tax_info') ) processor = PostTaxFormsProcessor(faker.pystr(), faker.pystr()) processor._data = data processor._post_tax_forms() assert ( mock_ows_payee.delete_tax_form_info.call_args_list == check_info.get_calls('delete_tax_form_info') ) assert mock_ows_payee.save_tax_form_info.call_args_list == check_info.get_calls( 'save_tax_form_info' ) assert mock_ows_abacus_state.set_state.call_args_list == check_info.get_calls( 'set_state' ) assert ( mock_ows_abacus_account.update_account_tax_info.call_args_list == check_info.get_calls('update_account_tax_info') ) assert processor._logs == {logging.ERROR: check_info.get_errors()}