"""Test handle_artwork handler.""" from io import BytesIO import json from unittest.mock import MagicMock from unittest.mock import patch from common.audio_exceptions import AudioValidationFailure from constants.audio import BAD_FORMATS from constants.audio import CD_AUDIO from constants.audio import DAT_AUDIO from constants.audio import GOOD_24_SAMPLERATES from constants.audio import GOOD_AUDIO_SPECS from constants.audio import GOOD_CHANNELS from constants.audio import GOOD_SUBTYPES from constants.audio import NEW_ROOT_SUBFOLDER from constants.audio import PCM_16_SUBTYPE from constants.audio import PCM_24_SUBTYPE from constants.audio import PRO_AUDIO from constants.audio import WAVE_PCM_FORMAT import pytest from src.index import handler from util.audio_utils import get_audio_stream_from_s3 from util.audio_utils import process_audio from util.audio_utils import save_audio_to_s3 from util.audio_utils import validate_file_key @patch('util.audio_utils.boto3.client') @patch('util.audio_utils.BytesIO') def test_get_audio_stream_from_s3(mock_bytes, mock_boto3_client): """Test getting audio data from S3.""" mock_boto3_client_return = MagicMock() mock_boto3_client.return_value = mock_boto3_client_return mock_get_object = MagicMock() mock_boto3_client_return.get_object = mock_get_object mock_get_object_body = MagicMock() mock_get_object_body.read.return_value = b'AudioData' mock_get_object.return_value = {'Body': mock_get_object_body} mock_bytes_object = MagicMock(name='BytesIO') mock_bytes.return_value = mock_bytes_object mock_key = 'key.jpg' mock_bucket = 'bucket' output = get_audio_stream_from_s3(mock_bucket, mock_key) mock_boto3_client.assert_called_once_with('s3') mock_get_object.assert_called_once_with( Bucket=mock_bucket, Key=mock_key ) assert output == mock_bytes_object @pytest.mark.parametrize('test_key, test_passes', [ ('file.wav', True), ('file.ogg', False), ('file.flac', True), ('.file', False), ]) def test_validate_file_key(test_key, test_passes): """Test validating a file is a wav.""" bucket = 'test_bucket' if test_passes: try: validate_file_key(bucket, test_key) except AudioValidationFailure: assert False else: with pytest.raises(AudioValidationFailure): validate_file_key(bucket, test_key) @patch('util.audio_utils.BytesIO') @patch('util.audio_utils.sf') def test_process_audio(mock_soundfile, mock_bytes): """Test processing wav stream.""" mock_out_buffer = MagicMock(name='out_buffer') mock_bytes.return_value = mock_out_buffer mock_wave_file = MagicMock( name='wav_file', samplerate=44000, endian='FILE', format='WAV', subtype='PCM_24' ) mock_in_buffer = MagicMock(name='in_buffer') mock_wave_file.read.return_value = mock_in_buffer process_audio(mock_wave_file) assert mock_soundfile.write.called_once_with( mock_out_buffer, mock_in_buffer, samplerate=mock_wave_file.samplerate, endian=mock_wave_file.endian, format=mock_wave_file.format, subtype=mock_wave_file.subtype ) assert mock_wave_file.read.called_once() assert mock_wave_file.close.called_once() assert mock_out_buffer.seek.called_once() @pytest.mark.parametrize('test_wav, test_bytes', [ (MagicMock(spec=BytesIO), True), (None, False) ]) @patch('util.audio_utils.boto3.client') @patch('util.audio_utils.get_current_logger') def test_save_wav_to_s3( mock_logger, mock_boto3_client, test_wav, test_bytes): """Test saving an audio asset.""" mock_boto3_client_return = MagicMock() mock_boto3_client.return_value = mock_boto3_client_return mock_boto3_client_return.head_object.return_value = {} mock_boto3_client_return.put_object.return_value = { 'ResponseMetadata': { 'HTTPStatusCode': 200 } } mock_boto3_client_return.copy_object.return_value = { 'ResponseMetadata': { 'HTTPStatusCode': 200 } } mock_key = 'key.wav' mock_target_bucket = 'target_bucket' mock_new_root = NEW_ROOT_SUBFOLDER mock_bucket = 'bucket' mock_target_key = mock_new_root + '/' + mock_bucket + '/' + mock_key mock_correlation_id = 'correlation_id' mock_metadata = {} mock_buffer_read = MagicMock() if test_bytes: test_wav.read.return_value = mock_buffer_read save_audio_to_s3( test_wav, mock_target_bucket, mock_target_key, mock_bucket, mock_key, mock_correlation_id) if test_bytes: mock_boto3_client_return.put_object.assert_called_once_with( Body=mock_buffer_read, Key=mock_target_key, Bucket=mock_target_bucket, ContentType='audio/x-wav', Metadata=mock_metadata ) else: mock_boto3_client_return.copy_object.assert_called_once_with( Key=mock_target_key, Bucket=mock_target_bucket, CopySource={'Bucket': mock_bucket, 'Key': mock_key}, Metadata=mock_metadata, MetadataDirective='REPLACE', TaggingDirective='REPLACE', ContentType='audio/x-wav' ) @patch('src.index.config.STAGING_S3_BUCKET', 'staging_bucket') @patch('src.index.sf') @patch('src.index.config.get_current_logger') @patch('src.index.save_audio_to_s3') @patch('src.index.process_audio') @patch('src.index.validate_file_key') @patch('src.index.get_audio_stream_from_s3') @pytest.mark.parametrize( 'audio_spec', ['pcm_16_48_st', 'pcm_24_96_st', 'flac_16_48_st', 'pcm_32_1764_st', 'wavex_24_96_st', 'aiff_24_96_st', 'pcm_16_12_st', 'pcm_16_96_st', 'pcm_24_12_st', 'pcm_16_48_mo', 'pcm_float_96_st'], ) def test_handler( mock_get_audio_stream_from_s3, mock_validate_file_key, mock_process_audio, mock_save_audio_to_s3, mock_current_logger, mock_soundfile, context_json_audio, audio_spec, request): """Test handler.""" mock_logger = MagicMock(name='info') mock_current_logger.return_value = mock_logger # TODO: parameterize this test for different cases format_params = request.getfixturevalue(audio_spec) mock_soundfile_object = MagicMock( name='audio_file', **format_params ) mock_soundfile.SoundFile.return_value = mock_soundfile_object mock_bucket = 'test_bucket' mock_target_bucket = 'staging_bucket' mock_key = 'test_path/test_filename.wav' mock_new_root = NEW_ROOT_SUBFOLDER mock_target_key = mock_new_root + '/' + mock_bucket + '/' + mock_key mock_correlation_id = 'test_correlation_id' mock_channels = format_params.get('channels') # Force to wav response if format_params.get('format') in BAD_FORMATS: mock_format = WAVE_PCM_FORMAT else: mock_format = format_params.get('format') # Acceptable subtypes response if format_params.get('subtype') in GOOD_SUBTYPES: mock_subtype = format_params.get('subtype') else: mock_subtype = PCM_24_SUBTYPE # Acceptable sample rates response test_samplerate = format_params.get('samplerate') if test_samplerate in GOOD_24_SAMPLERATES: mock_samplerate = test_samplerate # Up-sample bit depth response if mock_subtype == PCM_16_SUBTYPE and mock_samplerate > DAT_AUDIO: mock_subtype = PCM_24_SUBTYPE elif test_samplerate > PRO_AUDIO: # Bad samplerate response mock_samplerate = PRO_AUDIO else: mock_samplerate = CD_AUDIO mock_audio_stream = MagicMock(name='get_audio') mock_get_audio_stream_from_s3.return_value = mock_audio_stream mock_processed_stream = MagicMock(name='process_stream', spec=BytesIO) mock_process_audio.return_value = mock_processed_stream if mock_channels not in GOOD_CHANNELS: with pytest.raises(Exception): handler(json.loads(context_json_audio), None) else: handler(json.loads(context_json_audio), None) mock_validate_file_key.assert_called_once_with(mock_bucket, mock_key) mock_get_audio_stream_from_s3.assert_called_once_with( mock_bucket, mock_key) if tuple(format_params.values()) in GOOD_AUDIO_SPECS: mock_process_audio.assert_not_called() else: mock_process_audio.assert_called_once_with( mock_soundfile_object, audio_format=mock_format, channels=mock_channels, samplerate=mock_samplerate, subtype=mock_subtype ) mock_save_audio_to_s3.assert_called_once_with( mock_processed_stream, mock_target_bucket, mock_target_key, mock_bucket, mock_key, mock_correlation_id ) @patch('src.index.config.STAGING_S3_BUCKET', 'staging_bucket') @patch('src.index.config.get_current_logger') @patch('src.index.validate_file_key') def test_handler_fails( mock_validate_file_key, mock_current_logger, context_json_audio): """Test handler fails.""" mock_logger = MagicMock(name='info') mock_current_logger.return_value = mock_logger mock_validate_file_key.side_effect = AudioValidationFailure( 'Get Exception') with pytest.raises(AudioValidationFailure): handler(json.loads(context_json_audio), None) mock_logger.error.assert_called_once()