"""Model for track downloads.""" from ddtrace import tracer from sound_recordings.connectors import snowflake from sound_recordings.constants import cache from sound_recordings.utils import format_sql, store_availability from sound_recordings.utils.cache import cache_in_redis TRACK_DOWNLOADS_FIELDS = ["store_id", "date", "downloads"] TRACK_DOWNLOADS_BY_COUNTRY_FIELDS = ["country_code", "date", "downloads"] TRACK_DOWNLOADS_BY_PRODUCT = ["product_id", "date", "downloads"] SQLLoader = snowflake.SQLLoader(__file__) @tracer.wrap(name="get_downloads") @cache_in_redis(ttl=cache.SECONDS_PER_HOUR) def get_downloads( permissions_filter, isrc, distributors, countries=[], store_ids=[], start_date=None, end_date=None, ): """Get downloads for given ISRC by store. Args: permissions_filter (dict): dict containing resources users can access isrc (str): ISRC of track to fetch downloads for. distributors (list): List of distributors names countries(str[]): List of country codes to filter downloads by store_ids (list): List of store ids to filter by start_date (datetime.date): Start date end_date (datetime.date): End date Returns: list: track downloads per store. """ if len(store_ids) == 0: store_ids = store_availability.get_download_store_ids() else: store_ids = sorted( list( set(store_ids).intersection(store_availability.get_download_store_ids()) ) ) if len(store_ids) == 0: return [] params = { **permissions_filter, "isrc": isrc, "store_ids": store_ids, "start_date": start_date, "end_date": end_date, "distributors": distributors, } if len(countries) > 0: params["country_codes"] = countries query_table = "downloads_by_country" else: query_table = "downloads" sql = SQLLoader.load_query(query_table) sql = format_sql.format_with_permissions_filter( sql, permissions_filter, start_date=start_date, end_date=end_date ) stream_records = snowflake.fetchall(sql, params) return [dict(zip(TRACK_DOWNLOADS_FIELDS, record)) for record in stream_records] @tracer.wrap(name="get_downloads_by_country") @cache_in_redis(ttl=cache.SECONDS_PER_HOUR) def get_downloads_by_country( permissions_filter, isrc, distributors, countries=[], store_ids=[], start_date=None, end_date=None, ): """Get downloads for given ISRC by country. Args: permissions_filter (dict): dict containing resources users can access isrc (str): ISRC of track to fetch downloads for. distributors (list): List of distributors names countries(str[]): List of country codes to filter downloads by store_ids (list): List of store ids to filter by start_date (datetime.date): Start date end_date (datetime.date): End date Returns: list: track downloads per country. """ if len(store_ids) == 0: store_ids = store_availability.get_download_store_ids() else: store_ids = sorted( list( set(store_ids).intersection(store_availability.get_download_store_ids()) ) ) if len(store_ids) == 0: return [] params = { **permissions_filter, "isrc": isrc, "store_ids": store_ids, "start_date": start_date, "end_date": end_date, "distributors": distributors, } if len(countries) > 0: params["country_codes"] = countries query_table = "downloads_group_by_country_for_country" else: query_table = "downloads_group_by_country" sql = SQLLoader.load_query(query_table) sql = format_sql.format_with_permissions_filter( sql, permissions_filter, start_date=start_date, end_date=end_date ) stream_records = snowflake.fetchall(sql, params) return [ dict(zip(TRACK_DOWNLOADS_BY_COUNTRY_FIELDS, record)) for record in stream_records ] @tracer.wrap(name="get_downloads_by_product") @cache_in_redis(ttl=cache.SECONDS_PER_HOUR) def get_downloads_by_product( permissions_filter, isrc, distributors, countries=[], store_ids=[], start_date=None, end_date=None, ): """Get downloads by product for given ISRC. Args: permissions_filter (dict): dict containing resources users can access isrc (str): ISRC of track to fetch downloads for. distributors (list): List of distributors names countries(str[]): List of country codes to filter downloads by store_ids (list): List of store ids to filter by start_date (datetime.date): Start date end_date (datetime.date): End date Returns: List of dicts with keys product_id (int): Product ID date (str): Date downloads (int | None): Total downloads for date """ if len(store_ids) == 0: store_ids = store_availability.get_download_store_ids() else: store_ids = sorted( list( set(store_ids).intersection(store_availability.get_download_store_ids()) ) ) if len(store_ids) == 0: return [] params = { **permissions_filter, "isrc": isrc, "distributors": distributors, "store_ids": store_ids, "start_date": start_date, "end_date": end_date, } query_table = "downloads_by_product" if len(countries) > 0: params["country_codes"] = countries query_table = "downloads_by_product_by_country" sql = SQLLoader.load_query(query_table) sql = format_sql.format_with_permissions_filter( sql, permissions_filter, start_date=start_date, end_date=end_date ) records = snowflake.fetchall(sql, params) return [dict(zip(TRACK_DOWNLOADS_BY_PRODUCT, record)) for record in records]