"""Unit tests for snowflake_calculate_totals task.""" from unittest.mock import patch from lib import config from tasks.accounting_run_calculate_nr.snowflake_calculate_totals import \ calculate_run_totals_task import_path = 'tasks.accounting_run_calculate_nr.snowflake_calculate_totals' @patch(f'{import_path}.RoyaltySnowflakeHook') @patch(f'{import_path}.shared_helpers') @patch(f'{import_path}.helpers') @patch(f'{import_path}.calculate_totals_nr_template') def test_calculate_run_totals_task_success( mock_template, mock_helpers, mock_shared_helpers, mock_hook, mock_accounting_run_calculate_nr_dag_run ): """Test task that calculates NR run totals and stages them in snowflake.""" accounting_run_id = 123 accounting_period = { 'statement_period_id': 456 } accounting_run = { 'accounting_run_id': accounting_run_id, 'run_controller_name': 'Test NR Run' } mock_helpers.get_event_from_params.return_value.target_id = accounting_run_id mock_shared_helpers.get_event_records.return_value = \ accounting_period, accounting_run mock_template.return_value.render.return_value = \ 'CALCULATE RUN TOTALS AND INSERT RESULTS INTO ACCOUNTING_RUN_RESULTS_NR_STAGING' mock_hook.return_value.run.return_value = True calculate_run_totals_task(mock_accounting_run_calculate_nr_dag_run) mock_helpers.get_event_from_params.assert_called_once_with( mock_accounting_run_calculate_nr_dag_run ) mock_shared_helpers.get_event_records.assert_called_once_with(accounting_run_id) mock_hook.assert_called_once_with(snowflake_conn_id=config.SNOWFLAKE_CONN_NAME) mock_template.return_value.render.assert_called_once_with( accounting_run_id=accounting_run_id, accounting_run_name=accounting_run.get('run_controller_name'), schema=config.OWS_ENV, statement_period_id=accounting_period.get('statement_period_id') )