"""Common module for streams totals lambda functions.""" import csv import smart_open from sosmodels import streams_vendor_totals from sosmodels import streams_track_totals import consts # noqa import util def load(s3_key_url, totals_type): """Load totals data to passed model. Args: s3_key_url (str): S3 url for loading file. totals_type (str): Totals type (vendor or track). """ store = util.identify_retailer(s3_key_url) TOTALS_MODEL_MAPPING = { consts.VENDOR: { consts.SPOTIFY: streams_vendor_totals.StreamsVendorTotalsSpotify, consts.APPLE_MUSIC: streams_vendor_totals.StreamsVendorTotalsAppleMusic }, consts.TRACK: { consts.SPOTIFY: streams_track_totals.StreamsTrackTotalsSpotify, consts.APPLE_MUSIC: streams_track_totals.StreamsTrackTotalsAppleMusic } } try: totals_model = TOTALS_MODEL_MAPPING[totals_type][store] except KeyError: raise KeyError('Incorrect totals type.') hash_key_column = ( 'vendor_key' if totals_type == consts.VENDOR else 'vendor_track_key') columns = [ 'date', 'ttl', hash_key_column, 'overall_number_of_streams', 'streams_from_passive_discovery', 'streams_from_active_discovery', 'streams_from_collection'] with totals_model.batch_write() as batch: with smart_open.smart_open(s3_key_url) as lines: for totals in csv.reader(lines, dialect='escaped'): assert len(columns) == len(totals) item_data = dict(zip(columns, totals)) for column in columns: item_data[column] = getattr( totals_model, column ).deserialize(item_data[column]) batch.save(totals_model(**item_data)) def delete(s3_key_url): """Delete totals of passed model. Args: s3_key_url (str): S3 url for loading file. """ totals_type = util.identify_totals_type(s3_key_url) store = util.identify_retailer(s3_key_url) TOTALS_MODEL_MAPPING = { consts.VENDOR: { consts.SPOTIFY: streams_vendor_totals.StreamsVendorTotalsSpotify, consts.APPLE_MUSIC: streams_vendor_totals.StreamsVendorTotalsAppleMusic }, consts.TRACK: { consts.SPOTIFY: streams_track_totals.StreamsTrackTotalsSpotify, consts.APPLE_MUSIC: streams_track_totals.StreamsTrackTotalsAppleMusic } } try: totals_model = TOTALS_MODEL_MAPPING[totals_type][store] except KeyError: raise KeyError('Incorrect totals type.') with smart_open.smart_open(s3_key_url) as lines: models = totals_model.batch_get(csv.reader(lines, dialect='escaped')) for model in models: model.delete()