"""Model for track streams.""" from ddtrace import tracer from oto import response as oto_response 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 from sound_recordings.utils.streams import AGGREGATE_STREAMS_FIELDS, STREAMS_ALL_FIELDS BULK_TRACK_STREAMS_FIELDS = [ "store_id", "date", "isrc", "streams", "streams_with_skips", "skips", "saves", ] SQLLoader = snowflake.SQLLoader(__file__) @cache_in_redis(ttl=cache.SECONDS_PER_HOUR) @tracer.wrap(name="get_streams_bulk") def get_streams_bulk( permissions_filter, isrcs, distributors, countries=[], store_ids=[], start_date=None, end_date=None, ): """Get streams for given ISRC by store. Args: permissions_filter (dict): dict containing resources users can access isrcs ([str]): ISRCs of track to fetch streams for. distributors (str[]): List of distributors names countries(str[]): List of country codes to filter streams 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 streams per store. """ if len(store_ids) == 0: store_ids = store_availability.get_store_ids() else: store_ids = list( set(store_ids).intersection(store_availability.get_store_ids()) ) if len(store_ids) == 0: return oto_response.Response([]) params = { **permissions_filter, "isrcs": isrcs, "store_ids": store_ids, "start_date": start_date, "end_date": end_date, "distributors": distributors, } if len(countries) > 0: params["country_codes"] = countries query_table = "streams_by_store_by_country_bulk" else: query_table = "streams_by_store_bulk" 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 oto_response.Response( [dict(zip(BULK_TRACK_STREAMS_FIELDS, record)) for record in stream_records] ) @cache_in_redis(ttl=cache.SECONDS_PER_HOUR) @tracer.wrap(name="get_streams_all_time_bulk") def get_streams_all_time_bulk( permissions_filter, isrcs, distributors, countries=[], store_ids=[] ): """Get all time streams for given ISRCs by store. Args: permissions_filter (dict): dict containing resources users can access isrc ([str]): ISRCs of track to fetch streams for. distributors (str[]): List of distributors names countries(str[]): List of country codes to filter streams by store_ids (list): List of store ids to filter by Returns: list: aggregated all time streams. """ if len(store_ids) == 0: store_ids = store_availability.get_store_ids() else: store_ids = list( set(store_ids).intersection(store_availability.get_store_ids()) ) if len(store_ids) == 0: return oto_response.Response([]) params = { **permissions_filter, "isrcs": isrcs, "store_ids": store_ids, "distributors": distributors, } if len(countries) > 0: params["country_codes"] = countries query_table = "streams_all_time_by_country_bulk" else: query_table = "streams_all_time_bulk" sql = SQLLoader.load_query(query_table) sql = format_sql.format_with_permissions_filter(sql, permissions_filter) stream_records = snowflake.fetchall(sql, params) fields = [ "store_id", "isrc", "all_time", "growth_percentage", ] return oto_response.Response( [dict(zip(fields, record)) for record in stream_records] ) def _get_streams_all_result( permissions_filter, isrc, query, fields, distributors, countries=[], store_ids=[], start_date=None, end_date=None, ): params = { **permissions_filter, "isrc": isrc, "distributors": distributors, "store_ids": store_ids, "start_date": start_date, "end_date": end_date, } if len(countries) > 0: params["country_codes"] = countries query = "streams_all_by_country" sql = SQLLoader.load_query(query) 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(fields, record)) for record in records] @tracer.wrap(name="get_streams_all") @cache_in_redis(ttl=cache.SECONDS_PER_HOUR) def get_streams_all( permissions_filter, isrc, distributors, countries=[], store_ids=[], start_date=None, end_date=None, ): """Get aggregate streams timeseries for given ISRC. Args: permissions_filter (dict): dict containing resources users can access isrc (str): ISRC of track to fetch streams for. distributors (list): List of distributors names countries (list): List of country codes to filter 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 dictionaries with keys: date (str): Date in YYYY-MM-DD format streams (int): Total streams for date streams_with_skips (int): Total streams with skips for date saves (int): Total saves for date skips (int): Total skips for date """ if len(store_ids) == 0: store_ids = store_availability.get_store_ids() else: store_ids = sorted( list(set(store_ids).intersection(store_availability.get_store_ids())) ) if len(store_ids) == 0: return [] return _get_streams_all_result( permissions_filter, isrc, "streams_all", STREAMS_ALL_FIELDS, distributors, countries, store_ids, start_date, end_date, ) @tracer.wrap(name="get_aggregate_streams") @cache_in_redis(ttl=cache.SECONDS_PER_HOUR) def get_aggregate_streams( permissions_filter, isrcs, distributors, countries=[], store_ids=[] ): """Get aggregate streams for given ISRCs. Args: permissions_filter (dict): dict containing resources users can access isrcs ([str]): List of ISRCs to fetch streams for. distributors (list): List of distributors names countries (list): List of country codes to filter 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 dictionaries with keys: isrc (str): ISRC streams_all_time (int): Total streams for date streams_7_day_growth (int): Streams growth over 7 days """ if len(store_ids) == 0: store_ids = store_availability.get_store_ids() else: store_ids = sorted( list(set(store_ids).intersection(store_availability.get_store_ids())) ) if len(store_ids) == 0: return [] params = { **permissions_filter, "isrcs": isrcs, "distributors": distributors, "store_ids": store_ids, } query = "aggregate_streams_bulk" if len(countries) > 0: params["country_codes"] = countries query = "aggregate_streams_by_country_bulk" sql = SQLLoader.load_query(query) sql = format_sql.format_with_permissions_filter(sql, permissions_filter) records = snowflake.fetchall(sql, params) return [dict(zip(AGGREGATE_STREAMS_FIELDS, record)) for record in records]