from __future__ import annotations import time from datetime import datetime, timedelta from typing import Iterable, Mapping import jwt import requests from src import config from src.data_info import PlaylistInfo, TokenInfo, UserInfo from src.api_client.utils import request_in_chunks from ..errors import AppleMusicApiError from .base import BaseDSPApiClient __all__ = ["AppleMusicApiClient"] class AppleMusicApiClient(BaseDSPApiClient): default_error_cls = AppleMusicApiError base_url = "https://api.music.apple.com/v1" algorithm = "ES256" batch_size: int = 100 def __init__( self, *, key_id: str = config.APPLE_MUSICKIT_KEYID, secret_key: str = config.APPLE_MUSICKIT_KEY, team_id: str = config.APPLE_MUSICKIT_TEAM_ID, token_ttl: int = 24, image_base_url: str = config.IMAGE_BASE_URL, **kwargs, ): super().__init__(**kwargs) self._key_id = key_id self._secret_key = secret_key self._team_id = team_id self._token_ttl = token_ttl self._image_base_url = image_base_url def _get_token_from_code(self, code: str, redirect_uri: str, **extra) -> TokenInfo: access_token = self.get_access_token() return TokenInfo(refresh_token=code, access_token=access_token, expires_in=self._expires_in) def _refresh_access_token(self) -> tuple[str, int]: headers = {"alg": self.algorithm, "kid": self._key_id} expires_in = timedelta(hours=self._token_ttl) token_expired_at = int((datetime.utcnow() + expires_in).timestamp()) payload = {"iss": self._team_id, "iat": int(time.time()), "exp": token_expired_at} token = jwt.encode(payload, self._secret_key, algorithm=self.algorithm, headers=headers) return token, int(expires_in.total_seconds()) def _authorize_request(self, request: requests.Request): super()._authorize_request(request=request) request.headers["Music-User-Token"] = self._refresh_token @request_in_chunks(25, "isrc_list") def get_track_id_by_isrc(self, isrc_list: Iterable[str], storefront: str) -> Mapping[str, str]: if not isrc_list: return {} isrc_param = ",".join(isrc_list) self._logger.debug(f"Getting info for track with ISRC '{isrc_param}'") request = requests.Request( method="get", url=f"{self.base_url}/catalog/{storefront}/songs", params={"filter[isrc]": isrc_param} ) result = self.send_request_with_retry(request).json() if not result["data"]: return {} return {i["attributes"]["isrc"]: i["id"] for i in result["data"]} def get_playlist_track_ids(self, playlist_id: str, return_isrc: bool = False) -> Iterable[str]: self._logger.debug(f"Getting tracks for playlist with id '{playlist_id}'") page = 0 def handle_empty_playlist(response: requests.Response): if response.status_code == 404: return self._check_response(response) while True: request = requests.Request( method="get", url=f"{self.base_url}/me/library/playlists/{playlist_id}/tracks", params={ "limit": self.batch_size, "offset": page * self.batch_size, **({"include": "catalog"} if return_isrc else {}), }, ) result = self.send_request_with_retry(request, check_response=handle_empty_playlist).json() for track_item in result.get("data", []): if return_isrc: catalog_data = track_item.get("relationships", {}).get("catalog", {}).get("data") if not catalog_data: self._logger.debug(f"No catalog data for {playlist_id} {track_item['id']}") continue isrc = catalog_data[0].get("attributes", {}).get("isrc") if not isrc: self._logger.debug(f"No isrc for {playlist_id} {track_item['id']}") continue yield isrc else: catalog_id = track_item.get("attributes", {}).get("playParams", {}).get("catalogId") if not catalog_id: self._logger.debug(f"No catalog ID for {playlist_id} {track_item['id']}") continue yield catalog_id if result.get("next") is None: break page += 1 def get_storefront(self) -> str: # WARNING! undocumented API endpoint request = requests.Request( method="get", url=f"{self.base_url}/me/storefront", ) result = self.send_request_with_retry(request).json() if not result.get("data"): raise AppleMusicApiError("Can't get storefront") return result["data"][0]["id"] def get_playlist_id_map(self) -> Mapping[str, str]: mapping: dict[str, str] = {} page = 0 while True: request = requests.Request( method="get", url=f"{self.base_url}/me/library/playlists", params={"limit": self.batch_size, "offset": page * self.batch_size}, ) result = self.send_request_with_retry(request).json() for playlist_data in result["data"]: play_params = playlist_data.get("attributes", {}).get("playParams", {}) global_id, local_id = play_params.get("globalId"), play_params.get("id") if not global_id or not local_id: continue mapping[global_id] = local_id if result.get("next") is None: break page += 1 return mapping @request_in_chunks(config.ADD_TRACKS_CHUNK_SIZE, "track_ids", has_result=False) def insert_tracks(self, playlist_id: str, track_ids: Iterable[str]): request = requests.Request( method="post", url=f"{self.base_url}/me/library/playlists/{playlist_id}/tracks", json={"data": [{"id": track_id, "type": "songs"} for track_id in track_ids]}, ) self.send_request_with_retry(request) def get_playlist_info(self, playlist_id: str) -> PlaylistInfo: request = requests.Request( method="get", url=f"{self.base_url}/me/library/playlists/{playlist_id}", params={"include": "tracks"} ) response_data = self.send_request_with_retry(request).json() attributes = response_data["data"][0]["attributes"] relationships = response_data["data"][0]["relationships"] return PlaylistInfo( title=attributes["name"], description=attributes.get("description", {}).get("standard", ""), image_url=f"{self._image_base_url}/playlists/by_apple_id/{playlist_id}", total_tracks=relationships["tracks"].get("meta", {}).get("total", 0), user_id="", user_name=None, ) def get_user_info(self, **extra) -> UserInfo: # I can't find a way to get users info # there is an undocumented API endpoint '/me/account' but it's useless return UserInfo(user_identifier=extra.get("name", "AM account"), display_name=extra.get("name", "AM account"))