import datetime as dt from abc import abstractmethod from typing import Any, Collection, Dict, Generator, Iterable, List, Mapping, Optional, Tuple from slugify import slugify from app_types import Record, S3Path from constants import CHANGE_DAYS, CHART_DAYS from utils.interpolation import interpolate_series from utils.perf_counter import timeit from ... import BaseStep from ...mixins import CSVMixin, S3Mixin from ..constants import AVAILABLE_COUNTRIES, N_A_SLUG, SLUGIFY_PROPS __all__ = ["SearchTransformBase"] class SearchTransformBase(S3Mixin, CSVMixin, BaseStep): def __init__(self, **kwargs): super().__init__(**kwargs) if not CHART_DAYS or not CHANGE_DAYS: raise ValueError("Invalid CHART_DAYS or CHANGE_DAYS") self.charts_days: int = CHART_DAYS self.change_days: int = CHANGE_DAYS priority = 3 batch_size: int = 5000 csv_normalizers = {"id": str, "cm_artist_id": str, "name": str, "artist_name": str, "label": str} max_historical_days_diff: int = 60 @property @abstractmethod def raw_data_prefixes(self) -> Tuple[S3Path, ...]: pass @property @abstractmethod def data_path(self) -> S3Path: pass @property @abstractmethod def index_data_path(self) -> S3Path: pass @property @abstractmethod def dsp_to_interpolate(self) -> Tuple[str, ...]: pass @property @abstractmethod def dsp_to_extend(self) -> Tuple[str, ...]: pass @property @abstractmethod def dsp_change_trends_props_map(self) -> Mapping[str, Tuple[str, ...]]: pass @property @abstractmethod def top_props(self) -> Tuple[str, ...]: pass @property @abstractmethod def dsp_index_props_map(self) -> Mapping[str, Tuple[str, ...]]: pass @property @abstractmethod def index_props(self) -> Tuple[str, ...]: pass @property def skip_release_date(self) -> dict: return {} def process(self): self.wipe_folder(self.data_path) self.wipe_folder(self.index_data_path) self.logger.info("Processing search data") gens = {str(prefix).split("/")[-1]: self._read_raw_data(prefix) for prefix in self.raw_data_prefixes} records_queue = {} data_records = [] load_index = 0 while True: records = {dsp: next(gen) for dsp, gen in gens.items() if dsp not in records_queue} records.update(records_queue) # exclude non performing dsp for dsp in list(records): if records[dsp] is None: del gens[dsp], records[dsp] if not records: break # records are not ready to be inserted, split into 2 groups records_queue = {} min_id = str(min([int(res["id"]) for res in records.values()])) for dsp in list(records): if records[dsp]["id"] != min_id: records_queue[dsp] = records.pop(dsp) data_records.append(self._get_data_record(raw_records=records)) if len(data_records) >= self.batch_size: self._load_records(data_records, load_index) data_records = [] load_index += 1 self._load_records(data_records, load_index) def _get_data_record(self, raw_records: Dict[str, Record]) -> Record: data_record: Dict[str, Any] = {} release_date: Optional[dt.date] = None for dsp, raw_record in raw_records.items(): for prop in self.top_props: value = raw_record.pop(prop) if prop == "release_date": release_date = value if prop not in data_record: data_record[prop] = value data_record[dsp] = self._transform_raw_record(raw_record, dsp, release_date) self._transform_data_record(data_record) return data_record def _get_index_record(self, data_record: Record) -> Record: index_record = {} for key, value in data_record.items(): if key in self.index_props: index_record[key] = value continue if key in self.dsp_index_props_map: index_record[key] = {} for dsp_key in self.dsp_index_props_map[key]: index_record[key][dsp_key] = value.get(dsp_key) return index_record def _dump_record(self, record: Record, id_: str) -> Dict[str, str]: return {"id": id_, "value": self.format_json(record)} def _load_records(self, records: Iterable[Record], index: int): header = ("id", "value") data_records = [] index_records = [] for record in records: data_records.append(self._dump_record(record, record["id"])) index_records.append(self._dump_record(self._get_index_record(record), record["id"])) self._write_records(data_records, header, f"{self.data_path}/data_{index}.csv") self._write_records(index_records, header, f"{self.index_data_path}/data_{index}.csv") def _write_records(self, records: List[Dict[str, Any]], header: Collection[str], key: S3Path): self.put_s3_object(key, self.write_csv(records, header)) def _transform_raw_record(self, raw_record: Record, dsp: str, release_date: Optional[dt.date]) -> Record: record = raw_record.copy() charts = [k for k, v in record.items() if k.endswith("_chart")] if dsp in self.dsp_to_interpolate: self._interpolate_charts(record, charts) if dsp in self.dsp_to_extend: self._extend_charts(record, raw_record, charts, dsp, release_date) return record def _transform_data_record(self, record: Record): self._slugify_props(record) self._validate_country(record) self._add_slugified_name(record) @timeit() def _slugify_props(self, record: Record): for to_slugify in SLUGIFY_PROPS: if to_slugify not in record: continue if isinstance(record[to_slugify], str): record[to_slugify] = slugify(record[to_slugify]) elif isinstance(record[to_slugify], list): record[to_slugify] = list(set(map(slugify, record[to_slugify]))) @timeit() def _validate_country(self, record: Record): if "country" not in record: self.logger.warning(f"Country field not in record: {record.get('id')}") return if record["country"] not in AVAILABLE_COUNTRIES: self.logger.warning(f"{record['country']} is not a valid country, falling back to {N_A_SLUG}") record["country"] = N_A_SLUG @timeit() def _interpolate_charts(self, record: Record, charts: Collection[str]): for chart in charts: dates = record.pop(f"{chart}_dates", None) historical_date = record.pop(f"{chart}_historical_date", None) historical_value = record.pop(f"{chart}_historical_value", None) values = record[chart] if not dates or not values or len(values) >= self.charts_days: continue if historical_date is not None and self.key.timestamp.to_date() - historical_date > dt.timedelta( self.max_historical_days_diff ): historical_date = None end_date = self.key.timestamp.to_date() start_date = end_date - dt.timedelta(days=self.charts_days - 1) record[chart] = interpolate_series(start_date, end_date, dates, values, historical_date, historical_value) def _get_last_chart_date(self) -> dt.datetime: """ last_chart_date shouldn't be today's date, it should be yesterday """ return self.key.timestamp.to_date() - dt.timedelta(days=1) def change_trend_counter(self, values: list, release_date: Optional[dt.date], days_gap: int, min_len_values: int): """ min_len_values: int min_len_values - should be used to count change parameter correctly if release_date between 0 and 7(28) days, when release_date in days_gap """ change = None trend = None trend_chart_value = None change_avg = None if values and len(values) >= min_len_values: # count change here # release_date could be from future if ( release_date and len(values) < days_gap + 1 and dt.timedelta(0) < self.key.timestamp.to_date() - release_date < dt.timedelta(days_gap + 1) ): change = values[-1] else: trend_chart_value = values[-min(days_gap + 1, len(values))] change = values[-1] - trend_chart_value # Trend and change_avg should be calculated for zero change value! # Calculate trend # only when we have "full" chart data for 7 (28) days if trend_chart_value and change is not None and len(values) >= days_gap: # last days trend: trend = round(change / trend_chart_value * 100, 4) # Calculate change_avg # for 7 (28) days we need at least 2 (8) items in chart values list if change is not None and len(values) >= min_len_values: # last days average change: change_avg = round(sum(values[-days_gap:]) / min(len(values), days_gap), 4) return change, trend, change_avg @timeit() def _extend_charts( self, record: Record, raw_record: Record, charts: Collection[str], dsp: str, release_date: Optional[dt.date] ): for chart in charts: chart_base = chart.replace("_chart", "") values = record[chart] # Fill last value chart_base_value = values[-1] if values else None # if no data in values try to get historical value at least if not chart_base_value: # use get because historical_value is optional key chart_base_value = raw_record.get(f"{chart}_historical_value") # fill chart_base value record[chart_base] = chart_base_value # Fill last chart date record["last_chart_date"] = self._get_last_chart_date() # Fill "change" and "trend" if chart_base not in self.dsp_change_trends_props_map.get(dsp, ()): continue if release_date and chart_base in self.skip_release_date.get(dsp, ()): release_date = None change, trend, change_avg = self.change_trend_counter( values=values, release_date=release_date, days_gap=self.change_days, min_len_values=2 ) change_28, trend_28, change_avg_28 = self.change_trend_counter( values=values, release_date=release_date, days_gap=self.charts_days, min_len_values=8 ) record[f"{chart_base}_change_7"] = change record[f"{chart_base}_trend_7"] = trend record[f"{chart_base}_change_avg_7"] = change_avg record[f"{chart_base}_change_28"] = change_28 record[f"{chart_base}_trend_28"] = trend_28 record[f"{chart_base}_change_avg_28"] = change_avg_28 def _read_raw_data(self, prefix: str) -> Generator: for key in self.get_s3_keys(prefix): for record in self.read_csv(self.get_s3_object(key)): yield record yield None def _add_slugified_name(self, record: Record): if isinstance(record["name"], str): slugified_name = slugify(record["name"], separator="") else: self.logger.warning( f"name field has type {type(record['name'])} which is incompatible. Falling back to None" ) slugified_name = None record["slugified_name"] = slugified_name