from contextlib import contextmanager from unittest.mock import call from unittest.mock import MagicMock from unittest.mock import patch import pymysql import pytest from processing_accounting.util import db def test_ResultIter(monkeypatch): mock_cursor = MagicMock() mock_results = [] mock_result_obj_1 = MagicMock() mock_result_obj_2 = MagicMock() mock_results.append(mock_result_obj_2) mock_results.append(mock_result_obj_1) monkeypatch.setattr( mock_cursor, 'fetchmany', MagicMock( return_value=(mock_result_obj_1, mock_result_obj_2))) i = 0 for actual_result in db.ResultIter(mock_cursor): expected_result = mock_results.pop() assert actual_result == expected_result i += 1 # After 2 iteration, reset the mock object, simulating fetchmany cant # fetch anything if i == 2: monkeypatch.setattr( mock_cursor, 'fetchmany', MagicMock( return_value=None)) @pytest.mark.parametrize('sql,params', [ ('SELECT * from fact_sales', {'id': 1234}), ('UPDATE dim_track', {'track_id': 42, 'title': '🦍'}), ("Robert'); DROP TABLE Students;--", {}) ]) @patch('processing_accounting.util.db.SnowflakeSQLExecutor') def test_execute(snowflake_sql_mock, sql, params): """Test executing a SQL query on Snowflake.""" executer_mock = MagicMock() enter_mock = MagicMock() enter_mock.__enter__.return_value = executer_mock snowflake_sql_mock.return_value = enter_mock db.snowflake_execute(sql, **params) executer_mock.execute.assert_called_with(sql, params=params) @patch('processing_accounting.util.db.sf_connector') @patch('processing_accounting.util.db.ResultIter') def test_snowflake_query(resultiter_mock, sf_conn_mock): resultiter_mock.return_value = [''] actual_generator = db.snowflake_query('sql') for _ in actual_generator: sf_conn_mock.assert_has_calls([ call.connect().__enter__().cursor().__enter__().execute('sql')]) @pytest.fixture def mock_mysql_query_result(): """Return mock result from mysql fetchmany """ return [ { 'column1': 'paulo', 'column2': 'kuong', 'column3': 'is', 'column4': 'awesome' }, { 'column1': 'john', 'column2': 'is', 'column3': 'awesome', 'column4': 'too' } ] @contextmanager def mock_mysql_cursor(): """Mock cursor context manager """ mock_cursor = MagicMock() mock_cursor.execute = lambda sql: sql mock_cursor.description = [ ('column1', 'blah'), ('column2', 'blah'), ('column3', 'blah'), ('column4', 'blah')] yield mock_cursor def test_mysql_query(monkeypatch, mock_mysql_query_result): """Test mysql_query method """ mock_connection = MagicMock() monkeypatch.setattr( mock_connection, 'cursor', MagicMock( return_value=(mock_mysql_cursor()))) monkeypatch.setattr( db, 'ResultIter', MagicMock(return_value=(mock_mysql_query_result))) monkeypatch.setattr( pymysql, 'connect', MagicMock(return_value=mock_connection)) results = db.mysql_query('select * from artists') assert list(results) == [ { 'column1': 'paulo', 'column2': 'kuong', 'column3': 'is', 'column4': 'awesome' }, { 'column1': 'john', 'column2': 'is', 'column3': 'awesome', 'column4': 'too' }]