"""Unit tests for tasks of the sql2sf workflow.""" from datetime import datetime import json import subprocess import tempfile from unittest.mock import call from unittest.mock import MagicMock from unittest.mock import Mock from unittest.mock import patch from freezegun import freeze_time from garcon_contrib.aws.utils import garcon_s3 from garcon_contrib.snowflake import garcon_snowflake import pytest from snowflake_etl.flows.s32sf import helpers as s32sf_helpers from snowflake_etl.flows.sql2sf import helpers from snowflake_etl.flows.sql2sf import tasks @freeze_time('2015-05-10') def test_bootstrap(monkeypatch): """Test bootstrap task.""" activity_mock = MagicMock() sources_mock = MagicMock() sources_mock.is_table_supports_incremental_sync = MagicMock( return_value=False) sources_mock.is_schema_validation_needed = MagicMock( return_value=False) sources_mock.get_query = MagicMock( return_value='SELECT * FROM some_table') sources_mock.get_table_params = MagicMock(return_value={}) sources_mock.get_s3_path_data = MagicMock( return_value='some_s3_path_data/z/f/k/j/') sources_mock.get_s3_path_schema = MagicMock( return_value='some_s3_path_schema/z/f/k/j/') monkeypatch.setattr(tasks, 'sources', sources_mock) extract_bucket_path_mock = MagicMock( return_value=('some_bucket', 'some_s3_key')) monkeypatch.setattr( garcon_s3, 'extract_bucket_path', extract_bucket_path_mock) result = tasks.bootstrap( activity_mock, 'dev', 'mysql', 'reportsar04', 'art_relations', 'vendor_contract', 'snapshot') assert result['db_type'] == 'mysql' assert result == { 'source_schema': 'art_relations', 'unload_source_sql': 'SELECT * FROM some_table', 'source_table': 'vendor_contract', 'source_db_host': 'reportsar04', 'destination_s3_bucket': 'some_bucket', 'destination_s3_key': 'some_s3_key', 's3_path_data': 'some_s3_path_data/z/f/k/j/', 's3_prefix_to_clear': 'some_s3_path_data/z/f/', 's3_prefix_to_load': 'some_s3_path_data/z/f/k/j', 'db_type': 'mysql', 's3_path_schema': 'some_s3_path_schema/z/f/k/j/', 'query_type': 'snapshot', 'file_format': None, 'start_period': None, 'end_period': None, 'validate_schema': str(False)} @pytest.fixture def sf_config_mock(): """Fixture returning a dict with a mock Snowflake config.""" return { 'role': 'testrole', 'warehouse': 'testwh', 'db': 'testdb', 'schema': 'testschema', 'user': 'testuser', 'password': 'testpass', 'account': 'testacc'} def test_pipe_mysql_from_stdin_to_stdout(monkeypatch): """Test customized pipe_mysql_from_stdin_to_stdout task.""" activity_mock = MagicMock() sources_mock = MagicMock() sources_mock.get_db_credentials = MagicMock( return_value=('some_host', 0, 'some_user', 'some_password')) pipe_mock = MagicMock() monkeypatch.setattr(tasks, 'sources', sources_mock) subprocess_mock = MagicMock() popen_mock = MagicMock(return_value=subprocess_mock) monkeypatch.setattr(subprocess, 'Popen', popen_mock) fd = tempfile.TemporaryFile(mode='w+t') tempfile_mock = MagicMock(side_effect=[fd]) monkeypatch.setattr(tempfile, 'TemporaryFile', tempfile_mock) result = tasks.pipe_mysql_from_stdin_to_stdout( activity_mock, pipe_mock, 'some_schema', 'some_table') assert result == {'pipe': subprocess_mock, 'mysql_stderr': fd} fd.close() def test_pipe_mysql_from_stdin_to_stdout_empty_pipe(monkeypatch): """Test customized pipe_mysql_from_stdin_to_stdout task.""" activity_mock = MagicMock() sources_mock = MagicMock() sources_mock.get_db_credentials = MagicMock( return_value=('some_host', 0, 'some_user', 'some_password')) monkeypatch.setattr(tasks, 'sources', sources_mock) popen_mock = MagicMock(return_value='pipe') monkeypatch.setattr(subprocess, 'Popen', popen_mock) with pytest.raises(Exception): tasks.pipe_mysql_from_stdin_to_stdout( activity_mock, None, 'some_schema', 'some_table') def test_get_chunk_query_not_chunked(monkeypatch): """Test get_chunk_query task if not unload in chunks.""" activity_mock = MagicMock() sources_mock = MagicMock() sources_mock.is_unload_in_chunks = MagicMock(return_value=False) monkeypatch.setattr(tasks, 'sources', sources_mock) popen_mock = MagicMock(return_value='pipe') monkeypatch.setattr(subprocess, 'Popen', popen_mock) sources_mock.get_stripspecialchars_func_name = MagicMock( return_value='') result = tasks.get_chunk_query( activity_mock, 'some_host', 'some_schema', 'some_table', 'some_s3_key.gz', 1, 50, [('sometable1', True), ('sometable2', None)]) assert result == dict(pipe='pipe', destination_s3_key='some_s3_key.gz') def test_get_chunk_query_sql(monkeypatch): """Test get_chunk_query SQL query.""" activity_mock = MagicMock() sources_mock = MagicMock() sources_mock.is_unload_in_chunks = MagicMock(return_value=False) sources_mock.get_table_params.return_value = { 'wrap_to_convert_unicode': ['sometable2']} sources_mock.get_stripspecialchars_func_name = MagicMock( return_value='stripSpecialChars') monkeypatch.setattr(tasks, 'sources', sources_mock) popen_mock = MagicMock(return_value='pipe') monkeypatch.setattr(subprocess, 'Popen', popen_mock) tasks.get_chunk_query( activity_mock, 'some_host', 'some_schema', 'some_table', 'some_s3_key.gz', 1, 50, [('sometable1', True), ('sometable2', None)]) activity_mock.assert_has_calls([ call.logger.info( 'Setting up query: SELECT stripSpecialChars(sometable1), ' 'CONVERT((sometable2), CHAR(5000) UNICODE) FROM `some_table`')]) def test_get_chunk_query_sql_stripspecialcharslong(monkeypatch): """Test get_chunk_query SQL query.""" activity_mock = MagicMock() sources_mock = MagicMock() sources_mock.is_unload_in_chunks = MagicMock(return_value=False) sources_mock.get_table_params.return_value = { 'wrap_to_convert_unicode': ['sometable2']} sources_mock.get_stripspecialchars_func_name = MagicMock( return_value='stripSpecialCharsLong') monkeypatch.setattr(tasks, 'sources', sources_mock) popen_mock = MagicMock(return_value='pipe') monkeypatch.setattr(subprocess, 'Popen', popen_mock) tasks.get_chunk_query( activity_mock, 'some_host', 'some_schema', 'some_table', 'some_s3_key.gz', 1, 50, [('sometable1', True), ('sometable2', None)]) activity_mock.assert_has_calls([ call.logger.info( 'Setting up query: SELECT stripSpecialCharsLong(sometable1), ' 'CONVERT((sometable2), CHAR(5000) UNICODE) FROM `some_table`')]) def test_get_chunk_query_chunked(monkeypatch): """Test get_chunk_query task if unload in chunks.""" activity_mock = MagicMock() sources_mock = MagicMock() sources_mock.is_unload_in_chunks = MagicMock(return_value=True) sources_mock.get_primary_key = MagicMock(return_value='some_pk') monkeypatch.setattr(tasks, 'sources', sources_mock) sources_mock.get_stripspecialchars_func_name = MagicMock( return_value='') popen_mock = MagicMock(return_value='pipe') monkeypatch.setattr(subprocess, 'Popen', popen_mock) result = tasks.get_chunk_query( activity_mock, 'some_host', 'some_schema', 'some_table', 'some_s3_key.gz', 1, 50, [('sometable1', True), ('sometable2', None)]) assert result == dict(pipe='pipe', destination_s3_key='some_s3_key1-50.gz') def test_describe_mysql_table(monkeypatch): """Test describe_mysql_table task.""" execute_mock = MagicMock(return_value={}) monkeypatch.setattr(helpers, 'execute_with_mysql', execute_mock) sources_mock = MagicMock() sources_mock.get_db_credentials = MagicMock( return_value=('some_host', 0, 'some_user', 'some_password')) monkeypatch.setattr(tasks, 'sources', sources_mock) activity_mock = MagicMock() result = tasks.describe_mysql_table( activity_mock, 'reportsar', 'art_relations', 'vendor_contact') assert json.loads(result['describe_json']) == { 'table': 'vendor_contact', 'columns': {}, 'schema': 'art_relations'} def test_sanity_check_rows_count(monkeypatch): """Test sanity_check_rows_count task.""" execute_mock = MagicMock(return_value=[0]) monkeypatch.setattr(helpers, 'execute_with_mysql', execute_mock) monkeypatch.setattr(garcon_snowflake, 'execute_with_py_conn', MagicMock( return_value={'results': [0]})) monkeypatch.setattr(s32sf_helpers, 'extract_db_and_schema', MagicMock(return_value=('somedb', 'someschema'))) sources_mock = MagicMock() sources_mock.get_db_credentials = MagicMock( return_value=('some_host', 0, 'some_user', 'some_password')) sources_mock.get_sync_sanity_threshold = MagicMock( return_value=0) monkeypatch.setattr(tasks, 'sources', sources_mock) activity_mock = MagicMock() assert not tasks.sanity_check_rows_count( activity_mock, 'mysql', {}, 'snapshot', 'reportsar', 'art_relations', 'vendor_contact') def test_sanity_check_rows_skip(monkeypatch): """Test sanity_check_rows_count task skip if not MySQL.""" execute_mock = MagicMock(return_value=[0]) monkeypatch.setattr(helpers, 'execute_with_mysql', execute_mock) monkeypatch.setattr(garcon_snowflake, 'execute_with_py_conn', MagicMock( return_value={'results': [0]})) monkeypatch.setattr(s32sf_helpers, 'extract_db_and_schema', MagicMock(return_value=('somedb', 'someschema'))) sources_mock = MagicMock() sources_mock.get_db_credentials = MagicMock( return_value=('some_host', 0, 'some_user', 'some_password')) sources_mock.get_sync_sanity_threshold = MagicMock( return_value=0) monkeypatch.setattr(tasks, 'sources', sources_mock) activity_mock = MagicMock() assert tasks.sanity_check_rows_count( activity_mock, 'redshift', {}, 'snapshot', 'reportsar', 'art_relations', 'vendor_contact') is None def test_sanity_check_rows_count_fail(monkeypatch): """Test sanity_check_rows_count task.""" execute_mock = MagicMock(return_value=[200]) monkeypatch.setattr(helpers, 'execute_with_mysql', execute_mock) monkeypatch.setattr(garcon_snowflake, 'execute_with_py_conn', MagicMock( return_value={'results': [1000000]})) monkeypatch.setattr(s32sf_helpers, 'extract_db_and_schema', MagicMock(return_value=('somedb', 'someschema'))) sources_mock = MagicMock() sources_mock.get_db_credentials = MagicMock( return_value=('some_host', 0, 'some_user', 'some_password')) sources_mock.get_sync_sanity_threshold = MagicMock( return_value=100) monkeypatch.setattr(tasks, 'sources', sources_mock) activity_mock = MagicMock() result = tasks.sanity_check_rows_count( activity_mock, 'mysql', {}, 'snapshot', 'reportsar', 'art_relations', 'vendor_contact') assert result == { 'message': 'Sync strategy is snapshot, but source table row count ' 'is 200, and staging destination row count is 1000000', 'stop': 'True'} @freeze_time(datetime.strptime('2017-06-15T14:02:00', '%Y-%m-%dT%H:%M:%S')) @patch('snowflake_etl.flows.sql2sf.tasks.boto3') def test_set_sync_status_snapshot(boto3_mock): """Test set_sync_status task for the snapshot mode.""" params = { 'activity': Mock(), 'load_strategy': 'snapshot', 'end_period': None, 'status': 'INGESTED', 'source_schema': 'test_schema', 'source_table': 'test_table' } dynamodb_mock = Mock() boto3_mock.client.return_value = dynamodb_mock tasks.set_sync_status(**params) call, *_ = dynamodb_mock.put_item.call_args_list # first call (*_, data_arg) = call # first argument in positional arguments assert data_arg == {'TableName': 'sql2sf_table_sync_status', 'Item': {'schema_name': 'test_schema', 'table_name': 'test_table', 'last_sync_timestamp': '2017-06-15T14:02:00', 'synced_to_timestamp': '2017-06-15T14:02:00', 'status': 'INGESTED'}} @freeze_time(datetime.strptime('2017-06-15T14:02:00', '%Y-%m-%dT%H:%M:%S')) @patch('snowflake_etl.flows.sql2sf.tasks.boto3') def test_set_sync_status_incremental(boto3_mock): """Test set_sync_status task for the incremental mode.""" params = { 'activity': Mock(), 'load_strategy': 'incremental', 'end_period': '2017-06-14T15:00:00', 'status': 'INGESTED', 'source_schema': 'test_schema', 'source_table': 'test_table' } dynamodb_mock = Mock() boto3_mock.client.return_value = dynamodb_mock tasks.set_sync_status(**params) call, *_ = dynamodb_mock.put_item.call_args_list # first call (*_, data_arg) = call # first argument in positional arguments assert data_arg == {'TableName': 'sql2sf_table_sync_status', 'Item': {'schema_name': 'test_schema', 'table_name': 'test_table', 'last_sync_timestamp': '2017-06-15T14:02:00', 'synced_to_timestamp': '2017-06-14T15:00:00', 'status': 'INGESTED'}}