"""Test for Track Model.""" from datetime import datetime import pytest from pytest import raises from sqlalchemy import event from backend.connectors import mysql from backend.constants import track_field as tf from backend.models.track_additional_isrc import TrackAdditionalIsrc from backend.models.track_additional_isrc import TrackAdditionalIsrcType from backend.models.track_query import TrackQuery from backend.models.track_spatial import TrackSpatial from tests.testutils import db # Global used to keep track of number of SQL statements run for # test_track_query_eager_loading test. This needs to be a global due # to scoping issues. _num_statements = 0 @db.test_schema def test_get_by_tuid(test_track): """Test getting last track for product is successful.""" tuid = test_track[tf.TUID] with mysql.db_session() as session: track = TrackQuery.get_by_tuid(tuid, session) assert track.tuid == tuid @db.test_schema def test_get_all_by_isrc_and_type(test_track): """Test getting by isrc and type is successful.""" with mysql.db_session() as session: results = TrackQuery.get_all_by_isrc_and_type( test_track['isrc'], tf.TRACK_TYPE_MUSIC, session ) assert len(results) != 0 for track in results: assert track.isrc == test_track['isrc'] assert track.track_type == tf.TRACK_TYPE_MUSIC @db.test_schema def test_get_by_tuids(test_tuids): """Test getting tracks by list of tuids.""" retrieved_tuids = [] with mysql.db_session() as session: tracks = TrackQuery.get_by_tuids(test_tuids, session) for track in tracks: assert track.__class__.__name__ == 'Track' retrieved_tuids.append(track.tuid) assert set(test_tuids) == set(retrieved_tuids) @db.test_schema def test_get_tuids_by_product_ids_sorted_volume_number_order( test_product_ids, ): """Test getting tracks by list of tuids.""" retrieved_tuids = [] expected_tuids = [28, 29, 30, 31, 32, 33, 34, 35, 36, 37] with mysql.db_session() as session: tracks = TrackQuery.get_tuids_by_product_ids_with_order( test_product_ids, [{'column_name': 'volume_number', 'order': 'asc'}], session, ) # order_by_product_id=True by default for track in tracks: assert track.__class__.__name__ == 'Track' retrieved_tuids.append(track.tuid) assert retrieved_tuids == expected_tuids @db.test_schema def test_get_tuids_by_product_ids_sorted_track_number_order( test_product_ids, ): """Test getting tracks by list of tuids.""" retrieved_tuids = [] expected_tuids = [28, 37, 36, 35, 34, 33, 32, 31, 30, 29] with mysql.db_session() as session: tracks = TrackQuery.get_tuids_by_product_ids_with_order( test_product_ids, [{'column_name': 'track_number', 'order': 'asc'}], session, ) # order_by_product_id=True by default for track in tracks: assert track.__class__.__name__ == 'Track' retrieved_tuids.append(track.tuid) assert retrieved_tuids == expected_tuids @db.test_schema def test_get_tuids_by_product_ids_sorted_track_number_desc_order( test_product_ids, ): """Test getting tracks by list of tuids.""" retrieved_tuids = [] expected_tuids = [29, 30, 31, 32, 33, 34, 35, 36, 28, 37] with mysql.db_session() as session: tracks = TrackQuery.get_tuids_by_product_ids_with_order( test_product_ids, [{'column_name': 'track_number', 'order': 'desc'}], session, ) # order_by_product_id=True by default for track in tracks: assert track.__class__.__name__ == 'Track' retrieved_tuids.append(track.tuid) assert retrieved_tuids == expected_tuids @db.test_schema def test_get_tuids_by_product_ids_dict_column_name(test_product_ids): """Test illegal column name.""" with raises(ValueError): with mysql.db_session() as session: TrackQuery.get_tuids_by_product_ids_with_order( test_product_ids, [{'column_name': '__dict__', 'order': 'desc'}], session, ) @db.test_schema def test_get_by_tuids_for_product(test_tuids): """Test getting tracks by list of tuids belonging to product works.""" with mysql.db_session() as session: tracks = TrackQuery.get_by_tuids( test_tuids, session, belongs_to_product_id=1) retrieved_tuids = [track.tuid for track in tracks] assert set(test_tuids) == set(retrieved_tuids) @db.test_schema def test_get_by_tuids_for_product_missing(test_tuids_multiple_products): """Test getting tracks by list of tuids belonging to product fails.""" with mysql.db_session() as session: with pytest.raises(ValueError): TrackQuery.get_by_tuids( test_tuids_multiple_products, session, belongs_to_product_id=1) @db.test_schema def test_get_by_tuids_with_for_update(test_tuids): """Test getting tracks by list of tuids with_for_update.""" with mysql.db_session() as session: tracks = TrackQuery.get_by_tuids( test_tuids, session, with_for_update=True) assert tracks @db.test_schema def test_get_by_tuids_with_nones(test_tuids): """Test getting tracks by list of tuids in order with nones.""" retrieved_tuids = [] with mysql.db_session() as session: tracks = TrackQuery.get_by_tuids_with_nones(test_tuids, session) for track in tracks: assert track.__class__.__name__ == 'Track' retrieved_tuids.append(track.tuid) assert test_tuids == retrieved_tuids @db.test_schema def test_get_all_by_product_id(test_multi_tracks_product): """Test getting all tracks by product_id is successful.""" test_product_id = test_multi_tracks_product['product_id'] with mysql.db_session() as session: items = TrackQuery.get_all_by_product_id(test_product_id, session) last_volume = 1 last_track = 0 # Makse sure all tracks belong to product and are in order for track in items: assert track.product_id == test_product_id assert track.volume_number >= last_volume if track.volume_number > last_volume: track.volume_number = last_volume last_track = 0 assert track.track_number > last_track @db.test_schema def test_get_all_by_product_id_for_overview(test_multi_tracks_product): """Test getting all tracks by product_id is successful.""" test_product_id = test_multi_tracks_product['product_id'] with mysql.db_session() as session: items = TrackQuery.get_all_by_product_id( test_product_id, session, is_overview=True) last_volume = 1 last_track = 0 for track in items: assert track.product_id == test_product_id assert track.volume_number >= last_volume if track.volume_number > last_volume: track.volume_number = last_volume last_track = 0 assert track.track_number > last_track @db.test_schema def test_get_all_by_product_id_light(test_multi_tracks_product): """Test getting all tracks by product_id light is successful.""" test_product_id = test_multi_tracks_product['product_id'] with mysql.db_session() as session: items = TrackQuery.get_all_by_product_id_light(test_product_id, session) last_volume = 1 last_track = 0 # Makse sure all tracks belong to product and are in order for track in items: assert track.product_id == test_product_id assert track.volume_number >= last_volume if track.volume_number > last_volume: track.volume_number = last_volume last_track = 0 assert track.track_number > last_track def test_claim_new_isrcs(mocker): """Test claim new isrcs.""" mock_response = mocker.MagicMock() mock_response.fetchall.return_value = [('test_isrc',)] mock_session = mocker.MagicMock() mock_session.execute.return_value = mock_response TrackQuery.claim_new_isrcs(mock_session, 2) mock_session.execute.assert_called_with( TrackQuery.SP_CLAIM_ISRCS_FOR_USE, {'number_of_isrcs': 2}) result = TrackQuery.claim_new_isrcs(mock_session) mock_session.execute.assert_called_with( TrackQuery.SP_CLAIM_ISRCS_FOR_USE, {'number_of_isrcs': 1}) assert result == ['test_isrc'] @db.test_schema def test_track_query_eager_loading_for_single_tuid(test_track): """Test eager loading is loading only one to one tables.""" global _num_statements try: event.listen( mysql.db_engine, 'before_cursor_execute', _inc_num_statements) with mysql.db_session() as session: track = TrackQuery.get_by_tuid(1, session) _num_statements = 0 track.ownership_rights assert _num_statements == 0 track.original_rights_holder_country_id assert _num_statements == 0 len(track.artists) assert _num_statements == 1 len(track.publishers) assert _num_statements == 2 len(track.writers) assert _num_statements == 3 finally: event.remove( mysql.db_engine, 'before_cursor_execute', _inc_num_statements) @db.test_schema def test_track_query_eager_loading_for_multiple_tuids(test_track): """Test eager loading is loading one to many tables.""" global _num_statements try: event.listen( mysql.db_engine, 'before_cursor_execute', _inc_num_statements) with mysql.db_session() as session: tracks = TrackQuery.get_all_by_product_id(1, session)[:2] _num_statements = 0 for track in tracks[:2]: track.ownership_rights assert _num_statements == 0 track.original_rights_holder_country_id assert _num_statements == 0 len(track.artists) assert _num_statements == 0 len(track.publishers) assert _num_statements == 0 len(track.writers) assert _num_statements == 0 finally: event.remove( mysql.db_engine, 'before_cursor_execute', _inc_num_statements) def _inc_num_statements(*args): """Increment global _num_statements each time query is run. Note that this function must be attached to the proper SQLAlchemy event. See http://docs.sqlalchemy.org/en/latest/core/events.html """ global _num_statements _num_statements += 1 @db.test_schema_no_seed def test_track_add_spatial(track_factory): """Test add_spatial adds a TrackSpatial relationship and its mirror.""" track = track_factory.create(tuid=100, product_id=1) spatial = track.add_spatial('US1234567890') assert isinstance(spatial, TrackSpatial) assert spatial.track_id == track.tuid assert spatial.isrc == 'US1234567890' assert track._spatial == spatial assert len(track._additional_isrcs) == 1 mirror = track._additional_isrcs[0] assert mirror.type == TrackAdditionalIsrcType.ATMOS assert mirror.isrc == 'US1234567890' @db.test_schema def test_add_spatial_cascade_persists_track_additional_isrc_mirror(): """add_spatial's mirror is cascade-inserted when the track is committed. Covers the clone path (Track.add_spatial), which writes track_spatial via the relationship cascade instead of create_track_spatial -- the mirror must ride the same cascade so cloned spatial tracks aren't missed by the dual-write. """ with mysql.db_session() as session: track = TrackQuery.get_by_tuid(1, session, eager_loading=False) track.add_spatial('US1234567890') session.add(track) session.commit() mirror = session.query(TrackAdditionalIsrc).filter( TrackAdditionalIsrc.track_id == 1).one() assert mirror.type == TrackAdditionalIsrcType.ATMOS assert mirror.isrc == 'US1234567890' assert mirror.deleted_at is None @db.test_schema def test_track_query_eager_loading_for_single_tuid_loads_spatial(): """Test eager loading includes spatial one-to-one table.""" global _num_statements try: event.listen( mysql.db_engine, 'before_cursor_execute', _inc_num_statements) with mysql.db_session() as session: track = TrackQuery.get_by_tuid(1, session, eager_loading=False) track.add_spatial('US1234567890') session.add(track) session.commit() session.expire_all() track = TrackQuery.get_by_tuid(1, session) _num_statements = 0 assert track._spatial is not None assert track._spatial.isrc == 'US1234567890' assert track._spatial.track_id == 1 finally: event.remove( mysql.db_engine, 'before_cursor_execute', _inc_num_statements) @db.test_schema_no_seed def test_get_spatial_isrc_map_by_product_id(track_factory): """Test returns mapping of track_id to ISRC for tracks with spatial data.""" product_id = 1 tracks = [track_factory(product_id=product_id), track_factory(product_id=product_id)] db.merge_model_objects(tracks) db.merge_model_objects([TrackSpatial(track_id=tracks[0].tuid, isrc='US1234567890')]) with mysql.db_session() as session: result = TrackQuery.get_spatial_isrc_map_by_product_id(product_id, session) assert result == {tracks[0].tuid: 'US1234567890'} @db.test_schema_no_seed def test_get_spatial_isrc_map_by_product_id_no_spatial(track_factory): """Test returns empty dict when product tracks have no spatial records.""" product_id = 1 track = track_factory(product_id=product_id) db.merge_model_objects([track]) with mysql.db_session() as session: result = TrackQuery.get_spatial_isrc_map_by_product_id(product_id, session) assert result == {} @db.test_schema_no_seed def test_get_spatial_isrc_map_by_product_id_excludes_soft_deleted(track_factory): """Soft-deleted (deleted_at set) spatial records are excluded from the map.""" product_id = 1 tracks = [track_factory(product_id=product_id), track_factory(product_id=product_id)] db.merge_model_objects(tracks) db.merge_model_objects([ TrackSpatial( track_id=tracks[0].tuid, isrc='US1234567890', deleted_at=datetime(2026, 1, 1)), TrackSpatial(track_id=tracks[1].tuid, isrc='US0987654321'), ]) with mysql.db_session() as session: result = TrackQuery.get_spatial_isrc_map_by_product_id(product_id, session) assert result == {tracks[1].tuid: 'US0987654321'} @db.test_schema_no_seed def test_get_spatial_isrc_map_by_product_id_unknown_product(): """Test returns empty dict for a product with no tracks.""" with mysql.db_session() as session: result = TrackQuery.get_spatial_isrc_map_by_product_id(99999, session) assert result == {}