"""Unit tests for snowflake_calculate_mechanical_royalty_amounts_task.""" from unittest.mock import patch from lib import config from tasks.accounting_period_mechanicals.snowflake_calculate_mech_royalty_amounts \ import snowflake_calculate_mechanical_royalty_amounts_task import_path = \ 'tasks.accounting_period_mechanicals.snowflake_calculate_mech_royalty_amounts' @patch(f'{import_path}.RoyaltySnowflakeHook') @patch(f'{import_path}.helpers') @patch(f'{import_path}.calculate_mechanical_royalty_amounts') def test_snowflake_calculate_mechanical_royalty_amounts_task( mock_template, mock_helpers, mock_hook, mock_accounting_period_mechanicals_event, mock_accounting_period_mechanicals_dag_run ): """Test task that calculates mechanical deductions (royalty_amounts).""" accounting_period_id = mock_accounting_period_mechanicals_event['target_id'] mock_helpers.get_event_from_params.return_value.target_id = accounting_period_id mock_template.return_value.render.return_value = \ 'SUM DETAIL ROYALTY_AMOUNTS AND UPDATE SNAPSHOT_MECHANICAL_TRANSACTIONS' mock_hook.return_value.run.return_value = True snowflake_calculate_mechanical_royalty_amounts_task( mock_accounting_period_mechanicals_dag_run ) mock_helpers.get_event_from_params.assert_called_once_with( mock_accounting_period_mechanicals_dag_run ) mock_hook.assert_called_once_with(snowflake_conn_id=config.SNOWFLAKE_CONN_NAME) mock_template.return_value.render.assert_called_once_with( accounting_period_id=accounting_period_id, schema=config.OWS_ENV )