"""Unit tests for snowflake_rollover_unallocated_sales task.""" from unittest.mock import patch from lib import config from lib.constants import CONTRACT_TYPES from tasks.accounting_period_close.snowflake_rollover_unallocated_sales \ import snowflake_rollover_unallocated_sales_task import_path = 'tasks.accounting_period_close.snowflake_rollover_unallocated_sales' @patch(f'{import_path}.RoyaltySnowflakeHook') @patch(f'{import_path}.get_event_from_params') @patch(f'{import_path}.get_accounting_period_details') @patch(f'{import_path}.insert_stmt_db_sales_distro_staging') @patch(f'{import_path}.insert_stmt_db_sales_nr_staging') def test_snowflake_rollover_unallocated_sales_task_distro( mock_nr_template, mock_distro_template, mock_get_accounting_period_details, mock_get_event_from_params_helper, mock_snowflake_hook, mock_accounting_period_close_dag_run, mock_accounting_period_close_event ): """Unit test for rolling over unallocated distribution sales.""" accounting_period_id = mock_accounting_period_close_event['target_id'] statement_period_id = 310 mock_accounting_period = { 'accounting_period_id': accounting_period_id, 'contract_type': CONTRACT_TYPES.DISTRIBUTION, 'statement_period_id': statement_period_id } mock_get_event_from_params_helper.return_value.target_id = accounting_period_id mock_get_accounting_period_details.return_value = mock_accounting_period mock_distro_template.return_value.render.return_value = 'ROLLOVER DISTRO SALES' mock_snowflake_hook.return_value.run.return_value = True snowflake_rollover_unallocated_sales_task( mock_accounting_period_close_dag_run ) mock_get_event_from_params_helper.assert_called_once_with( mock_accounting_period_close_dag_run ) mock_get_accounting_period_details.assert_called_once_with(accounting_period_id) mock_snowflake_hook.assert_called_once_with( snowflake_conn_id=config.SNOWFLAKE_CONN_NAME ) mock_distro_template.return_value.render.assert_called_once_with( accounting_period_id=accounting_period_id, statement_period_id=statement_period_id, schema=config.OWS_ENV ) mock_nr_template.assert_not_called() mock_snowflake_hook.return_value.run.assert_called_once() @patch(f'{import_path}.RoyaltySnowflakeHook') @patch(f'{import_path}.get_event_from_params') @patch(f'{import_path}.get_accounting_period_details') @patch(f'{import_path}.insert_stmt_db_sales_distro_staging') @patch(f'{import_path}.insert_stmt_db_sales_nr_staging') def test_snowflake_rollover_unallocated_sales_task_nr( mock_nr_template, mock_distro_template, mock_get_accounting_period_details, mock_get_event_from_params_helper, mock_snowflake_hook, mock_accounting_period_close_dag_run, mock_accounting_period_close_event ): """Unit test for rolling over unallocated neighbouring_rights sales.""" accounting_period_id = mock_accounting_period_close_event['target_id'] statement_period_id = 310 mock_accounting_period = { 'accounting_period_id': accounting_period_id, 'contract_type': CONTRACT_TYPES.NEIGHBOURING_RIGHTS, 'statement_period_id': statement_period_id } mock_get_event_from_params_helper.return_value.target_id = accounting_period_id mock_get_accounting_period_details.return_value = mock_accounting_period mock_nr_template.return_value.render.return_value = 'ROLLOVER NR SALES' mock_snowflake_hook.return_value.run.return_value = True snowflake_rollover_unallocated_sales_task( mock_accounting_period_close_dag_run ) mock_get_event_from_params_helper.assert_called_once_with( mock_accounting_period_close_dag_run ) mock_get_accounting_period_details.assert_called_once_with(accounting_period_id) mock_snowflake_hook.assert_called_once_with( snowflake_conn_id=config.SNOWFLAKE_CONN_NAME ) mock_nr_template.return_value.render.assert_called_once_with( accounting_period_id=accounting_period_id, statement_period_id=statement_period_id, schema=config.OWS_ENV ) mock_distro_template.assert_not_called() mock_snowflake_hook.return_value.run.assert_called_once() @patch(f'{import_path}.RoyaltySnowflakeHook') @patch(f'{import_path}.get_event_from_params') @patch(f'{import_path}.get_accounting_period_details') @patch(f'{import_path}.insert_stmt_db_sales_distro_staging') @patch(f'{import_path}.insert_stmt_db_sales_nr_staging') def test_snowflake_rollover_unallocated_sales_task_other( mock_nr_template, mock_distro_template, mock_get_accounting_period_details, mock_get_event_from_params_helper, mock_snowflake_hook, mock_accounting_period_close_dag_run, mock_accounting_period_close_event ): """Unit test for rolling over unallocated sales with unknown contract_type.""" accounting_period_id = mock_accounting_period_close_event['target_id'] statement_period_id = 310 mock_accounting_period = { 'accounting_period_id': accounting_period_id, 'contract_type': 'unknown', 'statement_period_id': statement_period_id } mock_get_event_from_params_helper.return_value.target_id = accounting_period_id mock_get_accounting_period_details.return_value = mock_accounting_period mock_snowflake_hook.return_value.run.return_value = True snowflake_rollover_unallocated_sales_task( mock_accounting_period_close_dag_run ) mock_get_event_from_params_helper.assert_called_once_with( mock_accounting_period_close_dag_run ) mock_get_accounting_period_details.assert_called_once_with(accounting_period_id) mock_snowflake_hook.assert_called_once_with( snowflake_conn_id=config.SNOWFLAKE_CONN_NAME ) mock_nr_template.assert_not_called() mock_distro_template.assert_not_called() mock_snowflake_hook.return_value.run.assert_not_called()