"""Test snowflake_load_sales task.""" from unittest.mock import MagicMock, patch import pytest from lib import config from tasks.sales_ingest.snowflake_load_sales import load_sales_to_staging_table_task import_path = 'tasks.sales_ingest.snowflake_load_sales' @patch(f'{import_path}.RoyaltySnowflakeHook') @patch(f'{import_path}.insert_sales_to_nr_staging_table') @patch(f'{import_path}.get_abacus_event') def test_load_sales_to_nr_staging_table_task( mock_get_abacus_event, mock_snowflake_nr_load_sales_template, mock_hook, ): """Test using snowflake hook to load sales to NR staging table.""" dag_run = MagicMock() task_instance = MagicMock() mock_get_abacus_event.return_value.event_name = 'ingest_sales_nr' mock_get_abacus_event.return_value.target_id = 456 task_instance.xcom_pull.return_value = 'nr' kwargs = {'task_instance': task_instance} mock_snowflake_nr_load_sales_template.return_value.render.return_value = \ 'INSERT INTO NR STAGING' mock_hook.return_value.run.return_value = True load_sales_to_staging_table_task(dag_run, **kwargs) mock_get_abacus_event.assert_called_once_with(dag_run, **kwargs) mock_snowflake_nr_load_sales_template.return_value.render.assert_called_once_with( batch_id=456, schema=config.OWS_ENV.upper(), stage=config.ABACUS_TSV_STAGE, s3_path=f'extract-sales/nr/456/parts/' ) mock_hook.assert_called_once_with(snowflake_conn_id=config.SNOWFLAKE_CONN_NAME) task_instance.xcom_pull.assert_called_once_with( task_ids='determine_sales_ingest_type', key='sales_ingest_type' ) @patch(f'{import_path}.RoyaltySnowflakeHook') @patch(f'{import_path}.insert_sales_to_distro_staging_table') @patch(f'{import_path}.get_abacus_event') def test_load_sales_to_distro_staging_table_task( mock_get_abacus_event, mock_snowflake_distro_load_sales_template, mock_hook, ): """Test using snowflake hook to load sales to DISTRO staging table.""" dag_run = MagicMock() task_instance = MagicMock() mock_get_abacus_event.return_value.event_name = 'ingest_sales_distro' mock_get_abacus_event.return_value.target_id = 456 task_instance.xcom_pull.return_value = 'distro' kwargs = {'task_instance': task_instance} mock_snowflake_distro_load_sales_template.return_value.render.return_value = \ 'INSERT INTO DISTRO STAGING' mock_hook.return_value.run.return_value = True load_sales_to_staging_table_task(dag_run, **kwargs) mock_get_abacus_event.assert_called_once_with(dag_run, **kwargs) mock_snowflake_distro_load_sales_template.return_value \ .render.assert_called_once_with( batch_id=456, schema=config.OWS_ENV.upper(), stage=config.ABACUS_TSV_STAGE, s3_path=f'extract-sales/distro/456/parts/' ) mock_hook.assert_called_once_with(snowflake_conn_id=config.SNOWFLAKE_CONN_NAME) task_instance.xcom_pull.assert_called_once_with( task_ids='determine_sales_ingest_type', key='sales_ingest_type' ) @patch(f'{import_path}.RoyaltySnowflakeHook') @patch(f'{import_path}.insert_sales_to_distro_staging_table') @patch(f'{import_path}.get_abacus_event') def test_load_sales_to_distro_staging_table_error( mock_get_abacus_event, mock_snowflake_distro_load_sales_template, mock_hook, ): """Test throws an error if sales_ingest_type is not found.""" dag_run = MagicMock() task_instance = MagicMock() mock_get_abacus_event.return_value.event_name = 'ingest_sales_distro' mock_get_abacus_event.return_value.target_id = 456 task_instance.xcom_pull.return_value = None kwargs = {'task_instance': task_instance} with pytest.raises(Exception) as e: load_sales_to_staging_table_task(dag_run, **kwargs) assert str(e.value) == 'Unable to get the type for the sales insertion.' mock_get_abacus_event.assert_called_once_with(dag_run, **kwargs) mock_snowflake_distro_load_sales_template.assert_not_called() mock_hook.assert_not_called()