"""Unit tests for Apple Music Streams Snowflake SQL executor.""" from unittest.mock import ANY from unittest.mock import MagicMock from unittest.mock import patch import pytest from snowflake import connector from feed_ingestion.flows.apple_music_streams import config from feed_ingestion.flows.apple_music_streams import vendor_accounts from feed_ingestion.flows.apple_music_streams.snowflake_executor \ import AppleMusicStreams from tests.conftest import SubstringMatcher @pytest.fixture def mock_sql_loader(): """Return sql_loader mock.""" sql_loader_path = ( 'feed_ingestion.flows.apple_music_streams.snowflake_executor.' 'sql_loader') with patch(sql_loader_path) as sql_loader: yield sql_loader @pytest.fixture def mock_executor(sf_config_mock, monkeypatch): """Yield executor context.""" connect_mock = MagicMock() monkeypatch.setattr(connector, 'connect', connect_mock) executor = AppleMusicStreams(sf_config_mock) aws_params = { 'aws_key_id': 'test_id', 'aws_secret_key': 'test_secret', 'aws_token': 'test_token'} with patch.object(executor, 'execute', wraps=executor.execute) as \ executor.ex_mock, patch.object(executor, 'get_aws_params', return_value=aws_params): executor.fetchall = MagicMock() yield executor def test_create_temp_staging_raw_table(mock_executor): """Test create_temp_staging_raw_table method.""" for report_name in config.reports: # deprecated reports if report_name in ('amStreams', 'amNonRoyaltyStreams'): continue mock_executor.create_temp_staging_raw_table( temp_staging_raw_table='temp_table', report_name=report_name, date='2021-05-01', licensor='theorchard') mock_executor.ex_mock.assert_any_call( SubstringMatcher( containing=[ 'CREATE OR REPLACE TRANSIENT', 'test_db.test_schema.temp_table']), params={}) assert mock_executor.ex_mock.call_count == len(config.reports) - 2 def test_load_temp_staging_raw_table(mock_executor, aws_config_mock): """Test load_temp_staging_raw_table method.""" mock_executor.load_temp_staging_raw_table( temp_staging_raw_table='load_temp_staging_raw_table', aws={}, key_dir='s3://somepath', error_limit=1, date='2018-01-01', file_pattern='file_pattern') assert mock_executor.fetchall.call_count == 1 mock_executor.fetchall.assert_any_call( SubstringMatcher( containing=['COPY INTO', 'test_db.test_schema']), params={ 'aws_key_id': 'test_id', 'aws_secret_key': 'test_secret', 'aws_token': 'test_token', 's3_path': 's3://somepath', 'file_pattern': 'file_pattern'}) @pytest.mark.parametrize('report_name', ['amStreamsSummary', 'amContainer']) def test_staging_raw_location_consumer(mock_executor, report_name): """Consumer-reporting reports resolve to the consumer db/schema.""" db, schema = mock_executor.staging_raw_location(report_name) assert db == config.consumer_sf['db'] assert schema == config.consumer_sf['schema'] @pytest.mark.parametrize( 'report_name', ['amStreams', 'amContentDemographics', 'amSongs']) def test_staging_raw_location_default(mock_executor, report_name): """Reports without the flag stay on the flow's db/schema.""" db, schema = mock_executor.staging_raw_location(report_name) assert db == mock_executor.sf_config['db'] assert schema == mock_executor.sf_config['schema'] @pytest.mark.parametrize( 'report_name, staging_raw_table', [ ('amStreamsSummary', 'staging_raw_apple_music_summary_streams'), ('amContainer', 'staging_raw_apple_music_container')]) def test_load_staging_raw_table_to_consumer( mock_executor, report_name, staging_raw_table): """Migrated reports INSERT into the consumer db; temp tables stay local.""" mock_executor.load_staging_raw_table( '2025-06-01', staging_raw_table, 'temp_table', report_name=report_name, licensor='theorchard', processed_datetime='2025-06-01T00:00:00', filename='f', temp_table_amcontent='temp_amcontent', temp_table_amsubreference='temp_amsub', vendor='ORCHARD', apple_id_mapping_table='apple_id_mapping') mock_executor.ex_mock.assert_any_call( SubstringMatcher( containing=[ 'INSERT INTO {}.{}.{}'.format( config.consumer_sf['db'], config.consumer_sf['schema'], staging_raw_table), 'test_db.test_schema.temp_table']), params=ANY) def test_update_dimension_table(mock_executor, mock_get_vendors): """Test update_dimension_table method.""" mock_executor.fetchone = MagicMock() mock_executor.update_dimension_table('2021-02-01', 'dim_playlist') assert mock_executor.fetchone.call_count == 1 mock_executor.fetchone.assert_any_call( SubstringMatcher( containing=[ 'test_db.test_schema.dim_playlist', 'test_db.test_schema.staging_raw_apple_music_streams' ]), params={ 'date': '2021-02-01', 'storeid': 1, 'feedid': 4, 'consumer_db': config.consumer_sf['db'], 'consumer_schema': config.consumer_sf['schema'], 'theorchard_vendor_ids': ['80028967', '80029727'], 'sme_vendor_ids': ['80026921', '80030469'], 'awal_vendor_ids': ['80031998', ], }, dict_cursor=True) def test_update_dimension_table_after_switch(mock_executor, mock_get_vendors): """Test update_dimension_table method.""" mock_executor.fetchone = MagicMock() mock_executor.update_dimension_table('2021-04-14', 'dim_playlist') assert mock_executor.fetchone.call_count == 1 mock_executor.fetchone.assert_any_call( SubstringMatcher( containing=[ 'test_db.test_schema.dim_playlist', '{}.{}.staging_raw_apple_music_summary_streams'.format( config.consumer_sf['db'], config.consumer_sf['schema']) ]), params={ 'date': '2021-04-14', 'storeid': 1, 'feedid': 4, 'theorchard_vendor_ids': ['80028967', '80029727'], 'sme_vendor_ids': ['80026921', '80030469'], 'awal_vendor_ids': ['80031998', ], }, dict_cursor=True) def test_load_fact_data(mock_executor, mock_get_vendors): """Test load_fact_data method.""" for licensor in config.licensors: mock_executor.load_fact_data('2021-02-01', licensor=licensor) mock_executor.ex_mock.assert_any_call( SubstringMatcher( containing=[ 'test_db.test_schema.fact_analytics', 'test_db.test_schema.' 'staging_raw_apple_music_streams']), params={ 'storeid': 1, 'reportdate': '2021-02-01', 'feedid': 4, 'consumer_db': 'test_db', 'consumer_schema': 'test_schema', 'vendor_ids': vendor_accounts.get_vendors( '2017-11-16', licensor), 'licensor': licensor }) assert mock_executor.ex_mock.call_count == len(config.licensors) def test_load_fact_data_after_switch(mock_executor, mock_get_vendors): """Test load_fact_data method.""" for licensor in config.licensors: mock_executor.load_fact_data('2021-04-14', licensor=licensor) mock_executor.ex_mock.assert_any_call( SubstringMatcher( containing=[ 'test_db.test_schema.fact_analytics', 'test_db.test_schema.' 'staging_raw_apple_music_summary_streams']), params={ 'storeid': 1, 'reportdate': '2021-04-14', 'feedid': 4, 'consumer_db': config.consumer_sf['db'], 'consumer_schema': config.consumer_sf['schema'], 'vendor_ids': vendor_accounts.get_vendors( '2021-04-14', licensor), 'licensor': licensor }) assert mock_executor.ex_mock.call_count == len(config.licensors) def test_load_fact_error_data(mock_executor, mock_get_vendors): """Test load_fact_error_data method.""" for licensor in config.licensors: mock_executor.load_fact_error_data('2021-02-01', licensor=licensor) mock_executor.ex_mock.assert_any_call( SubstringMatcher( containing=[ 'test_db.test_schema.fact_analytics_error', 'test_db.test_schema.' 'staging_raw_apple_music_streams']), params={ 'storeid': 1, 'reportdate': '2021-02-01', 'feedid': 4, 'consumer_db': 'test_db', 'consumer_schema': 'test_schema', 'vendor_ids': vendor_accounts.get_vendors( '2021-02-01', licensor), 'licensor': licensor }) assert mock_executor.ex_mock.call_count == len(config.licensors) def test_load_fact_error_data_after_switch(mock_executor, mock_get_vendors): """Test load_fact_error_data method.""" for licensor in config.licensors: mock_executor.load_fact_error_data('2021-04-14', licensor=licensor) mock_executor.ex_mock.assert_any_call( SubstringMatcher( containing=[ 'test_db.test_schema.fact_analytics_error', 'test_db.test_schema.' 'staging_raw_apple_music_summary_streams']), params={ 'storeid': 1, 'reportdate': '2021-04-14', 'feedid': 4, 'consumer_db': config.consumer_sf['db'], 'consumer_schema': config.consumer_sf['schema'], 'vendor_ids': vendor_accounts.get_vendors( '2021-04-14', licensor), 'licensor': licensor }) assert mock_executor.ex_mock.call_count == len(config.licensors)