"""Module provides methods for processing sales data files.""" from datetime import datetime from itertools import islice from math import ceil import re from sqlalchemy import inspect from constants import publishing from constants import sales_file_structure import data_validator import lambda_exceptions from models import phf_mechadmin_track from models import phf_publishing_escrow def identify_company_by_file_headers(file_headers): """Compare list of headers with a predefined one to find a company name. Args: file_headers (set): headers from sales data file. Returns: company_name (str): name of company based on headers of its sales data. Raises: UnexpectedFileHeaders: if headers don't match. """ if file_headers == sales_file_structure.SALES_DATA_HEADERS_PHONOFILE: return sales_file_structure.PHONOFILE elif file_headers == sales_file_structure.SALES_DATA_HEADERS_FINETUNES: return sales_file_structure.FINETUNES else: raise lambda_exceptions.UnexpectedFileHeaders def remove_leading_dash_underscore(value): """Return value without leading dash or underscore if any. Args: value (str|int): value to check the dash. Returns: value (str|int): value without the dash. """ if isinstance(value, str) and value.startswith('_'): return value.lstrip('_') if isinstance(value, str) and '-' in value: return value.split('-', 1)[1] else: return value def remove_non_digits(value): """Return value without non-digit characters. Args: value (str|int): value to check the dash or leading underscore. Returns: value (str|int): value without the dash or leading underscore. """ if isinstance(value, str): return re.sub('[^0-9]', '', value) else: return value def format_track_id_for_phf(original_track_id): """Cast track_id to string and add PHF prefix to it. Args: original_track_id (str or int): id should be updated to phf format. Returns: track_id (str): Provided id with prefix added. """ return 'PHF{}'.format(str(original_track_id)) def calculate_royalty_rate(length_minute, length_seconds): """Calculate royalty rate by length of the track. If track is <= 5 minutes (300 or less seconds) then it's rate is 0.091$. If track is > 5 minutes (301 or more seconds) then it's rate is 0.0175*length in minutes or fraction thereof for those over 5 minutes. Details about this formula are described here: https://secure.harryfox.com/public/RoyaltyRateCalculator.jsp. Args: length_minute (str): Full minutes of track length. length_seconds (str): Seconds of track length. Returns: int: royalty_rate_calculated. """ length = int(length_minute) + ceil(int(length_seconds) / 60) if length <= 5: royalty_rate = publishing.ROYALTY_RATE_SHORT_TRACK else: royalty_rate = publishing.ROYALTY_RATE_PER_MINUTE_LONG_TRACK * length return royalty_rate def extract_publisher(original_publishers): """Extract publisher field from original_publisher of sales data. The value of the original_publishers field trimmed by the first instance of “||” or “ / “ (space, slash, space) characters combination in the string (if present). If there is neither “||” nor “ / “ in the original value, then “publisher” field value will match “original_publishers” field value. Args: original_publishers (str): Publisher info from sales data. Returns: str: trimmed publisher data. """ if original_publishers: separators = { original_publishers.find( sales_file_structure.SEPARATOR_PIPE): sales_file_structure.SEPARATOR_PIPE, original_publishers.find( sales_file_structure.SEPARATOR_SLASH): sales_file_structure.SEPARATOR_SLASH} separators_position_in_line = [ position for position in separators.keys() if position >= 0] if separators_position_in_line: separator_position = min(separators_position_in_line) separator = separators[separator_position] split_publishers = original_publishers.split(separator) split_non_empty_publishers = [ pub for pub in split_publishers if bool(pub)] if len(split_non_empty_publishers): new_publisher = split_non_empty_publishers[0] return new_publisher else: return None return original_publishers def format_company_sales_data_record(record, company): """Reformat uploaded sales data dictionary to database suitable format. Args: record (dict): raw sales data from the file in company specific format. company (str): company name, required for fields mapping. Returns: dict: sales data dictionary in unified format. """ record = {k: v for k, v in record.items() if v is not None} required_fields = sales_file_structure.COMPANY_TO_HEADERS_MAPPING[company] if len(required_fields) != len(record): raise lambda_exceptions.ContentDoesNotMatchHeaders result = {} for db_field_name, company_specific_names in ( sales_file_structure.SALES_DATA_HEADERS_MAPPING.items()): file_field_name = getattr(company_specific_names, company) if file_field_name is None and file_field_name not in record: continue elif file_field_name is not None and file_field_name not in record: # If field should exist in file, but it's not present there. raise lambda_exceptions.ContentDoesNotMatchHeaders result[db_field_name] = record[file_field_name] return result def calculate_release_date(release_date): """Calculate release date from sales data. If the date in the release_date field is missing or not in a date format, then replace it with "Sep 29, 2017". Args: release_date (str): release date from sales file. Returns: str: formatted release date or default release date. """ try: if '/' in release_date: date_format = '%d/%m/%Y' elif '-' in release_date: date_format = '%Y-%m-%d' else: date_format = '%Y%m%d' release_datetime = datetime.strptime(release_date, date_format) if (release_datetime.year < 1900 or release_date.isdigit() and len(release_date) != 8): return publishing.DEFAULT_RELEASE_DATE return release_datetime.strftime('%Y-%m-%d') except (ValueError, TypeError): return publishing.DEFAULT_RELEASE_DATE def calculate_release_period(year, month, periods): """Calculate release date from sales data. If the date in the release_date field is missing or not in a date format, then replace it with "Sep 29, 2017". Args: year (str): release year. month (str): release month. periods (periods): list of dicts, every of each contains one period data. Returns: int: period_id. """ return next(( period['period_id'] for period in periods if period['year'] == int(year) and period['month'] == int(month)), None) def calculate_additional_sales_data_fields( sales_data, company, sales_file_name, periods): """Calculate additional fields by a data present in a file. Args: sales_data (dict): formatted sales data in unified format. company (str): 'phonofile' or 'finetunes' string. sales_file_name (str): name of a file with sales data. periods (list): list of dicts, every of each contains one period data. Returns: dict: sales data extended with additional fields: royalty_calculated, period_id, release_date_calculated, royalty_rate_calculated, track_id, publisher, company, sales_file_name). """ royalty_rate_calculated = calculate_royalty_rate( length_minute=sales_data['length_minute'], length_seconds=sales_data['length_seconds']) qty = int(sales_data['qty']) royalty_calculated = royalty_rate_calculated * qty publisher = extract_publisher(sales_data['original_publishers']) release_date_calculated = calculate_release_date( sales_data['release_date']) period_of_release = calculate_release_period( sales_data['period_year'], sales_data['period_month'], periods) track_id = format_track_id_for_phf(sales_data['original_track_id']) upc = remove_non_digits(sales_data['upc']) isrc = remove_leading_dash_underscore(sales_data['isrc']) calculated_data = { 'period_id': period_of_release, 'royalty_calculated': royalty_calculated, 'release_date_calculated': release_date_calculated, 'royalty_rate_calculated': royalty_rate_calculated, 'track_id': track_id, 'publisher': publisher, 'company': company, 'sales_file_name': sales_file_name, 'qty': qty, 'upc': upc, 'isrc': isrc} sales_data.update(calculated_data) return sales_data def extract_mechadmin_track_data_from_sales_data(sales_data): """Return phf_mechadmin_track related data from sales data. Using sqlalchemy.inspect() get all required columns of table and extract corresponding values from sales_data dictionary. Args: sales_data (dict): formatted sales data in unified. Returns: dict: fields from sales data related to phf_mechadmin_track model. """ result = {} escape_fields = {'id', 'last_modified'} columns_to_fill = { col.name for col in inspect(phf_mechadmin_track.PhfMechadminTrack).c if col.name not in escape_fields} for column in columns_to_fill: result[column] = sales_data.get(column) return result def extract_publishing_escrow_data_from_sales_data(sales_data): """Return phf_publishing_escrow related data from sales data. Using sqlalchemy.inspect() get all required columns of table and extract corresponding values from sales_data dictionary. Args: sales_data (dict): formatted sales data in unified format. Returns: dict: fields from sales data related to phf_publishing_escrow model. """ result = {} escape_fields = { 'phf_transaction_id', 'last_modified', 'ownership', 'active'} columns_to_fill = { col.name for col in inspect( phf_publishing_escrow.PhfPublishingEscrow).c if col.name not in escape_fields} for column in columns_to_fill: result[column] = sales_data.get(column) return result def validate_mandatory_fields(formatted_sales_data): """Validate if all mandatory fields of formatted data are filled.""" mandatory_not_filled = [] for mandatory_field in sales_file_structure.MANDATORY_FIELDS: if not formatted_sales_data.get(mandatory_field): mandatory_not_filled.append(mandatory_field) if mandatory_not_filled: raise lambda_exceptions.MandatoryFieldsMissing(mandatory_not_filled) return 200 def process_sales_data_record( record, company, sales_file_name, periods, line_number): """Process and save sales data record. This function takes raw record from sales data file and processes it: 1) Changes names of columns according to the defined for Talend format using existing fields mapping for provided company name; 2) Calculates additional fields (royalty_calculated, period_id, release_date_calculated, royalty_rate_calculated, track_id, publisher) 3) Extracts metadata for phf_mechadmin_track and phf_publishing_escrow models. 4) Saves track data (if track doesn't exist yet) and saves publishing data to corresponding database. Args: record (dict): row of sales data from uploaded file. company (str): company name, required for fields mapping. sales_file_name (str): name of file with sales data. periods (list): list of dicts, every of each contains one period data. Returns: tuple: track and escrow data extracted from record. """ formatted_sales_data = format_company_sales_data_record(record, company) validate_mandatory_fields(formatted_sales_data) data_validator.validate_phf_sales_data_record( formatted_sales_data, line_number) full_sales_data = calculate_additional_sales_data_fields( formatted_sales_data, company, sales_file_name, periods) mechadmin_track_data = ( extract_mechadmin_track_data_from_sales_data(full_sales_data)) escrow_data = extract_publishing_escrow_data_from_sales_data( full_sales_data) return mechadmin_track_data, escrow_data def bulk_insert_new_phf_metadata(tracks_data, escrows_data): """Insert required data to db. The functions cleans tracks data by removing already existing in db tracks, then it validates data, creates tracks records and publishing records for it. Args: tracks_data (dict): dict in format {'track_id': track_data}. escrows_data (list): list with escrows data that should be inserted. Returns: int: status of inserting. """ existing_tracks = phf_mechadmin_track.get_tracks_by_list_of_track_ids( list(tracks_data.keys())) track_ids = [t['track_id'] for t in existing_tracks] # remove existing tracks from to_create list. tracks_data = { t_id: tracks_data[t_id] for t_id in tracks_data.keys() if t_id not in track_ids} phf_mechadmin_track.bulk_insert_phf_mechadmin_track( list(tracks_data.values())) phf_publishing_escrow.bulk_insert_phf_publishing_escrow(escrows_data) return 200 def guess_file_segment_encoding(file_path, start_index, stop_index): """Try to parse sales file with expected encodings for choosing one. Args: file_path (str): file that should be processed. start_index (int): on what line to start file processing. stop_index (int): on what line to finish file processing. Returns: str: fitting encoding. Raises: FileEncodingNotIdentified: if fitting encoding wasn't found. """ detected_encoding = None for encoding in sales_file_structure.ENCODINGS_TO_TRY: try: with open( file_path, encoding=encoding, errors='strict', newline='\r\n') as sales_file: file_segment = islice(sales_file, start_index, stop_index) # Iterate through every record. The file is opened in a strict # mode, so if one of records contains encoding issues then # UnicodeDecodeError will be raised immediately. all(file_segment) detected_encoding = encoding break except UnicodeDecodeError: continue if not detected_encoding: raise lambda_exceptions.FileEncodingNotIdentified return detected_encoding def extract_file_headers(file_path): """Get headers of file. Args: file_path (str): file that should be processed. Returns: str: file headers. """ with open( file_path, encoding=sales_file_structure.HEADERS_DEFAULT_ENCODING, errors='strict', newline='\r\n') as sales_file: headers = re.sub(r'(?