import pytest from datetime import date, datetime, timedelta, timezone from flask import url_for from http import HTTPStatus from typing import List, Optional, Union from apollo_main_db.spotify.models import ViewPlaylistTypeEnum from tests.legacy.spotify import factories from tests.legacy.util import generate_view_playlists def get_date( days_diff: int, base_date: Optional[Union[date, datetime]] = None, as_str: bool = False, with_time: bool = False, trim_seconds: bool = False, ) -> Union[date, datetime, str]: if not base_date: base_date = datetime.now(timezone.utc).replace(microsecond=0) if with_time else date.today() result = base_date - timedelta(days=days_diff) if as_str: result = result.isoformat() if trim_seconds: result = result[0:-6] return result def get_dates(*args, as_str: bool = False) -> List[date] or List[str]: return [get_date(i, as_str=as_str) for i in args] def get_nmf_playlists( count: int = 5, zero_followers_index: int = 4, dates_str: bool = False, with_name: bool = False ) -> List[dict]: return [ { "playlistId": f"pl_{i}", "countryCode": f"m_{i}", "rank": 2 * i, "fridayLastUpdatedDate": get_date(i, as_str=dates_str), "trackLastAdded": get_date(i, with_time=True, as_str=dates_str), "followers": 0 if zero_followers_index == i else (100 * i + 1), **({"name": f"pl_name_{i}"} if with_name else {}), } for i in range(1, count + 1) ] @pytest.mark.parametrize( "playlists_data,expected_result", ( (get_nmf_playlists(with_name=False), get_nmf_playlists(dates_str=True, with_name=False)), ( get_nmf_playlists(count=1, zero_followers_index=1, with_name=False), get_nmf_playlists(count=1, zero_followers_index=1, dates_str=True, with_name=False), ), ( get_nmf_playlists(count=9, zero_followers_index=1, with_name=False), get_nmf_playlists(count=9, zero_followers_index=1, dates_str=True, with_name=False), ), ), ) def test_playlists_nmf(playlists_data, expected_result, db_session, user_id, client, patch_auth_user): patch_auth_user(user_id) generate_view_playlists( playlists_data, ViewPlaylistTypeEnum.NMF, playlist_id_field="playlistId", market_code_field="countryCode", last_date_field="fridayLastUpdatedDate", last_added_ts_field="trackLastAdded", add_extra=False, ) for item in playlists_data: if item["followers"]: factories.SpotifyPlaylistFollowersFactory.create( playlist_id=item["playlistId"], followers=item["followers"] ) response = client.get(url_for("playlists_nmf.get_playlists_nmf")) assert response.status_code == HTTPStatus.OK assert response.json == expected_result def get_nmf_tracklists( count: int = 5, offset: int = 0, limit: int = 5, dates_str: bool = False, is_response: bool = False, with_image: bool = False, ) -> List[dict]: if limit > count: limit = count return [ { "isrc": f"isrc_{i}", "trackId": f"tr_{i}", "topTenFeatureCount": min(1 + i, max(10 - 3 * i, 0)), "playlists": [ { "playlistId": f"pl_{i}_{j}", "countryCode": f"cc_{i}_{j}", "addedDate": get_date(10 * i + j, with_time=True, as_str=dates_str, trim_seconds=dates_str), "position": 3 * i + j, **({} if is_response else {"rank": 1000 - 2 * i - j}), **({"playlistImageUrl": f"playlist-pl_{i}_{j}.jpeg"} if with_image else {}), } for j in (range(1 + i, 0, -1) if is_response else range(1, 2 + i)) ], } for i in (range(count - offset, count - limit - offset, -1) if is_response else range(offset + 1, limit + 1)) ] @pytest.mark.parametrize( "params,expected_status,track_data,expected_result", ( ({}, HTTPStatus.BAD_REQUEST, None, None), ( {"nmf_date": date.today()}, HTTPStatus.OK, get_nmf_tracklists(with_image=False), get_nmf_tracklists(dates_str=True, is_response=True, with_image=True), ), ( {"nmf_date": date.today()}, HTTPStatus.OK, get_nmf_tracklists(count=1, with_image=False), get_nmf_tracklists(count=1, dates_str=True, is_response=True, with_image=True), ), ( {"nmf_date": date.today(), "limit": 2}, HTTPStatus.OK, get_nmf_tracklists(with_image=False), get_nmf_tracklists(limit=2, dates_str=True, is_response=True, with_image=True), ), ( {"nmf_date": date.today(), "offset": 1, "limit": 2}, HTTPStatus.OK, get_nmf_tracklists(with_image=False), get_nmf_tracklists(offset=1, limit=2, dates_str=True, is_response=True, with_image=True), ), ), ) def test_playlists_nmf_tracks( params, expected_status, track_data, expected_result, db_session, user_id, client, patch_auth_user ): patch_auth_user(user_id) nmf_date = params.get("nmf_date") if nmf_date: for i in range(2): factories.SpotifyNewMusicFridayDateFactory.create(id=(i + 1), date=(nmf_date - timedelta(days=i))) generate_view_playlists( [p for t in track_data for p in t["playlists"]], ViewPlaylistTypeEnum.NMF, playlist_id_field="playlistId", market_code_field="countryCode", ) for index, track in enumerate(track_data): for playlist in track["playlists"]: for j in range(2): factories.SpotifyNewMusicFridayPlaylistTrackHistoryFactory.create( date_id=(j + 1), playlist_id=playlist["playlistId"], playlist_index=playlist["position"] - 1, track_id=track["trackId"], added=playlist["addedDate"], isrc=track["isrc"], ) response = client.get(url_for("playlists_nmf.get_playlists_nmf_tracks", **params)) assert response.status_code == expected_status if expected_status == HTTPStatus.OK: assert response.json["items"] == expected_result