"""Track Queries.""" from typing import Callable from sqlalchemy import asc from sqlalchemy import desc from sqlalchemy import inspect from sqlalchemy import orm from sqlalchemy.orm import load_only from backend.constants import track_field from backend.enums import SortingOrderEnum from backend.models.track import Track from backend.models.track_spatial import TrackSpatial class TrackQuery(object): """Common queries for Track. This returns the query results instead of dicts. """ SP_CLAIM_ISRCS_FOR_USE = 'CALL `sp_claim_isrcs_for_use`(:number_of_isrcs);' @classmethod def get_by_tuid(cls, tuid, session, eager_loading=True): """Get track by its unique id. Args: tuid (int): unique id of track session (object): session object eager_loading (bool): Use eager loading for related data Returns: Track: Will return None if tuid not found in database """ return cls._base_query(session, joinedload=eager_loading).get(tuid) @classmethod def get_by_upc_and_isrc(cls, upc, isrc, session): """Get track by upc and isrc. Args: upc (int): UPC of the track's product isrc (str): ISRC of the track session (object): session object Returns: Track: Will return None if not found """ return session.query(Track).filter_by( upc=upc, isrc=isrc ).first() @classmethod def get_all_by_isrc_and_type(cls, isrc, track_type, session): """Get all results by isrc and track_type.""" return session.query(Track).filter_by( isrc=isrc, track_type=track_type ).all() @classmethod def get_by_tuids( cls, tuids, session, belongs_to_product_id=None, with_for_update=False, eager_loading=True): """Get list of tracks by its unique id. Args: tuids (list): List of track unique ids session (object): session object belongs_to_product_id (int): Product Id track must belong to eager_loading (bool): Use eager loading for related data Returns: List: Returns list of found Track objects """ query = cls._base_query( session, joinedload=eager_loading, subqueryload=eager_loading).\ filter(Track.tuid.in_(tuids)) if belongs_to_product_id: query = query.filter_by(product_id=belongs_to_product_id) query = query.order_by(Track.tuid) if with_for_update: query = query.with_for_update() tracks = [] for track in query: tracks.append(track) if set(tuids) != set(track.tuid for track in tracks): raise ValueError('Not all tracks were found') return tracks @classmethod def get_by_tuids_with_nones( cls, tuids, session, eager_loading=True): """Get list of tracks by its unique id. Args: tuids (list): List of track unique ids session (object): session object belongs_to_product_id (int): Product Id track must belong to eager_loading (bool): Use eager loading for related data Returns: List: Returns list of found Track objects, or None if the track does not exist. """ query = cls._base_query( session, joinedload=eager_loading, subqueryload=eager_loading).\ filter(Track.tuid.in_(tuids)) tracks_by_tuid = {} for track in query: tracks_by_tuid[track.tuid] = track tracks = [ tracks_by_tuid.get(tuid, None) for tuid in tuids ] return tracks @classmethod def get_all_by_product_id( cls, product_id, session, with_for_update=False, eager_loading=True, is_overview=False): """Get list of all tracks by product_id. Orders by volume and track number in ascending order. Args: product_id (int): id of product session (object): session object eager_loading (bool): Use eager loading for related data Returns: Query Object """ if is_overview: query = session.query(Track)\ .options(load_only(*track_field.OVERVIEW_MODEL_FIELDS), orm.subqueryload(Track.artists))\ .filter_by(product_id=product_id) else: query = cls._base_query( session, joinedload=eager_loading, subqueryload=eager_loading).\ filter_by(product_id=product_id).\ order_by(Track.volume_number, Track.track_number) if with_for_update: query = query.with_for_update() return query @classmethod def get_tuids_by_product_ids_with_order( cls, product_ids, order_by_fields, session, ): """Get list of all track IDs by product_ids with given sorting order. Args: product_ids (list): ids of product order_by_fields (iterable): iterable of fields names for ordering session (object): session object Returns: Query Object """ if not order_by_fields: raise ValueError('order by fields must be provided') track_mapper = inspect(Track) order_by_objects = [] for order_by_field in order_by_fields: if order_by_field['column_name'] not in track_mapper.columns: raise ValueError(f'Invalid order_by field "{order_by_field}"') order_by_objects.append( cls._get_order_function( order_str=order_by_field['order'], )( getattr(Track, order_by_field['column_name']), ), ) query = session.query(Track).options( load_only( track_field.TUID, track_field.PRODUCT_ID, ), ).filter( Track.product_id.in_(set(product_ids)), ).order_by(*order_by_objects) return query @classmethod def get_all_track_isrc_by_product_id( cls, product_id, session): """Get list of all track ISRCs by product_id. Args: product_id (int): ID of product session (object): session object Returns: Query Object """ query = session.query(Track).\ options(load_only(track_field.ISRC)).\ filter_by(product_id=product_id).all() return query @classmethod def get_spatial_isrc_map_by_product_id(cls, product_id, session): """Get a mapping of track_id to spatial ISRC for all tracks in a product. Args: product_id (int): id of product session (object): session object Returns: dict: {track_id: isrc} for tracks that have spatial data """ records = session.query(TrackSpatial.track_id, TrackSpatial.isrc).join( Track, Track.tuid == TrackSpatial.track_id ).filter( Track.product_id == product_id, TrackSpatial.deleted_at.is_(None) ).all() return {track_id: isrc for track_id, isrc in records} @classmethod def get_all_focus_track_by_product_id( cls, product_id, session, with_for_update=False, eager_loading=True, is_overview=False): """Get list of all focus tracks by product_id. Orders by volume and track number in ascending order. Args: product_id (int): id of product session (object): session object eager_loading (bool): Use eager loading for related data Returns: Query Object """ if is_overview: query = session.query(Track)\ .options(load_only(*track_field.OVERVIEW_MODEL_FIELDS), orm.subqueryload(Track.artists))\ .filter_by(product_id=product_id) else: query = cls._base_query( session, joinedload=eager_loading, subqueryload=eager_loading).\ filter_by(product_id=product_id).filter(Track._focus_track.__ne__(None)).\ order_by(Track.volume_number, Track.track_number) if with_for_update: query = query.with_for_update() return query @classmethod def get_all_by_product_id_light(cls, product_id, session): """Get list of all tracks by product_id. Orders by volume and track number in ascending order. Args: product_id (int): id of product session (object): session object Returns: Query Object """ query = session.query(Track)\ .options(load_only(*track_field.LIGHT_BASIC_MODEL_FIELDS))\ .filter_by(product_id=product_id) return query @classmethod def get_all_by_product_ids_medium(cls, product_ids, session): """Get list of all tracks by product_ids. Orders by volume and track number in ascending order. Args: product_ids (list): ids of product session (object): session object Returns: Query Object """ query = session.query(Track)\ .options(load_only(*track_field.MEDIUM_BASIC_MODEL_FIELDS))\ .filter( Track.product_id.in_(set(product_ids)), ) return query @classmethod def claim_new_isrcs(cls, session, number_of_isrcs=1): """Acquire new ISRCs from the database. Args: session (object): SQLAlchemy database session number_of_isrcs (int): Number of ISRCs to claim Returns: list(str): list of new ISRCs """ # Get a list of ISRCs from the database isrc_response = session.execute( cls.SP_CLAIM_ISRCS_FOR_USE, {'number_of_isrcs': number_of_isrcs}) result = isrc_response.fetchall() return [isrc[0] for isrc in result] @classmethod def _base_query(cls, session, joinedload=False, subqueryload=False): """Build intial Track query. Args: session (object): session object joinedload (bool): Load related data (one to one) subqueryload (bool): Load related data (one to many) Returns: Query Object """ query = session.query(Track) options = [] if joinedload: options += [ orm.joinedload(Track._master_rights), orm.joinedload(Track._producer_nationality), orm.joinedload(Track.language), orm.joinedload(Track._focus_track), orm.joinedload(Track._spatial)] if subqueryload: options += [ orm.subqueryload(Track.artists), orm.subqueryload(Track.publishers), orm.subqueryload(Track.writers)] if options: query = query.options(*options) return query @classmethod def _get_order_function(cls, order_str: str) -> Callable: """Return correct orm ordering function. Args: order_str: string representation of ordering: (ask, desc). Returns: orm.asc if order_str == 'asc' and orm.desc if order_str == 'desc' """ order_enum = SortingOrderEnum(order_str) return asc if order_enum == SortingOrderEnum.ASC else desc