"""Tests for contract logic.""" from unittest.mock import MagicMock, call, patch import pytest from marshmallow import ValidationError from abacus_contract.constants.constants import CONTRACT_TYPES from abacus_contract.constants.error import ERROR_INVALID_CONTRACT_STATUS from abacus_contract.logic.contract_search import ( _execute_contract_query, _execute_paged_query, _get_search_term_from_params, _validate_contract_statuses, get_contracts, ) from abacus_contract.schemas import ContractDetailSchema from abacus_contract.tests.utils.factories import ContractFactory def test__execute_paged_query(): """Test pagination.""" query = MagicMock() query.order_by.return_value = query.limit.return_value = query.all.return_value = ( query.offset.return_value ) = query _execute_paged_query(query, {'limit': 100, 'offset': 0}) query.limit.assert_called_once_with(100) query.offset.assert_called_once_with(0) query.all.assert_called_once() def test__get_search_term_from_params(): """Test getting contract name value from params dict.""" params1 = {'contract_name': 'Test%20Contract'} res1 = _get_search_term_from_params(params1) params2 = {'contract_name': 'Test Contract'} res2 = _get_search_term_from_params(params2) assert res1 == res2 == 'Test Contract' params3 = {'contract_name': ''} res3 = _get_search_term_from_params(params3) params4 = {'other_parameter': 'false'} res4 = _get_search_term_from_params(params4) assert res3 == res4 == '' @patch('abacus_contract.models.contract.Contract.get_filtered_query') @patch('abacus_contract.logic.contract_search._validate_contract_statuses') @patch('abacus_contract.logic.contract_search._execute_paged_query') @patch('abacus_contract.logic.contract_search._get_account_ids_from_params') @patch('abacus_contract.logic.contract_search._get_search_term_from_params') @patch('abacus_contract.logic.contract.models') def test__execute_contract_query( mock_models, mock_get_search_term_from_params, mock_get_account_ids_from_params, mock_execute_paged_query, mock__validate_contract_statuses, mock_get_filtered_query, ): """Test contract searching function.""" contract_name = 'Test Contract' contract_type = CONTRACT_TYPES.DISTRIBUTION params = { 'contract_type': contract_type, 'contract_name': contract_name, 'search_term': 1234, 'is_excluded_from_accounting_run': 1, 'contract_statuses': 'init,terminated,active', 'run_controller_ids': '1,2', } mock_models.Contract.get_filtered_query.return_value = [] mock_get_search_term_from_params.side_effect = [ contract_type, 1, 'init,terminated,active', '1,2', contract_name, 1234, ] mock__validate_contract_statuses.return_value = True mock_get_account_ids_from_params.return_value = None _execute_contract_query(params) mock_get_search_term_from_params.assert_has_calls( [ call(params, key='contract_type'), call(params, key='is_excluded_from_accounting_run'), call(params, key='contract_statuses'), call(params, key='run_controller_ids'), call(params), call(params, key='search_term'), ] ) mock_get_account_ids_from_params.assert_called_once_with(params) mock_get_filtered_query.assert_called_once_with(**params) mock_execute_paged_query.assert_called_once() @patch('abacus_contract.logic.contract_search._execute_contract_query') def test_get_contracts(mock_execute_contract_query): """Test main search function.""" params = {'contract_name': 'Test Contract', 'account_ids': '1,2'} contracts = ContractFactory.create_batch(3) mock_execute_contract_query.return_value = contracts, len(contracts) res = get_contracts(params) mock_execute_contract_query.assert_called_once_with(params) assert res.message == { 'items': ContractDetailSchema().dump(contracts, many=True), 'total_count': len(contracts), } def test__validate_contract_statuses(): """Test _validate_contract_statuses method.""" statuses = 'init,active,to_be_terminated,terminated,in_collection_period' is_valid = _validate_contract_statuses(statuses) assert is_valid is True def test__validate_contract_statuses_error(): """Test to throw an error when status is invalid.""" with pytest.raises(ValidationError) as excinfo: _validate_contract_statuses('test,active') assert str(excinfo.value) == ERROR_INVALID_CONTRACT_STATUS.format( contract_lifecycle_status='init, active, to_be_terminated, terminated, in_collection_period' )