"""Tests for util.py module.""" from argparse import ArgumentTypeError import gzip import json import math import types from unittest.mock import Mock from unittest.mock import patch from uuid import uuid4 import pytest from flows import queries from flows import util @pytest.fixture def vendor_upcs(): """DB rows for vendor and UPCs, some multiple.""" return [ ['one', '1'], ['one', 'uno'], ['two', '2'], ['three', '3']] @pytest.fixture(params=[ None, 'my_workflow_id']) def workflow_id(request): """Define fixture for optional workflow_id parameter.""" return request.param def test_correlation_id_hex(): """Test correlation_id_hex function.""" for i in range(5): cid = uuid4() assert cid.hex == util.correlation_id_hex(str(cid)) def test_upc_cli_type(): """Test upc_cli_type function.""" valids = ['012345678901', '111111111111', '2222222222222'] invalids = ['+123', '-123', 'abc', '1.1', '1234'] for valid in valids: util.upc_cli_type(valid) for invalid in invalids: with pytest.raises(ArgumentTypeError): util.upc_cli_type(invalid) def test_date_type(): """Test date_type function.""" util.date_cli_type('1234-56-78') invalids = ['1234-5-67', '12-12-21', '12345678'] for invalid in invalids: with pytest.raises(ArgumentTypeError): util.date_cli_type(invalid) @patch('flows.util.config') def test_create_swf_execution_params(config, workflow_id): """Test create_swf_execution_params function.""" config.SWF_DOMAIN = uuid4().hex timeout = 3600 workflow_name = uuid4().hex workflow_version = uuid4().hex correlation_id = str(uuid4()) if workflow_id: workflow_id_value = workflow_id else: workflow_id_value = '{name}_{cid}'.format( name=workflow_name, cid=correlation_id) expected = { 'domain': config.SWF_DOMAIN, 'executionStartToCloseTimeout': str(timeout), 'tagList': ['cid:{cid}'.format(cid=correlation_id)], 'taskList': {'name': workflow_name}, 'workflowId': workflow_id_value, 'workflowType': {'name': workflow_name, 'version': workflow_version}} result = util.create_swf_execution_params( correlation_id, timeout, workflow_name, workflow_version, workflow_id) assert result == expected @patch('flows.util.boto3') @patch('flows.util.create_swf_execution_params') def test_start_swf_execution(create_swf, boto3, workflow_id): """Test with explicit swf execution context.""" expected_swf_response = {'message': 'it has the success!'} swf = Mock() swf.start_workflow_execution.return_value = expected_swf_response boto3.client.return_value = swf swf_params = { 'some': 'params', 'for': 'the', 'amazon swf': 'service'} create_swf.return_value = swf_params correlation_id = str(uuid4()) context = { 'correlation_id': correlation_id, 'more': 'random', 'data': 'for', 'testing': 'purposes'} response = util.start_swf_execution(context, 0, 'foo', 'bar', workflow_id) swf.start_workflow_execution.assert_called_once_with( input=json.dumps(context), **swf_params) assert response == expected_swf_response @patch('flows.util.boto3') @patch('flows.util.create_swf_execution_params') @patch('flows.util.uuid') def test_start_empty_swf_execution(uuid, create_swf, boto3): """Test with empty swf execution context.""" expected_swf_response = {'message': 'it has the success!'} swf = Mock() swf.start_workflow_execution.return_value = expected_swf_response boto3.client.return_value = swf swf_params = {'some': 'params'} create_swf.return_value = swf_params # this is the default correlation ID due to the empty context correlation_uuid = uuid4() uuid.uuid1.return_value = correlation_uuid context = {'correlation_id': str(correlation_uuid)} response = util.start_swf_execution({}, 0, 'foo', 'bar') create_call = create_swf.call_args_list[0] assert create_call[0][0] == context['correlation_id'] swf.start_workflow_execution.assert_called_once_with( input=json.dumps(context), **swf_params) assert response == expected_swf_response @patch('flows.util.boto3') def test_send_sns_message(boto): """Test send_sns_message function.""" client = Mock() boto.client.return_value = client util.send_sns_message('foo', 'bar', 'baz') client.publish.assert_called_with( TopicArn='baz', Message='bar', Subject='foo') @patch('flows.util.get_serviced_vendor_id_upc_map') def test_get_serviced_upcs(mapping): """Test get_serviced_upcs function. Testing done with sets since natively data is in sets, and ordering is not consistent. """ mapping.return_value = { 'foo': {'123', '234', '345'}, 'bar': {'one', 'two', 'three'}} expected = {'123', '234', '345', 'one', 'two', 'three'} results = util.get_serviced_upcs() mapping.assert_called_once_with(False) assert expected == set(results) @patch('flows.util.art_relations') def test_get_serviced_upcs_query(art_relations, vendor_upcs): """Test get_serviced_upcs function with query. Testing done with sets since natively data is in sets, and ordering is not consistent. """ expected = {'1', 'uno', '2', '3'} art_relations.query.return_value = vendor_upcs results = util.get_serviced_upcs(True) assert expected == set(results) @patch('flows.util.art_relations') @patch('flows.util.config') @patch('flows.util.queries') def test_get_upc_vendor_ids(queries, config, art_relations): """Test get_upc_vendor_ids lookup function.""" batch_size = 2 upc_list = ['123', '234', '345', '456', '567'] upc_vendor_rows = [ [['123', '321'], ['234', '432']], [['345', '543'], ['456', '654']], [['567', '765']]] expected = [ {'123': '321', '234': '432'}, {'345': '543', '456': '654'}, {'567': '765'}] art_relations.query.side_effect = upc_vendor_rows util.get_upc_vendor_id.batch_size = batch_size results_raw = util.get_upc_vendor_id(upc_list) results = list(results_raw) call_count = math.ceil(len(upc_list) / batch_size) assert queries.upc_vendor_id_lookup.call_count == call_count assert art_relations.query.call_count == call_count assert results == expected @patch('flows.util.config') def test_batch_upc_calls(config): """Test batch_upc_calls decorator.""" config.ITERATION_BATCH_SIZE = 5 @util.batch_upc_calls def callee(upcs): return sum(upcs) upcs = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17] expected = [10, 35, 60, 48] results = callee(upcs) assert isinstance(results, types.GeneratorType) assert list(results) == expected def test_batch_upc_calls_skip(): """Test batch_upc_calls decorator when it doesn't change logic.""" @util.batch_upc_calls def callee(name): return 'Hello, ' + name + '!' result = callee('unit test') assert not isinstance(result, types.GeneratorType) assert result == 'Hello, unit test!' def test_batch_param_calls(): """Test batch_param_calls decorator.""" @util.batch_param_calls('upcs', 5) def callee(upcs): return sum(upcs) upcs = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17] expected = [10, 35, 60, 48] results = callee(upcs) assert isinstance(results, types.GeneratorType) assert list(results) == expected def test_batch_param_calls_skip(): """Test batch_param_calls decorator when it doesn't change logic.""" @util.batch_param_calls('upcs', 5) def callee(name): return 'Hello, ' + name + '!' result = callee('unit test') assert not isinstance(result, types.GeneratorType) assert result == 'Hello, unit test!' @patch('flows.util.art_relations') @patch('flows.util.constant') @patch('flows.util.queries') def test_get_serviced_vendor_id_upc_map(queries, constant, art_relations): """Test whitelist lookup from hard coded constant.""" constant.FILM_WHITELIST = 'foo' results = util.get_serviced_vendor_id_upc_map() assert results == constant.FILM_WHITELIST assert not queries.get_serviced_vendor_id_upc_map.called assert not art_relations.query.called @patch('flows.util.art_relations') def test_get_serviced_vendor_id_upc_map_query(art_relations, vendor_upcs): """Test whitelist lookup from query lookup.""" art_relations.query.return_value = vendor_upcs expected = {'one': {'1', 'uno'}, 'two': {'2'}, 'three': {'3'}} results = util.get_serviced_vendor_id_upc_map(query_based=True) assert results == expected art_relations.query.assert_called_once_with( queries.serviced_vendor_id_upc_map_sql) @patch('flows.util.s3') def test_read_gzip_csv_from_s3(s3_mock): """Test read_gzip_csv_from_s3 function.""" mock_file = Mock() csv_data = [['a1', 'b1', 'c1'], ['a2', 'b2', 'c2']] mock_file.read.return_value = gzip.compress(b'a1,b1,c1\na2,b2,c2') mock_s3_object = Mock() mock_s3_object.get.return_value = {'Body': mock_file} s3_mock.get_object.return_value = mock_s3_object result = util.read_gzip_csv_from_s3('s3://test-bucket/test-file.csv.gz') for each_row in result: assert each_row in csv_data