"""Test for tasks module.""" from unittest.mock import Mock from unittest.mock import patch from botocore.exceptions import ClientError from flows.digital import status from flows.digital import tasks @patch('flows.digital.queries.flows_config') @patch('flows.digital.queries.sql_upcs_condition_in') @patch('flows.digital.tasks.log') @patch('flows.digital.tasks.datawarehouse') def test_unload_from_datawarehouse( datawarehouse, log, sql_upcs_condition_in, flows_config): """Test unload_from_datawarehouse function.""" tasks.queries.get_unload_from_snowflake_sql.batch_size = 3 bucket = 'test-bucket' upcs = ['123', '234', '345', '456', '567'] upcs_batches = [['123', '234', '345'], ['456', '567']] correlation_id = '1234-3456-5678' sql_upcs_condition_in.side_effect = upcs_batches flows_config.AWS_CREDENTIALS = { 'aws_access_key_id': 'aws-key', 'aws_secret_access_key': 'aws-secret'} filler_mocks = ['activity', 'date_end', 'date_start'] params = {key: Mock() for key in filler_mocks} params['bucket'] = bucket params['correlation_id'] = correlation_id params['upcs'] = upcs result = tasks.unload_from_datawarehouse(**params) assert result['batch_count'] == len(upcs_batches) assert bucket in result['unload_path'] assert correlation_id in result['unload_path'] assert result['unload_path'].endswith('/') assert datawarehouse.execute.call_count == len(upcs_batches) execute_calls = datawarehouse.execute.call_args_list for i, execute_call in enumerate(execute_calls): called_sql = execute_call[0][0] assert correlation_id in called_sql assert 'data-{batch}.csv.gz'.format(batch=i) in called_sql log.update_status.assert_called_with(correlation_id, status.UNLOADED) @patch('flows.digital.tasks.datastore') @patch('flows.digital.tasks.log') def test_create_temp_table(log, datastore): """Test create_temp_table function.""" create = Mock() create.format.return_value = 'create a test table' temp_table_name = Mock() temp_table_name.format.return_value = 'test_table_cid1234' params = { 'activity': Mock(), 'correlation_id': 'cid-1234.1', 'create': create, 'temp_table_name': temp_table_name} result = tasks.create_temp_table(**params) temp_table_name.format.assert_called_with(correlation_hex='cid1234') assert result['table_name'] == 'test_table_cid1234' create.format.assert_called_with(table_name='test_table_cid1234') datastore.execute.assert_any_call('create a test table') log.update_status.assert_called_with( 'cid-1234.1', status.TEMP_TABLE_CREATED) @patch('flows.digital.tasks.datastore') @patch('flows.digital.tasks.etl_util') @patch('flows.digital.tasks.log') def test_insert_to_temp_table(log, util, datastore, database_context): """Test insert_to_temp_table function.""" insert_sql = 'some insert query' insert = Mock() insert.format.return_value = insert_sql correlation_id = '1234-3456-5678' params = { 'activity': Mock(), 'batch_count': 2, 'bucket': 'test-bucket', 'correlation_id': correlation_id, 'insert': insert, 'temp_table_name': 'TempyMcTempFace'} batch_1 = Mock() batch_2 = Mock() batch_rows = [batch_1, batch_2] util.read_csv_from_s3.side_effect = batch_rows datastore.context = database_context tasks.insert_to_temp_table(**params) assert database_context._cursor.executemany.call_count == len(batch_rows) database_context._cursor.executemany.assert_any_call(insert_sql, batch_1) database_context._cursor.executemany.assert_any_call(insert_sql, batch_2) log.update_status.assert_called_with( correlation_id, status.TEMP_TABLE_INSERTED) @patch('flows.digital.tasks.datastore') @patch('flows.digital.tasks.etl_util') @patch('flows.digital.tasks.log') def test_insert_to_temp_table_missing(log, util, datastore, database_context): """Test insert_to_temp_table function with some missing files.""" insert_sql = 'some insert query' insert = Mock() insert.format.return_value = insert_sql correlation_id = '1234-3456-5678' params = { 'activity': Mock(), 'batch_count': 2, 'bucket': 'test-bucket', 'correlation_id': correlation_id, 'insert': insert, 'temp_table_name': 'TempyMcTempFace'} batch_1 = Mock() batch_2 = ClientError({'Error': {'Code': 'NoSuchKey'}}, 'GetObject') batch_rows = [batch_1, batch_2] util.read_csv_from_s3.side_effect = batch_rows datastore.context = database_context tasks.insert_to_temp_table(**params) assert database_context._cursor.executemany.call_count == 1 database_context._cursor.executemany.assert_any_call(insert_sql, batch_1) log.update_status.assert_called_with( correlation_id, status.TEMP_TABLE_INSERTED) @patch('flows.digital.tasks.datastore') @patch('flows.digital.tasks.etl_util') @patch('flows.digital.tasks.log') def test_insert_to_temp_table_error(log, util, datastore, database_context): """Test insert_to_temp_table function with reraised error.""" insert_sql = 'some insert query' insert = Mock() insert.format.return_value = insert_sql correlation_id = '1234-3456-5678' params = { 'activity': Mock(), 'batch_count': 2, 'bucket': 'test-bucket', 'correlation_id': correlation_id, 'insert': insert, 'temp_table_name': 'TempyMcTempFace'} batch_1 = Mock() batch_2 = FileNotFoundError('catch this') batch_rows = [batch_1, batch_2] util.read_csv_from_s3.side_effect = batch_rows datastore.context = database_context tasks.insert_to_temp_table(**params) assert database_context._cursor.executemany.call_count == 1 database_context._cursor.executemany.assert_any_call(insert_sql, batch_1) assert not log.update_status.called @patch('flows.digital.tasks.log') def test_update_etl_status(log): """Test update_etl_status function.""" params = { 'activity': Mock(), 'correlation_id': 'cid1234', 'status': 'testing'} tasks.update_etl_status(**params) log.update_status.assert_called_with('cid1234', 'testing') @patch('flows.digital.tasks.datastore') @patch('flows.digital.tasks.log') def test_move_temp_table_to_digital_revenue( log, datastore, database_context): """Test move_temp_table_to_digital_revenue function.""" temp_table = 'temp-test-table' upcs = ['123', '234', '345', '456', '567'] upcs_batches = [['123', '234', '345'], ['456', '567']] correlation_id = '1234-5678-9101-1123' date_params = {'date_end': '2016-01-01', 'date_start': '2015-01-01'} datastore.context = database_context tasks.queries.get_delete_from_digital_revenue_sql.batch_size = 3 tasks.move_temp_table_to_digital_revenue( activity=Mock(), correlation_id=correlation_id, temp_table_name=temp_table, upcs=upcs, **date_params) # List of indexes of upcs_batches. If an index is in here, it means that # batch from upcs_batches has been included in a "delete from" sql call. deleted_batches = set() execute_calls = database_context._cursor.execute.call_args_list for call in execute_calls: sql = call[0][0].strip().lower() params = call[0][1] if len(call[0]) > 1 else None if sql.startswith('delete from'): assert params == date_params for batch, upcs in enumerate(upcs_batches): if all(upc in sql for upc in upcs): deleted_batches.add(batch) elif sql.startswith('insert into'): assert temp_table in sql assert params is None elif sql.startswith('drop table'): assert temp_table in sql assert params is None assert deleted_batches == set(range(len(upcs_batches))) log.update_status.assert_called_with( correlation_id, status.TEMP_TABLE_MADE_LIVE) @patch('flows.digital.tasks.etl_util') def test_send_sns(util): """Test send_sns function.""" activity = Mock() correlation_id = '1234-5678-9101-1120.1' sns_correlation_id = '1234-5678-9101-1120.1.1' date_end = 'foo' date_start = 'bar' upcs = ['123', '456'] tasks.send_sns(activity, correlation_id, date_end, date_start, upcs) util.send_success_notification.assert_called_with( sns_correlation_id, date_end, date_start, upcs) @patch('flows.digital.tasks.etl_util') def test_queue_build_cache(util): """Test queue_build_cache function.""" activity = Mock() correlation_id = '1234-5678-9101-1120.1' sqs_correlation_id = '1234-5678-9101-1120.1.1' date_end = 'foo' date_start = 'bar' upcs = ['123', '456'] tasks.queue_build_cache( activity, correlation_id, date_end, date_start, upcs) util.queue_build_cache.assert_called_with( sqs_correlation_id, date_end, date_start, upcs) @patch('flows.digital.tasks.config') @patch('flows.digital.tasks.garcon_feed_status') @patch('flows.digital.tasks.log') @patch('flows.digital.tasks.status') def test_set_dynamo_status(status, log, garcon_feed_status, config): """Test set_dynamo_status task function.""" correlation_id = '1234-5678-9101-1120.1' date = '2012-12-20' workflow_name = 'test-digital-etl' dynamo_status = 'test-etl-ingested' etl_status = 'test-etl-status-updated' config.SWF_WORKFLOW_NAME = workflow_name garcon_feed_status.STATUS_INGESTED = dynamo_status status.DYNAMO_STATUS_UPDATED = etl_status tasks.set_dynamo_status(Mock(), correlation_id, date) garcon_feed_status.set_overall_status.assert_called_with( workflow_name, date, dynamo_status) log.update_status.assert_called_with(correlation_id, etl_status)