"""Test update_sales_file_metadata task.""" from decimal import Decimal from unittest.mock import MagicMock, patch from lib import config from tasks.sales_get_eligible import update_sales_file_metadata @patch('tasks.sales_get_eligible.update_sales_file_metadata.get_event_from_params') @patch('tasks.sales_get_eligible.update_sales_file_metadata.ows') @patch( 'tasks.sales_get_eligible.update_sales_file_metadata._query_metadata_from_snowflake' ) @patch('tasks.sales_get_eligible.update_sales_file_metadata._format_update_params') def test_update_sales_metadata( mock_format_update_params, mock_query_snowflake, mock_ows, mock_get_event, mock_sales_get_eligible_event, mock_sales_get_eligible_dag_run ): """Test main update_sales_metadata task method.""" sales_file_id = mock_sales_get_eligible_event.get('target_id') mock_metadata = (Decimal('123.456'), 789) mock_update_params = { 'amount_usd': '123.46', 'row_count': 789 } mock_get_event.return_value.target_id = sales_file_id mock_query_snowflake.return_value = mock_metadata mock_format_update_params.return_value = mock_update_params update_sales_file_metadata.update_sales_metadata(mock_sales_get_eligible_dag_run) mock_get_event.assert_called_once_with(mock_sales_get_eligible_dag_run) mock_query_snowflake.assert_called_once_with(sales_file_id) mock_format_update_params.assert_called_once_with(mock_metadata) mock_ows.update_sales_file.assert_called_once_with( sales_file_id, **mock_update_params ) def test_format_update_params(): """Test _format_update_params returns a rounded total.""" mock_metadata = (Decimal('123.456'), 789) result = update_sales_file_metadata._format_update_params(mock_metadata) assert result.get('amount_usd') == '123.46' assert result.get('row_count') == mock_metadata[1] @patch('tasks.sales_get_eligible.update_sales_file_metadata.' 'ows.get_accounting_period_details') @patch('tasks.sales_get_eligible.update_sales_file_metadata.' 'ows.get_sales_file_details') @patch('tasks.sales_get_eligible.update_sales_file_metadata.RoyaltySnowflakeHook') @patch( 'tasks.sales_get_eligible.update_sales_file_metadata.get_distro_sales_file_metadata' ) def test_query_distribution_metadata_from_snowflake( mock_distribution_template: MagicMock, mock_hook: MagicMock, mock_ows_get_sales_file_details: MagicMock, mock_ows_get_accounting_period_detail: MagicMock, ): """Test using snowflake hook to query for distribution sales metadata.""" mock_sales_file_response = { 'sales_file_id': 123, 'accounting_period_id': 1 } mock_accounting_period_response = { 'accounting_period_id': 1, 'accounting_period_status': 'open', 'contract_type': 'distribution' } mock_ows_get_sales_file_details.return_value = mock_sales_file_response mock_ows_get_accounting_period_detail.return_value = mock_accounting_period_response sales_file_id = 123 mock_metadata = (Decimal('123.456'), 789) mock_query = 'SELECT * FROM SF;' mock_distribution_template.return_value.render.return_value = mock_query mock_hook.return_value.get_first.return_value = mock_metadata result = update_sales_file_metadata._query_metadata_from_snowflake(sales_file_id) assert result == mock_metadata mock_hook.assert_called_once_with(snowflake_conn_id=config.SNOWFLAKE_CONN_NAME) mock_distribution_template.assert_called_once() mock_distribution_template.return_value.render.assert_called_once_with( env=config.OWS_ENV, sales_file_id=sales_file_id ) mock_hook.return_value.get_first.assert_called_once_with(mock_query) mock_ows_get_sales_file_details.assert_called_once_with(sales_file_id) mock_ows_get_accounting_period_detail.assert_called_once_with( mock_sales_file_response['accounting_period_id'], ) @patch('tasks.sales_get_eligible.update_sales_file_metadata.' 'ows.get_accounting_period_details') @patch('tasks.sales_get_eligible.update_sales_file_metadata.' 'ows.get_sales_file_details') @patch('tasks.sales_get_eligible.update_sales_file_metadata.RoyaltySnowflakeHook') @patch( 'tasks.sales_get_eligible.update_sales_file_metadata.get_nr_sales_file_metadata' ) def test_query_nr_metadata_from_snowflake( mock_nr_template: MagicMock, mock_hook: MagicMock, mock_ows_get_sales_file_details: MagicMock, mock_ows_get_accounting_period_details: MagicMock, ): """Test using snowflake hook to query for NR sales metadata.""" mock_sales_file_response = { 'sales_file_id': 123, 'accounting_period_id': 2 } mock_accounting_period_response = { 'accounting_period_id': 2, 'accounting_period_status': 'open', 'contract_type': 'neighbouring_rights' } mock_ows_get_sales_file_details.return_value = mock_sales_file_response mock_ows_get_accounting_period_details.return_value =\ mock_accounting_period_response sales_file_id = 123 mock_metadata = (Decimal('123.456'), 789) mock_query = 'SELECT * FROM SF;' mock_nr_template.return_value.render.return_value = mock_query mock_hook.return_value.get_first.return_value = mock_metadata result = update_sales_file_metadata._query_metadata_from_snowflake(sales_file_id) assert result == mock_metadata mock_hook.assert_called_once_with(snowflake_conn_id=config.SNOWFLAKE_CONN_NAME) mock_nr_template.assert_called_once() mock_nr_template.return_value.render.assert_called_once_with( env=config.OWS_ENV, sales_file_id=sales_file_id ) mock_hook.return_value.get_first.assert_called_once_with(mock_query) mock_ows_get_sales_file_details.assert_called_once_with(sales_file_id) mock_ows_get_accounting_period_details.assert_called_once_with( mock_sales_file_response['accounting_period_id'])