"""set-pricing.""" import uuid from typing import Dict, List, Union from ddex_ingester_common.constants.release_type import VIDEO_RELEASE_TYPES from ddex_ingester_common.constants.status import (IN_CONTENT, TRANSFER_TO_CONTENT) from ddex_ingester_common.helpers.s3_ddex import load_ddex_json from ddex_ingester_common.lambda_exceptions import PricingException from ddex_ingester_common.logging import utils as logging_utils from ddex_ingester_common.models.s3.body import Body as S3Context from ddex_ingester_common.models.s3.deal import Deal as S3Deal from ddex_ingester_common.models.s3.deal_term import DealTerm as S3DealTerm from ddex_ingester_common.models.s3.track import Track as S3Track from ddex_ingester_common.models.state_machine.body import \ Body as StateMachineContext from ddex_ingester_common.models.state_machine.track import \ Track as StateMachineTrack from ddex_ingester_common.schemas.s3_schema import S3Schema from ddex_ingester_common.schemas.state_machine_schema import \ StateMachineSchema from lambdacommon.graphql import graphql import config from config import graphql_gateway from constants import queries from constants.constants import (MUSIC_ALBUM, MUSIC_TRACK, VIDEO, VIDEO_FRONT_TIER) logger = logging_utils.get_logger(config.app_logger) def handler(event, context): """Lambda entry point.""" try: context = StateMachineSchema().load(event) s3_data = S3Schema().load(load_ddex_json(event)) correlation_id = context.correlation_id or str(uuid.uuid4()) context.correlation_id = correlation_id logging_utils.update_logger_correlation_id(logger, correlation_id) logging_utils.update_logger_with_message_ids( logger, context.message_id, context.message_thread_id, context.execution_name ) graphql_gateway.set_headers( { 'Orchard-User-Id': config.OA_USER, 'Correlation-Id': correlation_id, } ) if context.product.release_type in VIDEO_RELEASE_TYPES: # TODO: confirm pricing family and tier values video_pricing_data = format_update_pricing_payload( product_id=context.product.product_id, pricing_family=VIDEO, orchard_pricing_tier_id=VIDEO_FRONT_TIER) # If we're in an unsupported workflow, skip the update. if context.product.status not in [IN_CONTENT, TRANSFER_TO_CONTENT]: graphql_gateway.execute( queries.UPDATE_PRODUCT_PRICING_TIER, video_pricing_data) else: update_music_album_product_pricing_tier(context, s3_data) update_music_track_pricing_tier(context, s3_data) except graphql.GraphQLError as err: raise PricingException('Graphql error') from err # Only one of the parallel steps needs to return the context return StateMachineSchema().dump(context) def update_music_album_product_pricing_tier( context: StateMachineContext, s3_data: S3Context): """Update the album pricing tier of a product.""" product_deal_terms = get_product_deal_terms(s3_data) or [] product_price_tier = next(( deal_term.price_type for deal_term in product_deal_terms if deal_term.price_type), None ) if not product_price_tier or not str.isdigit(product_price_tier): logger.info('Skipped updating product pricing tier') return payload = format_update_pricing_payload( product_id=context.product.product_id, pricing_family=MUSIC_ALBUM, orchard_pricing_tier_id=int(product_price_tier)) logger.info( f'Running update product pricing tier with payload:\n{payload}') graphql_gateway.execute( queries.UPDATE_PRODUCT_PRICING_TIER, payload) def update_music_track_pricing_tier( context: StateMachineContext, s3_data: S3Context): """Update the track pricing tier of a product.""" # Set the most common track price type at the product level # The others are set as track overrides price_tier = get_most_common_track_price_tier(s3_data) track_overrides = get_pricing_tier_track_overrides( context, s3_data, price_tier) if not price_tier or not str.isdigit(price_tier): logger.info('Skipped updating track pricing tier') return payload = format_update_pricing_payload( product_id=context.product.product_id, pricing_family=MUSIC_TRACK, orchard_pricing_tier_id=int(price_tier), track_overrides=track_overrides) logger.info( f'Running update track pricing tier with payload:\n{payload}') graphql_gateway.execute( queries.UPDATE_PRODUCT_PRICING_TIER, payload) def get_pricing_tier_track_overrides( context: StateMachineContext, s3_data: S3Context, common_price_tier: int) -> Dict: """Build track overrides for update product pricing tier mutation.""" track_overrides = [] for track in context.tracks or []: track_price_tier = get_track_price_tier(track, s3_data.deals) if track_price_tier and str.isdigit(track_price_tier) and \ track_price_tier != common_price_tier: track_overrides.append({ 'trackId': track.tuid, 'orchardPricingTier': int(track_price_tier), }) return track_overrides def get_most_common_track_price_tier(s3_data: S3Context) -> int: """Get the most common track price tier.""" price_tier_counter = {} for track in s3_data.tracks or []: price_tier = get_track_price_tier(track, s3_data.deals) if price_tier: if price_tier not in price_tier_counter: price_tier_counter[price_tier] = 0 price_tier_counter[price_tier] += 1 max_price_tier = None max_price_tier_counter = 0 for price_tier in price_tier_counter: if price_tier_counter[price_tier] > max_price_tier_counter: max_price_tier = price_tier max_price_tier_counter = price_tier_counter[price_tier] return max_price_tier def get_track_price_tier( track: Union[S3Track, StateMachineTrack], deals: List[S3Deal]) -> int: """Get a track's price type from the deal.""" for deal in deals: for deal_term in deal.deal_terms: if deal_term.price_type and \ track.release_reference in deal.release_references: return deal_term.price_type return None def format_update_pricing_payload( product_id: int, pricing_family: int, orchard_pricing_tier_id: int, track_overrides: List[Dict] = None) -> Dict: """Format pricing payload. Args: product_id (int): Product ID pricing_family (str): Pricing Family orchard_pricing_tier_id (int): Orchard pricing tier ID Returns: dict """ payload = { 'data': { 'productId': product_id, 'orchardPricingTier': orchard_pricing_tier_id, 'pricingFamily': pricing_family } } if track_overrides is not None: payload['data']['trackOverrides'] = track_overrides return payload def get_product_deal_terms(s3_data: S3Context) -> List[S3DealTerm]: """Get deal term for R0 if not then R1 else None.""" for deal in s3_data.deals: if 'R0' in deal.release_references: return deal.deal_terms if 'R1' in deal.release_references: return deal.deal_terms return None