"""DB utils.""" from contextlib import contextmanager from typing import Dict, Iterable, List from sqlalchemy import create_engine, union_all from sqlalchemy.orm import Query, scoped_session, sessionmaker import config from apollo_main_db.spotify.models import SpotifyPersonalizedPlaylist, SpotifyPersonalizedPlaylistTrackList, \ SpotifyTrack2 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_personalized_playlists_list() -> List[str]: """Returns list of all personalized playlists id.""" with session_scope() as session: return [p.playlist_id for p in session.query(SpotifyPersonalizedPlaylist)] def get_saved_tracks_isrc_list(playlist_id: str) -> List[str]: """Returns personalized tracks isrc list for particular playlist.""" with session_scope() as session: q = session.query(SpotifyPersonalizedPlaylistTrackList.isrc).filter( SpotifyPersonalizedPlaylistTrackList.playlist_id == playlist_id ) return [t.isrc.upper() for t in q] def _get_track_id_query(session, isrc: str) -> Query: """Returns query to get one track_id by isrc.""" return ( session.query(SpotifyTrack2.id.label("track_id"), SpotifyTrack2.isrc.label("isrc")) .filter(SpotifyTrack2.isrc == isrc) .limit(1) ) def get_tracks_ids_query(isrc_list: List[str]) -> Iterable: """Returns query to get track_id list by isrc_list""" if not isrc_list: return [] with session_scope() as session: return session.query(union_all(*[_get_track_id_query(session, _isrc) for _isrc in isrc_list]).alias()) def delete_playlist_tracks(playlist_id: str, isrc_list: List[str]): """Delete tracks with isrc from isrc_list from track list of particular playlist.""" if not isrc_list: return with session_scope() as session: return ( session.query(SpotifyPersonalizedPlaylistTrackList) .filter(SpotifyPersonalizedPlaylistTrackList.playlist_id == playlist_id) .filter(SpotifyPersonalizedPlaylistTrackList.isrc.in_(isrc_list)) .delete(synchronize_session=False) ) def add_playlist_tracks(playlist_id: str, isrc_to_track_map: Dict[str, list]): """Insert tracks to track list of particular playlist.""" if not isrc_to_track_map: return with session_scope() as session: objects = [ SpotifyPersonalizedPlaylistTrackList( playlist_id=playlist_id, isrc=_isrc, added_datetime=track_data[0], track_id=track_data[1] ) for _isrc, track_data in isrc_to_track_map.items() ] session.bulk_save_objects(objects) session.commit()