"""Tests for track_display_artist MySQL queries.""" from contextlib import nullcontext from pytest_mock import MockerFixture from contributor.queries.mysql.track_display_artist import ( TrackDisplayArtist, get_by_track_ids, ) class TestTrackDisplayArtistModel: def test_model_has_correct_tablename(self) -> None: assert TrackDisplayArtist.__tablename__ == "track_display_artists" def test_model_instantiation_with_all_fields(self) -> None: track_display_artist = TrackDisplayArtist( id=1, track_id=100, contributor_id=50, display_artist_role="primary", sequence_number=1, ) assert track_display_artist.id == 1 assert track_display_artist.track_id == 100 assert track_display_artist.contributor_id == 50 assert track_display_artist.display_artist_role == "primary" assert track_display_artist.sequence_number == 1 def test_model_instantiation_with_required_fields_only(self) -> None: track_display_artist = TrackDisplayArtist( track_id=100, contributor_id=50, display_artist_role="featuring", ) assert track_display_artist.track_id == 100 assert track_display_artist.contributor_id == 50 assert track_display_artist.display_artist_role == "featuring" def test_model_instantiation_with_sequence_number(self) -> None: track_display_artist = TrackDisplayArtist( track_id=100, contributor_id=50, display_artist_role="feature_to_primary", sequence_number=5, ) assert track_display_artist.sequence_number == 5 def test_model_has_all_required_attributes(self) -> None: required_fields = [ "id", "track_id", "contributor_id", "display_artist_role", "sequence_number", ] for field in required_fields: assert hasattr(TrackDisplayArtist, field) class TestGetByTrackIds: def test_returns_display_artists_grouped_by_track_id( self, mocker: MockerFixture ) -> None: session = mocker.Mock() execute_result = mocker.Mock() session.execute.return_value = execute_result execute_result.all.return_value = [ (101, "primary", 1, 10, "11111111-2222-4333-8444-555555555555"), (101, "featuring", 2, 11, "66666666-7777-4888-8999-aaaaaaaaaaaa"), (202, "primary", 1, 12, "bbbbbbbb-cccc-41dd-8fff-ffffffffffff"), ] result = get_by_track_ids(session=session, track_ids=[101, 202]) session.execute.assert_called_once() assert result == [ { "track": {"tuid": 101}, "display_artists": [ { "contributor": {"uuid": "11111111-2222-4333-8444-555555555555"}, "role": "primary", }, { "contributor": {"uuid": "66666666-7777-4888-8999-aaaaaaaaaaaa"}, "role": "featuring", }, ], }, { "track": {"tuid": 202}, "display_artists": [ { "contributor": {"uuid": "bbbbbbbb-cccc-41dd-8fff-ffffffffffff"}, "role": "primary", } ], }, ] def test_uses_wrapped_session_when_not_passed(self, mocker: MockerFixture) -> None: session = mocker.Mock() execute_result = mocker.Mock() session.execute.return_value = execute_result execute_result.all.return_value = [ (101, "primary", 1, 10, "11111111-2222-4333-8444-555555555555"), ] ar_db_session_mock = mocker.patch( "contributor.connectors.mysql.ar_db_session", return_value=nullcontext(session), ) result = get_by_track_ids(track_ids=[101]) ar_db_session_mock.assert_called_once_with() assert result == [ { "track": {"tuid": 101}, "display_artists": [ { "contributor": {"uuid": "11111111-2222-4333-8444-555555555555"}, "role": "primary", } ], } ]