"""Lambda test module.""" import datetime from types import SimpleNamespace import boto3 import config import index from moto import mock_aws import pytest import test_utils @pytest.fixture def ar_contract_details(): """Return get expiring_contracts query sample data.""" return [ SimpleNamespace( company='test', vendor_id=25824, cont_end=datetime.date(2018, 3, 27) ) ] @pytest.fixture def formatted_contract_details(): """Return formatted expiring_contracts message.""" return ( 'Greetings,\n\n' 'The contracts for the following labels ' 'are set to expire in 30 days (no rollovers enabled):' '\n\n' 'test (ID: 25824) \u2022 ' 'https://oa.theorchard.com/cont_mgmt/view_vendor.php' '?vendor_id=25824') @pytest.fixture def no_contracts_log(): """Return loger message for no expiring contracts.""" return 'There are no expiring Contract(s)' def test_handler(mocker, ar_contract_details): """Test main handler method.""" mocked_get_expiring_contracts = mocker.patch( 'index.get_expiring_contracts') mocked_get_expiring_contracts.return_value = ar_contract_details mocked_send_sns = mocker.patch('index.send_sns') mocked_send_sns.return_value = None result = index.handler(None, None) assert mocked_get_expiring_contracts.call_count == 1 assert mocked_send_sns.call_count == 1 assert result is None def test_handler_logger(mocker, no_contracts_log): """Test main handler method for no expiring contracts.""" mocked_get_expiring_contracts = mocker.patch( 'index.get_expiring_contracts') mocked_get_expiring_contracts.return_value = [] mocked_logger = mocker.patch( 'index.logger.error') mocked_logger.return_value = no_contracts_log result = index.handler(None, None) assert mocked_get_expiring_contracts.call_count == 1 assert mocked_logger.call_count == 1 assert result == no_contracts_log def test_get_expiring_contracts(mocker, ar_contract_details): """Test index.expiring_contracts function.""" # patch session scope mocked_session = test_utils.mock_db_session(mocker) mocked_session.execute.return_value.all.return_value = ar_contract_details sql_text = 'test' mocked_sqlalchemy_text = mocker.patch('index.sqlalchemy.text') mocked_sqlalchemy_text.return_value = sql_text result = index.get_expiring_contracts() assert result == ar_contract_details assert mocked_session.execute.call_args[0] == (sql_text,) def test_format_sns_message(ar_contract_details, formatted_contract_details): """Test format_sns_message method.""" result = index.format_sns_message(ar_contract_details) assert result == formatted_contract_details @mock_aws def test_send_sns(formatted_contract_details): """Test send_sns.""" sns = boto3.client('sns', region_name='us-east-1') topic_res = sns.create_topic(Name='some_topic') sns_topic_arn = topic_res['TopicArn'] config.SNS_ARN = sns_topic_arn config.AWS_REGION = 'us-east-1' message = formatted_contract_details result = index.send_sns(message) assert result['ResponseMetadata']['HTTPStatusCode'] == 200