""" Integration with Vendor API. Vendor API: https://github.com/filtr/vendor-api. """ from apollo_utils.service.exceptions import APIInvalidResponse from typing import Any, Dict, List, Optional, Tuple from server import config from server.client.base.client import ApiKeyClient from server.client.base.config import ApiKeyClientConfig from server.client.utils import prepare_list_arg, prepare_list_arg_for_names from server.constants import DSP, Service from server.constants.market import Market from server.utils.parallel import ListResult, request_in_chunks class VendorApiConfig(ApiKeyClientConfig): """Configuration class for Vendor api client.""" service = Service.VENDOR_API.value uri = config.VENDOR_URI auth_secret = config.VENDORAPI_AUTH data_ttl = config.VENDOR_DEFAULT_DATA_TTL class VendorApiClient(ApiKeyClient): """Client for sending requests to Vendor api app.""" config: VendorApiConfig @classmethod def prepare_request_data( cls, params: Dict or None, headers: Dict or None, body: Dict or None, data: Dict or None, **kwargs ): dsp = kwargs.get("dsp") request_data = params or body or data if dsp and "market" in request_data: request_data["market"] = cls._get_market(request_data["market"], dsp) super().prepare_request_data(params, headers, body, data, **kwargs) async def send_request( self, relative_url: str, method: str = "GET", params: Dict = None, headers: Dict = None, body: Dict or None = None, data: Dict or None = None, data_ttl: int or None = None, # maximum age of the requested data (seconds) **kwargs, ) -> Optional[Dict[str, Any]]: data_ttl = data_ttl or self.config.data_ttl headers = headers or {} if data_ttl and data_ttl.isdigit(): headers.update({"X-Data-Timeout": str(data_ttl)}) return await super().send_request( relative_url, method=method, params=params, headers=headers, body=body, data=data, **kwargs ) async def _get_images_base(self, ids: List[str], url: str, dsp: str, market: str, size: int = None): """Get images urls by list of ids. Args: ids: list of ids. url: url to get data from. dsp: 'apple' or 'spotify' market: Country 2 letter code or generic markets. Returns: List of image url data items. """ request_data = {"ids": ids, "market": self._get_market(market, DSP(dsp))} if size: request_data["image_size"] = size return await self.send_request(method="POST", relative_url=url, body=request_data) @staticmethod def _get_market(market: str, dsp: DSP) -> str: """Check market and return default if it is not valid. Args: market (str): Current market value. dsp (str): DSP enum value. Returns: str: Correct market code. """ if market and len(market) == 2: return market return Market.DEFAULT_PER_DSP[dsp.value] async def get_album(self, dsp: DSP, album_id: str, with_tracks: bool = False, **params): if with_tracks and dsp.value == DSP.SPOTIFY.value: params["fields"] = "tracks_full" response = await self.send_request(f"api/{dsp.value}/v1/albums/{album_id}", params=params, dsp=dsp) return response @prepare_list_arg("ids") @request_in_chunks(chunk_size=200, result_type=ListResult) async def get_albums(self, dsp: DSP, data_ttl: int = None, **params): response = await self.send_request(f"api/{dsp.value}/v1/albums", params=params, dsp=dsp, data_ttl=data_ttl) if dsp.value == DSP.SPOTIFY.value: return response["albums"] if dsp.value == DSP.APPLE.value: return response["data"] raise NotImplementedError() @prepare_list_arg("ids") @request_in_chunks(chunk_size=50, result_type=ListResult) async def get_episodes(self, data_ttl: int = None, **params): response = await self.send_request("api/spotify/v1/episodes", params=params, data_ttl=data_ttl) return response["episodes"] @prepare_list_arg("playlist_ids") async def get_playlists_images(self, playlist_ids: List[str], dsp: DSP, market: str = None, size: int = None): response = await self._get_images_base( playlist_ids, f"api/{dsp.value}/v1/playlists-images", dsp.value, market, size=size ) return response @prepare_list_arg("ids") @request_in_chunks(chunk_size=100, result_type=ListResult, items_key="ids") async def get_spotify_audio_features(self, **params): result = await self.send_request("api/spotify/v1/audio-features", params=params) return result["audio_features"] @prepare_list_arg("ids") async def get_stations(self, dsp: DSP, **params): response = await self.send_request( f"api/{dsp.value}/v1/stations", method="POST", body=params, data_ttl=config.VENDOR_STATIONS_DATA_TTL, dsp=dsp, ) if dsp.value == DSP.APPLE.value: return response["data"] else: raise NotImplementedError() @request_in_chunks(chunk_size=10, result_type=ListResult, items_key="isrc") @request_in_chunks(chunk_size=500, result_type=ListResult) async def get_tracks(self, dsp: DSP, data_ttl: int = None, **params): params = prepare_list_arg_for_names(["ids", "isrc"], params) response = await self.send_request(f"api/{dsp.value}/v1/tracks", params=params, dsp=dsp, data_ttl=data_ttl) if dsp.value == DSP.SPOTIFY.value: return response["tracks"] elif dsp.value == DSP.APPLE.value: return response["data"] raise NotImplementedError() async def get_tracks_v1( self, tracks: List[dict] = None, market: str = None, data_ttl: int = None, ) -> dict: request_data = {"tracks": tracks} if market: request_data["market"] = market return await self.send_request( method="POST", relative_url="api/v1/vendors/search/", body=request_data, data_ttl=data_ttl )