"""DB utils.""" from contextlib import contextmanager from datetime import date, datetime from typing import Dict, List, Optional from apollo_main_db.apollo.models import ApolloKeyValueStorage from apollo_main_db.spotify.models import (MarketRankTypeEnum, SpotifyAnalyticsAccountStreamInfo, SpotifyMarketRank, SpotifyNewMusicFridayDate, SpotifyNewMusicFridayPlaylistTrackHistory, SpotifyViewPlaylist, ViewPlaylistTypeEnum) from sqlalchemy import bindparam, create_engine, func, update from sqlalchemy.orm import scoped_session, sessionmaker import config engine = create_engine( "mysql+pymysql://{user}:{password}@{host}:{port}/{database_name}".format( user=config.MySQL.USER, password=config.MySQL.PASSWORD, host=config.MySQL.HOST, port=config.MySQL.PORT, database_name=config.MySQL.DATABASE, ), pool_recycle=config.MySQL.POOL_RECYCLE, pool_size=config.MySQL.POOL_SIZE, ) _session_factory = scoped_session(sessionmaker(bind=engine)) _session = _session_factory() @contextmanager def session_scope(): """Provide a transactional scope around a series of operations.""" try: yield _session _session.commit() except Exception: _session.rollback() raise finally: _session.close() def get_playlists() -> List[SpotifyViewPlaylist]: """Get IDs of all NMF playlists. Returns: List of NMF playlist ID. """ with session_scope() as session: result = session.query(SpotifyViewPlaylist).filter(SpotifyViewPlaylist.type == ViewPlaylistTypeEnum.NMF).all() session.expunge_all() return result def get_date(friday: date) -> int: """Get NMF date object. Args: friday: Friday date. Returns: NMF date object. """ with session_scope() as session: result = session.query(SpotifyNewMusicFridayDate.id).filter(SpotifyNewMusicFridayDate.date == friday).first() return result[0] if result else None def create_date(friday: date) -> int: """Create NMF date object. Args: friday: Friday date. Returns: NMF date object. """ with session_scope() as session: result = SpotifyNewMusicFridayDate(date=friday) session.add(result) session.flush() return result.id def get_playlist_tracks(playlist_id: str, date_id: int) -> List[SpotifyNewMusicFridayPlaylistTrackHistory]: """Get NMF tracks for chosen playlist and date. Args: playlist_id: Playlist ID. date_id: NMF date ID. Returns: NMF tracks. """ with session_scope() as session: result = ( session.query(SpotifyNewMusicFridayPlaylistTrackHistory) .filter(SpotifyNewMusicFridayPlaylistTrackHistory.playlist_id == playlist_id) .filter(SpotifyNewMusicFridayPlaylistTrackHistory.date_id == date_id) .all() ) session.expunge_all() return result def update_tracklist(playlist_id: str, date_id: int, tracks: List[SpotifyNewMusicFridayPlaylistTrackHistory]): """Update NMF tracklist for a specific playlist. Args: playlist_id: Playlist ID. date_id: NMF date ID. tracks: NMF track list. """ with session_scope() as session: ( session.query(SpotifyNewMusicFridayPlaylistTrackHistory) .filter(SpotifyNewMusicFridayPlaylistTrackHistory.playlist_id == playlist_id) .filter(SpotifyNewMusicFridayPlaylistTrackHistory.date_id == date_id) .delete(synchronize_session=False) ) if tracks: session.bulk_save_objects(tracks) def update_playlist_dates(playlist_id: str, last_date: date, last_added_ts: Optional[datetime]): """Update last NMF date and last playlist changed (track added) timestamp. Args: playlist_id: Playlist ID. last_date: Last NMF date. last_added_ts: Last track added to playlist. """ with session_scope() as session: ( session.query(SpotifyViewPlaylist) .filter(SpotifyViewPlaylist.playlist_id == playlist_id) .filter(SpotifyViewPlaylist.type == ViewPlaylistTypeEnum.NMF) .update( { SpotifyViewPlaylist.last_date: last_date, SpotifyViewPlaylist.last_added_ts: last_added_ts, }, synchronize_session=False, ) ) def get_value(key: str) -> str or None: """Get value by key from key-value storage. Args: key: Key. Returns: Value. """ with session_scope() as session: return session.query(ApolloKeyValueStorage.value).filter(ApolloKeyValueStorage.key == key).first()[0] def set_value(key: str, value: str): """Set value for a key to key-value storage. Args: key: Key. value: Value. """ with session_scope() as session: ( session.query(ApolloKeyValueStorage) .filter(ApolloKeyValueStorage.key == key) .update({ApolloKeyValueStorage.value: value}, synchronize_session=False) ) def get_markets_order_by_rank(date_from: date, date_to: date) -> List[str]: """Get markets ordered by rank. Args: date_from: Date from. date_to: Date to. Returns: Ordered list of markets. """ with session_scope() as session: return [ i[0] for i in ( session.query(SpotifyAnalyticsAccountStreamInfo.market) .filter(SpotifyAnalyticsAccountStreamInfo.account == 1) .filter(SpotifyAnalyticsAccountStreamInfo.date >= date_from) .filter(SpotifyAnalyticsAccountStreamInfo.date <= date_to) .group_by(SpotifyAnalyticsAccountStreamInfo.market) .order_by(func.sum(SpotifyAnalyticsAccountStreamInfo.total_streams).desc()) ) ] def get_last_streams_date() -> date: """Get max streams table date. Returns: Last available date. """ with session_scope() as session: return ( session.query(func.max(SpotifyAnalyticsAccountStreamInfo.date)) .filter(SpotifyAnalyticsAccountStreamInfo.account == 1) .first() )[0] def set_playlists_ranks( market_rank_mapping: Dict[str, int], column_name: str = "rank", record_type: str = MarketRankTypeEnum.NMF ): """Set markets ranks. Args: market_rank_mapping: Market code to rank mapping. column_name: Column name to update. record_type: Rank record type. """ with session_scope() as session: query = ( update(SpotifyMarketRank) .where(SpotifyMarketRank.market_code == bindparam("market_code")) .where(SpotifyMarketRank.type == record_type) .values({getattr(SpotifyMarketRank, column_name): bindparam(column_name)}) ) session.execute( query, [{"market_code": market_code, column_name: value} for market_code, value in market_rank_mapping.items()], ) def get_playlists_markets() -> List[str]: """Get NMF playlists market code list. Returns: Market code list. """ with session_scope() as session: market_list = ( session.query(SpotifyViewPlaylist.market_code) .filter(SpotifyViewPlaylist.type == ViewPlaylistTypeEnum.NMF) .all() ) return [i.market_code.lower() for i in market_list]