"""demographic stream data model.""" from sound_recordings.connectors import snowflake from sound_recordings.constants import cache from sound_recordings.constants.demographics import ( EIGHTEEN_TWENTYTWO, FEMALE, FORTYFIVE_FIFTYNINE, MALE, OVER_SIXTY, THIRTYFIVE_FORTYFOUR, TWENTYEIGHT_THIRTYFOUR, TWENTYTHREE_TWENTYSEVEN, UNDER_18, UNKNOWN_AGE, UNKNOWN_GENDER, ) from sound_recordings.utils import format_sql, store_availability from sound_recordings.utils.cache import cache_in_redis SQLLoader = snowflake.SQLLoader(__file__) DEMOGRAPHICS_FIELDS = [ UNDER_18, EIGHTEEN_TWENTYTWO, TWENTYTHREE_TWENTYSEVEN, TWENTYEIGHT_THIRTYFOUR, THIRTYFIVE_FORTYFOUR, FORTYFIVE_FIFTYNINE, OVER_SIXTY, UNKNOWN_AGE, MALE, FEMALE, UNKNOWN_GENDER, ] @cache_in_redis(ttl=cache.SECONDS_PER_HOUR) def get_demographics( permissions_filter, isrc, countries, store_ids, start_date, end_date, distributors ): """Get demographics for a given ISRC by store. Args: permissions_filter (dict): dict containing resources users can access isrc (str): ISRC of track to fetch downloads for. 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 distributors (list): List of distributors names Returns: list: stream demographics per store by ISRC """ store_ids = _get_store_ids(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, } query_table = "demographics" if len(countries) > 0: params["country_codes"] = countries query_table = "demographics_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 ) breakdown = snowflake.fetchone(sql, params) if not breakdown: return None return dict(zip(DEMOGRAPHICS_FIELDS, breakdown)) def _get_store_ids(store_ids): if len(store_ids) == 0: store_ids = store_availability.get_demographic_store_ids() else: store_ids = sorted( list( set(store_ids).intersection( store_availability.get_demographic_store_ids() ) ) ) return store_ids