"""Tests for main producer script.""" from collections import OrderedDict import csv import os import tempfile from unittest import mock from pymysql import err import pytest from snowflake.connector import errors from accounting import config from accounting import const from accounting import producer from accounting.models import sql from accounting.models import sql_fp DIR = os.path.dirname(os.path.realpath(__file__)) TEST_EXPORT_FILE = f'{DIR}/export.csv' TEST_FILE = f'{DIR}/test_file.csv' TEST_DIR = f'{DIR}/test_dir' TEST_GZIP = f'{DIR}/snowflake_data.csv.gz' @pytest.fixture def db_rows(): """Rows exported from stmt DB fixture.""" return [ {'isrc': 'TEST_ISRC1', 'territory': 'US'}, {'isrc': 'TEST_ISRC2', 'territory': 'CA'}, {'isrc': 'TEST_ISRC3', 'territory': 'RU'} ] @pytest.fixture def db_report_rows(): """Rows updated in stmt DB fixture.""" return [ { 'isrc': 'TEST_ISRC1', 'territory': 'US', 'tuid': '123', 'internal_conflict': '0' }, { 'isrc': 'TEST_ISRC2', 'territory': 'CA', 'tuid': '123', 'internal_conflict': '0' }, { 'isrc': 'TEST_ISRC3', 'territory': 'RU', 'tuid': None, 'internal_conflict': '1' } ] @mock.patch('accounting.producer.mysql.get_connection') def test_export_csv(connect_mock, db_rows): """Test export CSV file with ISRC, Territory data from stmt DB.""" cursor_mock = mock.MagicMock() cursor_mock.__iter__.return_value = iter(db_rows) connection_mock = mock.MagicMock() connection_mock.cursor.return_value.__enter__.return_value = cursor_mock connect_mock.return_value = connection_mock with tempfile.NamedTemporaryFile('w') as f: producer.export_csv(f.name) exported = open(f.name, 'r') reader = csv.DictReader( exported, fieldnames=[const.ISRC, const.TERRITORY]) with open(TEST_EXPORT_FILE) as export_file: csv_file = csv.DictReader( export_file, fieldnames=[const.ISRC, const.TERRITORY]) for test, expected in zip(reader, csv_file): assert test == expected connection_mock.close.assert_called() @mock.patch( 'accounting.producer.mysql.get_connection', side_effect=err.MySQLError) def test_export_csv_connection_failed(connect_mock): """Test export CSV failed due to connection error.""" with pytest.raises(err.MySQLError): producer.export_csv('export_failed.csv') @mock.patch('accounting.producer.sys.exit') @mock.patch('accounting.producer.mysql.get_connection') def test_export_csv_db_read_failed(connect_mock, sys_exit_mock): """Test export CSV failed.""" cursor_mock = mock.MagicMock( execute=mock.MagicMock(side_effect=err.MySQLError)) connection_mock = mock.MagicMock() connection_mock.cursor.return_value.__enter__.return_value = cursor_mock connect_mock.return_value = connection_mock with tempfile.NamedTemporaryFile('w') as f: producer.export_csv(f.name) connection_mock.close.assert_called() sys_exit_mock.assert_called_with(1) @mock.patch('accounting.producer.snowflake.get_snowflake_connection') def test_process_snowflake(connect_mock): """Test process data in Snowflake.""" cursor_mock = mock.MagicMock() connection_mock = mock.MagicMock() connection_mock.cursor.return_value = cursor_mock connect_mock.return_value = connection_mock producer.process_snowflake(TEST_FILE, TEST_DIR) cursor_mock.execute.assert_has_calls([ mock.call(sql.CREATE_TMP_TABLE), mock.call(sql.PUT_CSV.format(file_name=TEST_FILE)), mock.call(sql.COPY_CSV), mock.call(sql.UPDATE_TMP_TABLE), mock.call(sql.UNLOAD_CSV.format(max_file_size=config.MAX_FILE_SIZE)), mock.call(sql.GET_CSV.format(file_name=TEST_DIR)) ]) connection_mock.close.assert_called() @mock.patch('accounting.producer.sys.exit') @mock.patch('accounting.producer.snowflake.get_snowflake_connection') def test_process_snowflake_failed(connect_mock, sys_exit_mock): """Test process data in Snowflake failed.""" cursor_mock = mock.MagicMock( execute=mock.MagicMock(side_effect=errors.Error)) connection_mock = mock.MagicMock() connection_mock.cursor.return_value = cursor_mock connect_mock.return_value = connection_mock producer.process_snowflake(TEST_FILE, TEST_DIR) connection_mock.rollback.assert_called() connection_mock.close.assert_called() sys_exit_mock.assert_called_with(1) @mock.patch( 'accounting.producer.snowflake.get_snowflake_connection', side_effect=errors.Error) def test_process_snowflake_connection_failed(connect_mock): """Test process data in Snowflake failed due to connection error.""" with pytest.raises(errors.Error): producer.process_snowflake(TEST_FILE, TEST_DIR) @mock.patch('accounting.producer.mysql.get_connection') def test_import_csv(connect_mock, db_report_rows): """Test import an updated CSV file back into stmt DB.""" cursor_mock = mock.MagicMock() connection_mock = mock.MagicMock() connection_mock.cursor.return_value.__enter__.return_value = cursor_mock connect_mock.return_value = connection_mock query = sql.UPDATE_REPORT.format(table_name=config.MYSQL_TABLE_NAME) producer.import_csv(TEST_GZIP) cursor_mock.execute.assert_has_calls([ mock.call(query, OrderedDict(params)) for params in db_report_rows]) connection_mock.close.assert_called() @mock.patch('accounting.producer.tempfile.TemporaryDirectory') @mock.patch('accounting.producer.import_csv') @mock.patch('accounting.producer.export_csv', return_value=3) @mock.patch('accounting.producer.snowflake.get_snowflake_connection') @mock.patch('accounting.producer.mysql.get_connection') @mock.patch('accounting.config.USE_MRR', True) def test_start_uses_mrr_queries_when_use_mrr_true( connect_mock, snowflake_connect_mock, export_csv_mock, import_csv_mock, tmpdir_mock, db_rows): """Uses sql (MRR) when USE_MRR is True.""" tmpdir_mock.return_value.__enter__.return_value = '/tmp/fake_dir' cursor_mock = mock.MagicMock() cursor_mock.__iter__.return_value = iter(db_rows) connection_mock = mock.MagicMock() connection_mock.cursor.return_value.__enter__.return_value = cursor_mock connect_mock.return_value = connection_mock snowflake_cursor_mock = mock.MagicMock() snowflake_connection_mock = mock.MagicMock() snowflake_connection_mock.cursor.return_value = snowflake_cursor_mock snowflake_connect_mock.return_value = snowflake_connection_mock producer.start() snowflake_cursor_mock.execute.assert_any_call(sql.CREATE_TMP_TABLE) snowflake_cursor_mock.execute.assert_any_call(sql.UPDATE_TMP_TABLE) assert mock.call(sql_fp.CREATE_TMP_TABLE) not in \ snowflake_cursor_mock.execute.call_args_list assert mock.call(sql_fp.UPDATE_TMP_TABLE) not in \ snowflake_cursor_mock.execute.call_args_list snowflake_connection_mock.close.assert_called() @mock.patch( 'accounting.producer.mysql.get_connection', side_effect=err.MySQLError) def test_import_csv_connection_failed(connect_mock): """Test import an updated CSV failed due to connection error.""" with pytest.raises(err.MySQLError): producer.import_csv(TEST_GZIP) @mock.patch('accounting.producer.sys.exit') @mock.patch('accounting.producer.mysql.get_connection') def test_import_csv_db_update_failed(connect_mock, sys_exit_mock): """Test import an updated CSV failed due to database error.""" cursor_mock = mock.MagicMock( execute=mock.MagicMock(side_effect=err.MySQLError)) connection_mock = mock.MagicMock() connection_mock.cursor.return_value.__enter__.return_value = cursor_mock connect_mock.return_value = connection_mock producer.import_csv(TEST_GZIP) connection_mock.rollback.assert_called() connection_mock.close.assert_called() sys_exit_mock.assert_called_with(1)