"""Product persister.""" from typing import List, Optional from sqlalchemy import String, and_, cast, column, distinct, func, or_, select, union from sqlalchemy.orm.session import Session from sqlalchemy.sql.elements import ColumnElement from collaborator.connectors import snowflake from collaborator.constants.product import STATUS_IN_CONTENT, DeletionStatus from collaborator.models.snowflake.label_participant import LabelParticipant from collaborator.models.snowflake.label_participant_product_particpation import ( LabelParticipantProductParticipation, ) from collaborator.models.snowflake.label_participant_track_particpation import ( LabelParticipantTrackParticipation, ) from collaborator.models.snowflake.product import Product from collaborator.models.snowflake.split import Split from collaborator.models.snowflake.track import Track class ProductPersister: """Handles product operations.""" @classmethod def _base_products_query(cls): """Create base products query.""" return ( select( Product.product_id, Product.vendor_id, Product.release_date, func.count(distinct(Track.id)).label("tracks_count"), func.count(distinct(Split.collaborator_id)).label( "collaborators_count" ), func.count(distinct(Split.identifier)).label( "tracks_with_splits_count" ), func.count(distinct(Split.split_id)).label("splits_count"), ) .where(Product.product_id.isnot(None)) .join(Track, and_(Product.upc == Track.upc, Track.tuid != 0), isouter=True) .join( Split, and_( cast(Track.tuid, String) == cast(Split.identifier, String), # Sqlite seems to require explicitly checking for IS NOT NULL # as well as != True, else it filters out results erroneously. or_( Split.fivetran_deleted.is_(None), Split.fivetran_deleted != True, # noqa: E712 ), ), isouter=True, ) .group_by( Product.product_id, Product.name, Product.vendor_id, Product.release_date, ) ) @classmethod @snowflake.db_session def get_products( cls, vendor_id: Optional[int], collaborator_id: Optional[int], limit: Optional[int], offset: Optional[int], min_tracks_with_splits: Optional[int], max_tracks_with_splits: Optional[int], min_tracks_without_splits: Optional[int], max_tracks_without_splits: Optional[int], term: Optional[str], label_participant_uuids: Optional[List[str]], subaccount_id: Optional[int], deleted: Optional[bool], sort_key: Optional[str], sort_direction: Optional[str], session: Session, ): """Get products.""" query = cls._base_products_query().where( Product.release_status == STATUS_IN_CONTENT ) if vendor_id is not None: query = query.where(Product.vendor_id == vendor_id) if collaborator_id is not None: query = query.where(Split.collaborator_id == collaborator_id) if term is not None: query = query.where( or_( Product.name.ilike(f"%{term}%"), cast(Product.display_upc, String).like(f"{term}%"), ) ) if label_participant_uuids is not None: query = query.where( Product.product_id.in_( cls._product_ids_for_label_participants_subquery( label_participant_uuids ) ) ) if subaccount_id is not None: query = query.where(Product.subaccount_id == subaccount_id) if deleted is not None: query = query.where( Product.deletions == (DeletionStatus.YES if deleted else DeletionStatus.NO) ) tracks_with_splits: ColumnElement = column("tracks_with_splits_count") tracks_without_splits: ColumnElement = column("tracks_count") - column( "tracks_with_splits_count" ) if min_tracks_with_splits is not None: query = query.having(tracks_with_splits >= min_tracks_with_splits) if max_tracks_with_splits is not None: query = query.having(tracks_with_splits <= max_tracks_with_splits) if min_tracks_without_splits is not None: query = query.having(tracks_without_splits >= min_tracks_without_splits) if max_tracks_without_splits is not None: query = query.having(tracks_without_splits <= max_tracks_without_splits) count = session.scalar(select(func.count()).select_from(query.subquery())) order = getattr(Product, sort_key) if sort_key else func.upper(Product.name) order = order.desc() if sort_direction == "DESC" else order.asc() query = query.order_by(order) if limit: query = query.limit(limit) if offset: query = query.offset(offset) rows = session.execute(query).all() return rows, count @classmethod @snowflake.db_session def get_counts_for_products(cls, product_ids: List[int], session: Session): """Get counts for products.""" query = cls._base_products_query().where(Product.product_id.in_(product_ids)) rows = session.execute(query).all() return rows @classmethod @snowflake.db_session def get_tracks_for_template( cls, vendor_id: int, session: Session ) -> list[tuple[int, str, str, str, str, str]]: """Return track rows for all vendor tracks. Each row is (product_id, product_upc, product_title, tuid_str, track_name, isrc). """ rows = session.execute( select( Product.product_id, Product.display_upc, Product.name, Track.tuid, Track.track_name, Track.isrc, ) .select_from(Product) .join(Track, Track.upc == Product.upc) .where( Product.vendor_id == vendor_id, Track.tuid != 0, ) .distinct() .order_by(Product.product_id, Track.tuid) ).all() return [ (product_id, display_upc, name, str(tuid), track_name, isrc) for product_id, display_upc, name, tuid, track_name, isrc in rows ] @classmethod def _product_ids_for_label_participants_subquery( cls, label_participant_uuids: List[str] ): """Get subquery for IDs of products with participants at product or track level.""" product_participations = ( select(Product.product_id) .join( LabelParticipantProductParticipation, Product.product_id == LabelParticipantProductParticipation.product_id, ) .join( LabelParticipant, and_( LabelParticipantProductParticipation.label_participant_id == LabelParticipant.id, LabelParticipant.uuid.in_(label_participant_uuids), ), ) ) track_participations = ( select(Product.product_id) .join(Track, Product.upc == Track.upc) .join( LabelParticipantTrackParticipation, Track.id == LabelParticipantTrackParticipation.track_id, ) .join( LabelParticipant, and_( LabelParticipantTrackParticipation.label_participant_id == LabelParticipant.id, LabelParticipant.uuid.in_(label_participant_uuids), ), ) ) return union(product_participations, track_participations).scalar_subquery() @classmethod @snowflake.db_session def get_product_ids_for_tuids( cls, tuids: set[str], session: Session ) -> dict[str, int]: """Return {tuid_str: product_id} for TUIDs found in DIM_TRACK.""" int_tuids = [int(t) for t in tuids if t.isdigit()] if not int_tuids: return {} rows = session.execute( select(Track.tuid, Product.product_id) .select_from(Track) .join(Product, Track.upc == Product.upc) .where(Track.tuid.in_(int_tuids)) ).all() return {str(tuid): product_id for tuid, product_id in rows} @classmethod @snowflake.db_session def get_vendor_map_by_tuids( cls, tuids: set[str], session: Session ) -> dict[str, int]: """Return {tuid_str: vendor_id} for TUIDs found in DIM_TRACK. Uses an inner join to Product so each TUID maps to its product's vendor. TUIDs absent from the result do not exist in Snowflake. """ int_tuids = [int(t) for t in tuids if t.isdigit()] rows = session.execute( select(Track.tuid, Product.vendor_id) .select_from(Track) .join(Product, Track.upc == Product.upc) .where(Track.tuid.in_(int_tuids)) ).all() return {str(tuid): vendor_id for tuid, vendor_id in rows} @classmethod @snowflake.db_session def get_ids_for_vendor( cls, product_ids: set[int], vendor_id: int, session: Session ) -> set[int]: """Return the subset of product IDs that belong to the vendor.""" rows = session.execute( select(Product.product_id).where( Product.product_id.in_(product_ids), Product.vendor_id == vendor_id, ) ).all() return {pid for (pid,) in rows} @classmethod @snowflake.db_session def get_tuid_product_pairs( cls, product_ids: set[int], session: Session ) -> list[tuple[str, int]]: """Return (tuid_str, product_id) pairs for tracks matching the given product IDs.""" rows = session.execute( select(Track.tuid, Product.product_id) .select_from(Product) .join(Track, Track.upc == Product.upc) .where(Product.product_id.in_(product_ids), Track.tuid != 0) ).all() return [(str(tuid), product_id) for tuid, product_id in rows]