from marshmallow import Schema, ValidationError, fields, validates_schema from playlist.connectors.snowflake import SnowflakeQuery from playlist.queries.schema import ISRC class TotalVsPlaylistTimeSeriesQuerySchema(Schema): isrc = ISRC() global_participant_id = fields.Str() start_date = fields.Date() end_date = fields.Date() stream_countries = fields.List(fields.Str) store_ids = fields.List(fields.Str) permission_label_ids = fields.List(fields.String()) permission_subaccount_ids = fields.List(fields.String()) permission_label_participant_ids = fields.List(fields.String()) transfer_product_ownership_enabled = fields.Boolean(load_default=False) @validates_schema def validate_main_query_arg(self, data, **kwargs): if "isrc" not in data and "global_participant_id" not in data: raise ValidationError("isrc or global_participant_id required") class TotalVsPlaylistTimeSeriesQuery(SnowflakeQuery): query_schema = TotalVsPlaylistTimeSeriesQuerySchema filename = "total_vs_playlist_streams_time_series.sql" @property def table_params(self) -> dict: default_params = { "placement_streams_table": "V_STREAMS_BY_TRACK_PLAYLIST_COUNTRY_FEED_DISTRIBUTOR_DAILY", "streams_table": "V_STREAMS_BY_TRACK_COUNTRY_FEED_DISTRIBUTOR_DAILY", } if "global_participant_id" in self.user_params: default_params["placement_streams_table"] = ( "V_STREAMS_BY_PARTICIPANT_TRACK_PLAYLIST_COUNTRY_FEED_DISTRIBUTOR_DAILY" ) default_params["streams_table"] = ( "V_STREAMS_BY_PARTICIPANT_TRACK_COUNTRY_FEED_DISTRIBUTOR_DAILY" ) return default_params