"""Unit Tests for contract_advance logic.""" from datetime import date from unittest.mock import patch import pytest from abacus_contract.constants import error from abacus_contract.constants.constants import ( ADVANCE_STATUSES, ADVANCE_STATUSES_EXTRA_VALIDATION_REQUIRED ) from abacus_contract.logic import contract_advance as logic from tests.utils.factories import ContractAdvanceFactory from tests.utils.factories import ContractFactory from tests.utils.factories import ReferencePaymentTypeFactory @patch('abacus_contract.logic.contract_advance._calculate_contract_advance_fields') @patch('abacus_contract.logic.contract_advance.models') def test_create_contract_advance(mock_models, mock_calculate_contract_advance_fields): """Test create_contract_advance function.""" reference_payment_type = ReferencePaymentTypeFactory.create() contract_advance = ContractAdvanceFactory.create( reference_payment_type=reference_payment_type, amount_after_withholding_and_vat=None ) mock_models.ContractAdvance.build.return_value = contract_advance res = logic.create_contract_advance( contract_advance.contract_id, contract_advance.advance_description, contract_advance.amount, contract_advance.currency_code, contract_advance.milestone, contract_advance.milestone_date, contract_advance.milestone_description, contract_advance.reference_payment_type_id, contract_advance.advance_status, contract_advance.note, contract_advance.vat_amount, contract_advance.withholding_tax_amount, contract_advance.us_source_income_rate ) assert res.status == 201 mock_models.ContractAdvance.build.assert_called_once_with( contract_id=contract_advance.contract_id, advance_description=contract_advance.advance_description, amount=contract_advance.amount, currency_code=contract_advance.currency_code, milestone=contract_advance.milestone, milestone_date=contract_advance.milestone_date, milestone_description=contract_advance.milestone_description, advance_status=contract_advance.advance_status, note=contract_advance.note, reference_payment_type_id=contract_advance.reference_payment_type_id, # noqa vat_amount=contract_advance.vat_amount, withholding_tax_amount=contract_advance.withholding_tax_amount, us_source_income_rate=contract_advance.us_source_income_rate ) mock_calculate_contract_advance_fields.assert_called_once_with(contract_advance) mock_models.ContractAdvance.commit_changes.assert_called_once() @patch('abacus_contract.logic.contract_advance.models') def test_create_contract_advance_failure_no_milestone(mock_models): """Test create_contract advance failure. no milestone_date for qualifying advances """ reference_payment_type = ReferencePaymentTypeFactory.create() contract_advance = ContractAdvanceFactory.create( advance_status=ADVANCE_STATUSES.QUALIFIED, milestone_date=None, reference_payment_type=reference_payment_type ) mock_models.ContractAdvance.create.return_value = contract_advance res = logic.create_contract_advance( contract_advance.contract_id, contract_advance.advance_description, contract_advance.amount, contract_advance.currency_code, contract_advance.milestone, contract_advance.milestone_date, contract_advance.milestone_description, contract_advance.reference_payment_type_id, contract_advance.advance_status, contract_advance.note, contract_advance.vat_amount, contract_advance.withholding_tax_amount, contract_advance.us_source_income_rate ) assert res.status == 400 assert res.errors['message'] == \ error.ERROR_CONTRACT_ADVANCE_NO_MILESTONE_DATE.format( status=ADVANCE_STATUSES.QUALIFIED ) @patch('abacus_contract.logic.contract_advance.models') @patch('abacus_contract.logic.contract_advance.date') def test_create_contract_advance_failure_incorrect_milestone( mock_date, mock_models ): """Test create_contract advance failure. milestone_date is in the future """ reference_payment_type = ReferencePaymentTypeFactory.create() contract_advance = ContractAdvanceFactory.create( advance_status=ADVANCE_STATUSES.QUALIFIED, milestone_date='2022-09-22', reference_payment_type=reference_payment_type ) mock_models.ContractAdvance.create.return_value = contract_advance mock_date.today.return_value = date(2022, 9, 21) res = logic.create_contract_advance( contract_advance.contract_id, contract_advance.advance_description, contract_advance.amount, contract_advance.currency_code, contract_advance.milestone, contract_advance.milestone_date, contract_advance.milestone_description, contract_advance.reference_payment_type_id, contract_advance.advance_status, contract_advance.note, contract_advance.vat_amount, contract_advance.withholding_tax_amount, contract_advance.us_source_income_rate ) assert res.status == 400 assert res.errors['message'] == \ error.ERROR_CONTRACT_ADVANCE_MILESTONE_DATE @patch('abacus_contract.logic.contract_advance._calculate_contract_advance_fields') @patch('abacus_contract.logic.contract_advance.models') def test_create_contract_advance_success_note_non_paid_advance( mock_models, mock_calculate_contract_advance_fields ): """Test create_contract advance failure. note added when advance_status is not paid """ reference_payment_type = ReferencePaymentTypeFactory.create() contract_advance = ContractAdvanceFactory.create( advance_status=ADVANCE_STATUSES.QUALIFIED, note='test', reference_payment_type=reference_payment_type ) mock_models.ContractAdvance.build.return_value = contract_advance res = logic.create_contract_advance( contract_advance.contract_id, contract_advance.advance_description, contract_advance.amount, contract_advance.currency_code, contract_advance.milestone, contract_advance.milestone_date, contract_advance.milestone_description, contract_advance.reference_payment_type_id, contract_advance.advance_status, contract_advance.note, contract_advance.vat_amount, contract_advance.withholding_tax_amount, contract_advance.us_source_income_rate ) assert res.status == 201 mock_models.ContractAdvance.build.assert_called_once_with( contract_id=contract_advance.contract_id, advance_description=contract_advance.advance_description, amount=contract_advance.amount, currency_code=contract_advance.currency_code, milestone=contract_advance.milestone, milestone_date=contract_advance.milestone_date, milestone_description=contract_advance.milestone_description, advance_status=contract_advance.advance_status, note=contract_advance.note, reference_payment_type_id=contract_advance.reference_payment_type_id, # noqa vat_amount=contract_advance.vat_amount, withholding_tax_amount=contract_advance.withholding_tax_amount, us_source_income_rate=contract_advance.us_source_income_rate ) mock_calculate_contract_advance_fields.assert_called_once_with(contract_advance) mock_models.ContractAdvance.commit_changes.assert_called_once() @pytest.mark.parametrize('advance_status', ADVANCE_STATUSES_EXTRA_VALIDATION_REQUIRED) @patch('abacus_contract.logic.contract_advance.models') def test_create_contract_advance_failure_no_vat_amount( mock_models, advance_status ): """Test create_contract_advance failure with no vat_amount.""" reference_payment_type = ReferencePaymentTypeFactory.create() contract_advance = ContractAdvanceFactory.create( reference_payment_type=reference_payment_type, amount_after_withholding_and_vat=None, advance_status=advance_status ) mock_models.ContractAdvance.build.return_value = contract_advance res = logic.create_contract_advance( contract_advance.contract_id, contract_advance.advance_description, contract_advance.amount, contract_advance.currency_code, contract_advance.milestone, contract_advance.milestone_date, contract_advance.milestone_description, contract_advance.reference_payment_type_id, contract_advance.advance_status, contract_advance.note, None, contract_advance.withholding_tax_amount, contract_advance.us_source_income_rate ) assert res.status == 400 assert res.errors['message'] == error.ERROR_CONTRACT_ADVANCE_TAXES_FIELDS_REQUIRED @pytest.mark.parametrize('advance_status', ADVANCE_STATUSES_EXTRA_VALIDATION_REQUIRED) @patch('abacus_contract.logic.contract_advance.models') def test_create_contract_advance_failure_no_withholding_tax_amount( mock_models, advance_status ): """Test create_contract_advance failure with no withholding_tax_amount.""" reference_payment_type = ReferencePaymentTypeFactory.create() contract_advance = ContractAdvanceFactory.create( reference_payment_type=reference_payment_type, amount_after_withholding_and_vat=None, advance_status=advance_status ) mock_models.ContractAdvance.build.return_value = contract_advance res = logic.create_contract_advance( contract_advance.contract_id, contract_advance.advance_description, contract_advance.amount, contract_advance.currency_code, contract_advance.milestone, contract_advance.milestone_date, contract_advance.milestone_description, contract_advance.reference_payment_type_id, contract_advance.advance_status, contract_advance.note, contract_advance.vat_amount, None, contract_advance.us_source_income_rate ) assert res.status == 400 assert res.errors['message'] == error.ERROR_CONTRACT_ADVANCE_TAXES_FIELDS_REQUIRED @patch('abacus_contract.logic.contract_advance.models') def test_get_contract_advances(mock_models): """Test get_contract_advances function.""" contract = ContractFactory.create() contract_advance = ContractAdvanceFactory.create(contract=contract) contract_id = contract.contract_id status = 'pending' params = { 'limit': 10, 'offset': 0, 'reference_payment_type_id': 1 } mock_models.ContractAdvance.get_by_contract_id.return_value = \ ([contract_advance], 1) res = logic.get_contract_advances(contract_id, status, params) assert res.status == 200 mock_models.ContractAdvance \ .get_by_contract_id \ .assert_called_once_with( contract_id=1, contract_advance_status='pending', reference_payment_type_id=1, limit=10, offset=0 ) @patch('abacus_contract.logic.contract_advance._calculate_contract_advance_fields') @patch('abacus_contract.logic.contract_advance.models') def test_update_contract_advance( mock_models, mock_calculate_contract_advance_fields, mocker ): """Test update_contract_advance function.""" contract_advance = ContractAdvanceFactory.create( amount_after_withholding_and_vat=None ) mock_models.ContractAdvance.commit_changes.return_value = None update_attributes_spy = mocker.spy(contract_advance, 'update_attributes') result = logic.update_contract_advance( contract_advance, advance_status=ADVANCE_STATUSES.QUALIFIED) assert result.status == 200 update_attributes_spy.assert_called_once_with( advance_status=ADVANCE_STATUSES.QUALIFIED ) mock_calculate_contract_advance_fields.aasert_called_once_with(contract_advance) mock_models.ContractAdvance.commit_changes.assert_called_once() @pytest.mark.parametrize('advance_status', ADVANCE_STATUSES_EXTRA_VALIDATION_REQUIRED) @patch('abacus_contract.logic.contract_advance.models') def test_update_contract_advance_failure_no_vat_amount( mock_models, advance_status ): """Test update_contract_advance failure no vat_amount.""" contract_advance = ContractAdvanceFactory.create( amount_after_withholding_and_vat=None, advance_status=advance_status, vat_amount=None ) mock_models.ContractAdvance.commit_changes.return_value = None result = logic.update_contract_advance(contract_advance) assert result.status == 400 assert result.errors[ 'message' ] == error.ERROR_CONTRACT_ADVANCE_TAXES_FIELDS_REQUIRED @pytest.mark.parametrize('advance_status', ADVANCE_STATUSES_EXTRA_VALIDATION_REQUIRED) @patch('abacus_contract.logic.contract_advance.models') def test_update_contract_advance_failure_no_withholding_tax_amount( mock_models, advance_status ): """Test update_contract_advance failure no withholding_tax_amount.""" contract_advance = ContractAdvanceFactory.create( amount_after_withholding_and_vat=None, advance_status=advance_status, withholding_tax_amount=None ) mock_models.ContractAdvance.commit_changes.return_value = None result = logic.update_contract_advance(contract_advance) assert result.status == 400 assert result.errors[ 'message' ] == error.ERROR_CONTRACT_ADVANCE_TAXES_FIELDS_REQUIRED @patch('abacus_contract.logic.contract_advance.models') def test_validate_update_contract_advance_params(mock_models): """Test _validate_update_contract_advance_params.""" contract_advance = ContractAdvanceFactory.create() params = { 'advance_status': ADVANCE_STATUSES.QUALIFIED } result = logic._validate_update_contract_advance_params(contract_advance, **params) assert result def test_validate_update_contract_advance_params_failure_no_milestone_date(): """Test _validate_update_contract_advance_params. milestone not reached """ contract_advance = ContractAdvanceFactory.create(milestone_date=None) params = { 'advance_status': ADVANCE_STATUSES.PAID } with pytest.raises( Exception, match=error.ERROR_CONTRACT_ADVANCE_NO_MILESTONE_DATE.format( status=ADVANCE_STATUSES.PAID ) ): logic._validate_update_contract_advance_params(contract_advance, **params) @patch('abacus_contract.logic.contract_advance.date') def test_validate_update_contract_advance_params_failure_milestone_date( mock_date ): """Test _validate_update_contract_advance_params. invalid milestone_date """ contract_advance = ContractAdvanceFactory.create(milestone_date=None) mock_date.today.return_value = date(2022, 9, 21) params = { 'milestone_date': date(2022, 9, 22) } with pytest.raises( Exception, match=error.ERROR_CONTRACT_ADVANCE_MILESTONE_DATE ): logic._validate_update_contract_advance_params(contract_advance, **params) def test_validate_update_contract_advance_params_failure_deleted(): """Test _validate_update_contract_advance_params. object is deleted. """ contract_advance = ContractAdvanceFactory.create( advance_status=ADVANCE_STATUSES.DELETED ) params = { 'milestone_date': date(2022, 9, 22) } with pytest.raises( Exception, match=error.ERROR_CONTRACT_ADVANCE_DELETED ): logic._validate_update_contract_advance_params(contract_advance, **params) @pytest.mark.parametrize('advance_status', ADVANCE_STATUSES) def test_validate_update_contract_advance_params_success_note_with_feature_enabled( advance_status ): """Test _validate_update_contract_advance_params. Valid note if advance payment feature is enabled. """ contract_advance = ContractAdvanceFactory.create( advance_status=ADVANCE_STATUSES.QUALIFIED ) params = { 'advance_status': advance_status, 'note': 'test' } # from QUALIFIED status transition only to any other status except PENDING_PAYMENT if advance_status == ADVANCE_STATUSES.PENDING_PAYMENT: with pytest.raises( Exception, match=error.ERROR_CONTRACT_ADVANCE_IN_REVIEW_STATUS_REQUIRED ): logic._validate_update_contract_advance_params(contract_advance, **params) else: assert logic._validate_update_contract_advance_params( contract_advance, **params ) def test_validate_update_contract_advance_params_failure_status_transition(): """Test _validate_update_contract_advance_params. invalid status transition. """ contract_advance = ContractAdvanceFactory.create( advance_status=ADVANCE_STATUSES.NOT_QUALIFIED ) params = { 'advance_status': ADVANCE_STATUSES.IN_REVIEW, } with pytest.raises( Exception, match=error.ERROR_CONTRACT_ADVANCE_QUALIFIED_STATUS_REQUIRED ): logic._validate_update_contract_advance_params(contract_advance, **params) def test_validate_update_contract_advance_params_failure_status_transition_2(): """Test _validate_update_contract_advance_params. invalid status transition. """ contract_advance = ContractAdvanceFactory.create( advance_status=ADVANCE_STATUSES.NOT_QUALIFIED ) params = { 'advance_status': ADVANCE_STATUSES.PENDING_PAYMENT, } with pytest.raises( Exception, match=error.ERROR_CONTRACT_ADVANCE_IN_REVIEW_STATUS_REQUIRED ): logic._validate_update_contract_advance_params(contract_advance, **params) @patch('abacus_contract.logic.contract_advance.models') def test_get_paid_contract_advances(mock_models): """Test get_contract_advances function to get paid advances.""" mock_paid_advances = [{ 'contract_advance_id': 1, 'contract_id': 1, 'statement_period_id': 1, 'amount': '100.00', 'currency_code': 'USD', 'advance_amount_payee_currency': '100.00', 'advance_payee_currency_code': 'USD', 'milestone': 'recoupment', 'advance_status': 'paid', 'advance_description': 'Advance Description', 'milestone_description': 'Milestone Description', 'reference_payment_type_id': 1, 'note': 'Testing', 'milestone_date': date(2022, 9, 16), 'created_at': date(2022, 9, 1), 'date_paid': date(2022, 9, 30) }] contract_id = 1 status = 'paid' params = { 'limit': 10, 'offset': 0, 'reference_payment_type_id': 1 } mock_models.ContractAdvance.get_by_contract_id.return_value = \ (mock_paid_advances, 1) res = logic.get_contract_advances(contract_id, status, params) assert res.status == 200 assert res.message['items'] == [{ 'contract_advance_id': 1, 'contract_id': 1, 'statement_period_id': 1, 'amount': '100.00', 'currency_code': 'USD', 'advance_amount_payee_currency': '100.00', 'advance_payee_currency_code': 'USD', 'milestone': 'recoupment', 'advance_status': 'paid', 'advance_description': 'Advance Description', 'milestone_description': 'Milestone Description', 'note': 'Testing', 'reference_payment_type_id': 1, 'milestone_date': '2022-09-16', 'created_at': '2022-09-01', 'date_paid': '2022-09-30' }] mock_models.ContractAdvance \ .get_by_contract_id \ .assert_called_once_with( contract_id=contract_id, contract_advance_status='paid', reference_payment_type_id=1, limit=10, offset=0 ) @patch('abacus_contract.logic.contract_advance.models') def test_get_contract_advances_invalid_status(mock_models): """Test get_contract_advances function for invalid advance status.""" params = { 'limit': 10, 'offset': 0 } status = 'test' res = logic.get_contract_advances(1, status, params) assert res.status == 400 assert res.errors['message'] == error.ERROR_INVALID_CONTRACT_ADVANCE_STATUS.format( status=', '.join(['pending', 'paid', 'qualified', 'in_review']) ) def test_validate_request_params_success(): """Test _validate_request_params function for valid params.""" ContractFactory.create(contract_id=12456) contract_id = 12456 status = 'pending' params = { 'limit': 10, 'offset': 0 } res = logic._validate_request_params(contract_id, status, params) assert res def test_validate_request_params_failed(): """Test _validate_request_params function for invalid params.""" contract_id = 1 status = 'paid' request_params = { 'limit': 'test', 'offset': 0 } with pytest.raises(Exception, match=error.ERROR_INVALID_LIMIT_OFFSET): logic._validate_request_params(contract_id, status, request_params) status = 'test' request_params = { 'limit': 10, 'offset': 0 } error_msg = error.ERROR_INVALID_CONTRACT_ADVANCE_STATUS.format( status=', '.join(['pending', 'paid']) ) with pytest.raises(Exception, match=error_msg): logic._validate_request_params(contract_id, status, request_params) contract_id = 1 status = 'paid' request_params = { 'limit': 10, 'offset': 0, 'reference_payment_type_id': 'test' } with pytest.raises(Exception, match=error.ERROR_INVALID_REFERENCE_PAYMENT_TYPE_ID): logic._validate_request_params(contract_id, status, request_params) @patch('abacus_contract.logic.contract_advance.models') def test_delete_contract_advance(mock_models): """Test delete_contract_advance method.""" contract_advance = ContractAdvanceFactory.create() assert contract_advance.advance_status == ADVANCE_STATUSES.NOT_QUALIFIED mock_models.ContractAdvance.commit_changes.return_value = None response = logic.delete_contract_advance(contract_advance) assert response.status == 204 assert contract_advance.advance_status == ADVANCE_STATUSES.DELETED mock_models.ContractAdvance.delete_by_id_or_error.assert_called_once_with( contract_advance.contract_advance_id, soft_delete=True ) mock_models.ContractAdvance.commit_changes.assert_called_once() @patch('abacus_contract.logic.contract_advance.models') def test_delete_contract_advance_error(mock_models): """Test delete_contract_advance method for paid contract advance.""" contract_advance = ContractAdvanceFactory.create( advance_status=ADVANCE_STATUSES.PAID ) response = logic.delete_contract_advance(contract_advance) assert response.status == 400 assert response.errors['message'] == error.ERROR_CONTRACT_ADVANCE_PAID_DELETE mock_models.ContractAdvance.delete_by_id_or_error.assert_not_called() mock_models.ContractAdvance.commit_changes.assert_not_called() def test_calculate_contract_advance_fields_payment_enabled(): """Test contract advance fields calculation logic with enabled feature.""" contract_advance = ContractAdvanceFactory.create( amount_after_withholding_and_vat=None ) logic._calculate_contract_advance_fields(contract_advance) assert contract_advance.amount_after_withholding_and_vat == ( contract_advance.amount + contract_advance.withholding_tax_amount + contract_advance.vat_amount )