"""Tests for product_display_artist MySQL queries.""" from contextlib import nullcontext from pytest_mock import MockerFixture from contributor.queries.mysql.product_display_artist import ( ProductDisplayArtist, get_by_product_ids, ) class TestProductDisplayArtistModel: def test_model_has_correct_tablename(self) -> None: assert ProductDisplayArtist.__tablename__ == "product_display_artists" def test_model_instantiation_with_all_fields(self) -> None: product_display_artist = ProductDisplayArtist( id=1, product_id=100, contributor_id=50, display_artist_role="primary", sequence_number=1, ) assert product_display_artist.id == 1 assert product_display_artist.product_id == 100 assert product_display_artist.contributor_id == 50 assert product_display_artist.display_artist_role == "primary" assert product_display_artist.sequence_number == 1 def test_model_instantiation_with_required_fields_only(self) -> None: product_display_artist = ProductDisplayArtist( product_id=100, contributor_id=50, display_artist_role="featuring", ) assert product_display_artist.product_id == 100 assert product_display_artist.contributor_id == 50 assert product_display_artist.display_artist_role == "featuring" def test_model_instantiation_with_sequence_number(self) -> None: product_display_artist = ProductDisplayArtist( product_id=100, contributor_id=50, display_artist_role="feature_to_primary", sequence_number=5, ) assert product_display_artist.sequence_number == 5 def test_model_has_all_required_attributes(self) -> None: required_fields = [ "id", "product_id", "contributor_id", "display_artist_role", "sequence_number", ] for field in required_fields: assert hasattr(ProductDisplayArtist, field) class TestGetByProductIds: def test_returns_display_artists_grouped_by_product_id( self, mocker: MockerFixture ) -> None: session = mocker.Mock() execute_result = mocker.Mock() session.execute.return_value = execute_result execute_result.all.return_value = [ (12345, "primary", 1, 10, "11111111-2222-4333-8444-555555555555"), (12345, "featuring", 2, 11, "66666666-7777-4888-8999-aaaaaaaaaaaa"), (67890, "primary", 1, 12, "bbbbbbbb-cccc-41dd-8fff-ffffffffffff"), ] result = get_by_product_ids(session=session, product_ids=[12345, 67890]) session.execute.assert_called_once() assert result == [ { "product": {"product_id": 12345}, "display_artists": [ { "contributor": {"uuid": "11111111-2222-4333-8444-555555555555"}, "role": "primary", }, { "contributor": {"uuid": "66666666-7777-4888-8999-aaaaaaaaaaaa"}, "role": "featuring", }, ], }, { "product": {"product_id": 67890}, "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 = [ (12345, "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_product_ids(product_ids=[12345]) ar_db_session_mock.assert_called_once_with() assert result == [ { "product": {"product_id": 12345}, "display_artists": [ { "contributor": {"uuid": "11111111-2222-4333-8444-555555555555"}, "role": "primary", } ], } ]