"""Trigger tests.""" from collections import namedtuple from unittest.mock import MagicMock, patch import pytest from trigger import trigger from trigger.constants import SplitType from trigger.schemas import SplitRow TriggerReportGenerationTest = namedtuple( "TriggerReportGenerationTest", ["skip_update", "trigger_type", "invalid_collaborators"], ) @pytest.mark.parametrize( TriggerReportGenerationTest._fields, [ TriggerReportGenerationTest( skip_update=False, trigger_type="AUTO", invalid_collaborators=[], ), TriggerReportGenerationTest( skip_update=True, trigger_type="MANUAL", invalid_collaborators=[], ), TriggerReportGenerationTest( skip_update=False, trigger_type="MANUAL", invalid_collaborators=[], ), TriggerReportGenerationTest( skip_update=True, trigger_type="AUTO", invalid_collaborators=[], ), TriggerReportGenerationTest( skip_update=False, trigger_type="AUTO", invalid_collaborators=["123"], ), TriggerReportGenerationTest( skip_update=True, trigger_type="AUTO", invalid_collaborators=["123"], ), ], ) @patch("trigger.trigger.rds.get_report_run") @patch("trigger.trigger.validate_splits_for_report_run") @patch("trigger.trigger.sync_splits") @patch("trigger.trigger.sync_collaborators") @patch("trigger.trigger.snowflake.trigger_reports_query_async") @patch("trigger.trigger.rds.update_reports_status") @patch("trigger.trigger.rds.update_snowflake_query_id") def test_trigger_report_generation( mock_update_snowflake_query_id, mock_update_reports_status, mock_trigger_reports_query_async, mock_sync_collaborators, mock_sync_splits, mock_validate_splits_for_report_run, mock_get_report_run, skip_update, trigger_type, invalid_collaborators, ): report_run_uuid = "123" period_ids = ["1", "2", "3"] reports_query_id = "567" mock_get_report_run.return_value = (period_ids, trigger_type) mock_trigger_reports_query_async.return_value = reports_query_id mock_validate_splits_for_report_run.return_value = invalid_collaborators trigger.trigger_report_generation(report_run_uuid, skip_update) mock_get_report_run.assert_called_with(report_run_uuid) mock_sync_splits.assert_called_with(report_run_uuid, invalid_collaborators) mock_sync_collaborators.assert_called_with(report_run_uuid, invalid_collaborators) mock_trigger_reports_query_async.assert_called_with(report_run_uuid, period_ids) if skip_update: mock_update_snowflake_query_id.assert_not_called() mock_update_reports_status.assert_not_called() else: mock_update_snowflake_query_id.assert_called_with(report_run_uuid, reports_query_id) if len(invalid_collaborators) > 0: mock_update_reports_status.assert_called_with( report_run_uuid, "ERROR", with_collaborator_ids=invalid_collaborators ) @patch("trigger.trigger.csv.write_temporary_csv") @patch("trigger.trigger.snowflake.create_splits_table") @patch("trigger.trigger.snowflake.upload_file_to_table_stage") @patch("trigger.trigger.snowflake.copy_into_table_from_stage") @patch("os.remove") def test_sync_splits( mock_remove, mock_copy_into_table_from_stage, mock_upload_file_to_table_stage, mock_create_splits_table, mock_write_temporary_csv, ): report_run_uuid = "123" invalid_collaborators = [1] temp_file_path = "/path/to/temp_file.csv" level1_rows = [SplitRow(1, SplitType.TRACK, "tuid-A", 0.5, 42, "NET")] level2_rows = [SplitRow(-1, SplitType.SUBACCOUNT, "tuid-B", 0.3, 42, "NET")] mock_level1 = MagicMock() mock_level1.to_track_splits.return_value = (level1_rows, set()) mock_level2 = MagicMock() mock_level2.to_track_splits.return_value = (level2_rows, set()) mock_write_temporary_csv.return_value = temp_file_path with patch.object(trigger, "SPLIT_LEVEL_HIERARCHY", [mock_level1, mock_level2]): trigger.sync_splits(report_run_uuid, invalid_collaborators) mock_level1.get_splits.assert_called_with(report_run_uuid, invalid_collaborators) mock_level2.get_splits.assert_called_with(report_run_uuid, invalid_collaborators) mock_write_temporary_csv.assert_called_with(level1_rows + level2_rows) mock_create_splits_table.assert_called() mock_upload_file_to_table_stage.assert_called_with( temp_file_path, trigger.config.SNOWFLAKE_OBJECTS["temp_split"] ) mock_copy_into_table_from_stage.assert_called_with( trigger.config.SNOWFLAKE_OBJECTS["temp_split"] ) mock_remove.assert_called_with(temp_file_path) @patch("trigger.trigger.rds.get_collaborators_for_report_run") @patch("trigger.trigger.csv.write_temporary_csv") @patch("trigger.trigger.snowflake.create_collaborators_table") @patch("trigger.trigger.snowflake.upload_file_to_table_stage") @patch("trigger.trigger.snowflake.copy_into_table_from_stage") @patch("os.remove") def test_sync_collaborators( mock_remove, mock_copy_into_table_from_stage, mock_upload_file_to_table_stage, mock_create_collaborators_table, mock_write_temporary_csv, mock_get_collaborators_for_report_run, ): report_run_uuid = "123" collaborators = ["collaborator1", "collaborator2", "collaborator3"] temp_file_path = "/path/to/temp_file.csv" invalid_collaborators = [1] mock_get_collaborators_for_report_run.return_value = collaborators mock_write_temporary_csv.return_value = temp_file_path trigger.sync_collaborators(report_run_uuid, invalid_collaborators) mock_create_collaborators_table.assert_called() mock_get_collaborators_for_report_run.assert_called_with(report_run_uuid, invalid_collaborators) mock_write_temporary_csv.assert_called_with(collaborators) mock_upload_file_to_table_stage.assert_called_with( temp_file_path, trigger.config.SNOWFLAKE_OBJECTS["temp_collaborator"] ) mock_copy_into_table_from_stage.assert_called_with( trigger.config.SNOWFLAKE_OBJECTS["temp_collaborator"] ) mock_remove.assert_called_with(temp_file_path)