"""Test report_fetch module.""" import base64 import json from uuid import uuid4 import pytest from aws_kinesis_agg.aggregator import RecordAggregator from src.common.logic import report_fetch from src.common.logic.report_fetch import encode_records from src.common.models.metadata import ARSyncEvent from src.common.models.metadata import InconsistencyType from src.common.models.metadata import ReleaseArtistTable from src.common.models.metadata import TrackArtistTable from src.common.models.metadata import TrackTable from src.common.models.metadata import TrackWriterTable class TestFilterReportType: """Test filter_report_type function.""" @pytest.mark.parametrize( 'inconsistency_type, expected_result', [ ( InconsistencyType.BAD_PRODUCT_PARTICIPATIONS, 'product_participation_inconsistency_report' ), ( InconsistencyType.BAD_TRACK_PARTICIPATIONS, 'track_participation_inconsistency_report' ), ( InconsistencyType.BAD_LABEL_SOUND_RECORDINGS, 'label_sound_recording_inconsistency_report' ), ] ) def test_filter_report_type( self, inconsistency_reports, inconsistency_type, expected_result, product_participation_inconsistency_report, track_participation_inconsistency_report, label_sound_recording_inconsistency_report ): """Test filter_report_type function.""" fixture_map = { 'product_participation_inconsistency_report': product_participation_inconsistency_report, 'track_participation_inconsistency_report': track_participation_inconsistency_report, 'label_sound_recording_inconsistency_report': label_sound_recording_inconsistency_report, } filtered_report = report_fetch.filter_report_type( inconsistency_reports, inconsistency_type) assert filtered_report == fixture_map[expected_result] def test_filter_report_type_failure(self, product_participation_inconsistency_report_only): """Test filter_report_type function.""" with pytest.raises(ValueError): report_fetch.filter_report_type( product_participation_inconsistency_report_only, InconsistencyType.BAD_TRACK_PARTICIPATIONS) class TestGetIdsPerTable: """Test get_ids_per_table function.""" @pytest.mark.parametrize( 'inconsistency_report, table, expected_result', [ ( 'product_participation_inconsistency_report', ReleaseArtistTable, 'inconsistent_release_artist_id' ), ( 'track_participation_inconsistency_report', TrackArtistTable, 'inconsistent_track_artist_id' ), ( 'track_participation_inconsistency_report', TrackWriterTable, 'inconsistent_track_writer_id' ), ( 'label_sound_recording_inconsistency_report', TrackTable, 'inconsistent_track_id' ), ] ) def test_get_ids_per_table( self, inconsistency_report, table, expected_result, product_participation_inconsistency_report, track_participation_inconsistency_report, label_sound_recording_inconsistency_report, inconsistent_release_artist_id, inconsistent_track_artist_id, inconsistent_track_writer_id, inconsistent_track_id ): """Test get_ids_per_table function.""" fixture_map = { 'product_participation_inconsistency_report': product_participation_inconsistency_report, 'track_participation_inconsistency_report': track_participation_inconsistency_report, 'label_sound_recording_inconsistency_report': label_sound_recording_inconsistency_report, 'inconsistent_release_artist_id': inconsistent_release_artist_id, 'inconsistent_track_artist_id': inconsistent_track_artist_id, 'inconsistent_track_writer_id': inconsistent_track_writer_id, 'inconsistent_track_id': inconsistent_track_id, } result = report_fetch.get_ids_per_table(fixture_map[inconsistency_report], table) assert result == [fixture_map[expected_result]] class TestEncodeRecords: """Test encode_records function.""" @pytest.fixture def ar_sync_events(self): """List of ARSyncEvent objects.""" return [ ARSyncEvent( data={'key1': 'value1'}, table='my_table' ), ARSyncEvent( data={'key2': 'value2'}, table='my_table' ), ] @pytest.fixture def uuids(self, ar_sync_events): """List of uuids for AR events.""" return [ str(uuid4()) for _ in ar_sync_events ] @pytest.fixture def uuid4_mock(self, mocker, ar_sync_events, uuids): """Mock uuid4 to return uuids.""" return mocker.patch.object(report_fetch, 'uuid4', side_effect=uuids) @pytest.fixture def ar_sync_events_encoded(self, ar_sync_events, uuids): """Fixture for encoded ARSyncEvent records.""" aggregator = RecordAggregator() for uuid, record in zip(uuids, ar_sync_events): aggregator.add_user_record( uuid, json.dumps(record.dict(), default=str) ) agg_content = aggregator.current_record.get_contents()[-1] encoded_records = base64.b64encode(agg_content).decode() return encoded_records def test_encode_records(self, ar_sync_events, ar_sync_events_encoded, uuid4_mock): """Test encode_records function.""" result = encode_records(ar_sync_events) assert result == ar_sync_events_encoded