"""Tests for main producer script.""" 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 as sql_mrr from accounting.models import sql_fp as sql DIR = os.path.dirname(os.path.realpath(__file__)) TEST_EXPORT_FILE = f'{DIR}/fp_export.csv' TEST_FILE = f'{DIR}/test_file.csv' TEST_DIR = f'{DIR}/test_dir' TEST_GZIP = f'{DIR}/fp_data.csv.gz' @pytest.fixture def db_rows(): """Rows exported from stmt DB fixture.""" return [ { 'service': 'TEST_SERVICE', 'date': '2024-01-01', 'isrc': 'TEST_ISRC1', 'territory': 'US' }, { 'service': 'TEST_SERVICE', 'date': '2024-01-01', 'isrc': 'TEST_ISRC2', 'territory': 'CA' }, { 'service': 'TEST_SERVICE', 'date': '2024-01-01', 'isrc': 'TEST_ISRC3', 'territory': 'RU' }, ] @pytest.fixture def db_report_rows(): """Rows updated in stmt DB fixture.""" return [ { 'service': 'TEST_SERVICE', 'date': '2024-01-01', 'isrc': 'TEST_ISRC1', 'territory': 'US', 'tuid': '123', 'internal_conflict': '0', 'rules_summary': '[{"policy": "standard"}]', }, { 'service': 'TEST_SERVICE', 'date': '2024-01-01', 'isrc': 'TEST_ISRC2', 'territory': 'CA', 'tuid': '123', 'internal_conflict': '0', 'rules_summary': '', }, { 'service': 'TEST_SERVICE', 'date': '2024-01-01', 'isrc': 'TEST_ISRC3', 'territory': 'RU', 'tuid': None, 'internal_conflict': '1', 'rules_summary': '', }, ] @mock.patch('accounting.producer.mysql.get_connection') @mock.patch('accounting.config.USE_MRR', False) 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.SERVICE, const.DATE_COLUMN_NAME, const.ISRC, const.TERRITORY]) with open(TEST_EXPORT_FILE) as export_file: csv_file = csv.DictReader( export_file, fieldnames=[ const.SERVICE, const.DATE_COLUMN_NAME, 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) @mock.patch('accounting.config.USE_MRR', False) 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') @mock.patch('accounting.config.USE_MRR', False) 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') @mock.patch('accounting.config.USE_MRR', False) 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_for_fingerprinting(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') @mock.patch('accounting.config.USE_MRR', False) 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_for_fingerprinting(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) @mock.patch('accounting.config.USE_MRR', False) 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_for_fingerprinting(TEST_FILE, TEST_DIR) @mock.patch('accounting.producer.mysql.get_connection') @mock.patch('accounting.config.USE_MRR', False) 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 bulk_update_sql = sql.BULK_UPDATE_REPORT.format( table_name=config.MYSQL_TABLE_NAME) expected_rows = [ (r[const.SERVICE], r[const.DATE_COLUMN_NAME], r[const.ISRC], r[const.TERRITORY], r[const.TUID], r[const.INTERNAL_CONFLICT], r.get(const.RULES_SUMMARY)) for r in db_report_rows ] producer.import_csv(TEST_GZIP) cursor_mock.execute.assert_has_calls([ mock.call(sql.CREATE_TMP_UPDATE_TABLE), mock.call(sql.TRUNCATE_TMP_UPDATE_TABLE), mock.call(bulk_update_sql), ]) cursor_mock.executemany.assert_called_once_with( sql.INSERT_TMP_UPDATE, expected_rows) connection_mock.close.assert_called() @mock.patch( 'accounting.producer.mysql.get_connection', side_effect=err.MySQLError) def test_import_csv_connection_failed(connect_mock, monkeypatch): """Test import an updated CSV failed due to connection error.""" monkeypatch.setenv('USE_MRR', 'False') with pytest.raises(err.MySQLError): producer.import_csv(TEST_GZIP) @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', False) def test_start_uses_fp_queries_when_use_mrr_false( connect_mock, snowflake_connect_mock, export_csv_mock, import_csv_mock, tmpdir_mock, db_rows): """Uses sql_fp when USE_MRR is False.""" 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_mrr.CREATE_TMP_TABLE) not in \ snowflake_cursor_mock.execute.call_args_list assert mock.call(sql_mrr.UPDATE_TMP_TABLE) not in \ snowflake_cursor_mock.execute.call_args_list snowflake_connection_mock.close.assert_called() @mock.patch('accounting.producer.sys.exit') @mock.patch('accounting.producer.mysql.get_connection') @mock.patch('accounting.config.USE_MRR', False) def test_import_csv_db_update_failed(connect_mock, sys_exit_mock, monkeypatch): """Test import an updated CSV failed due to database error.""" monkeypatch.setenv('USE_MRR', 'False') 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)