import datetime as dt from pathlib import PurePosixPath from typing import Any, Iterable, Collection, Generator, Mapping, Tuple from airflow.models import BaseOperator from airflow.providers.amazon.aws.hooks.s3 import S3Hook from airflow.utils.context import Context from slugify import slugify from common.macros.charts import get_max_chart_date from common.utils.types import Record, S3Path, Records from common.utils.json import format_json from common.config.app import CSV_DELIMITER, CSV_NULL_VALUE, CSV_QUOTE_CHAR, S3_ARTIFACTS_BUCKET, N_A_SLUG, \ AVAILABLE_COUNTRIES, CSV_NORMALIZERS from common.config import search as search_config __all__ = ["TransformSearchOperator"] class TransformSearchOperator(BaseOperator): def __init__( self, raw_data_path: S3Path, agg_data_path: S3Path, index_data_path: S3Path, top_props: Tuple[str, ...], index_props: Tuple[str, ...], dsp_index_props_map: Mapping[str, Any], dsp_to_interpolate: Tuple[str, ...], dsp_to_extend: Tuple[str, ...], dsp_change_trends_props_map: Mapping[str, Any], aws_conn_id: str = S3Hook.default_conn_name, **kwargs ): super().__init__(**kwargs) self.aws_conn_id = aws_conn_id self.batch_size: int = 5000 self.raw_data_path = raw_data_path self.agg_data_path = agg_data_path self.index_data_path = index_data_path self.top_props = top_props self.index_props = index_props self.dsp_index_props_map = dsp_index_props_map self.dsp_to_interpolate = dsp_to_interpolate self.dsp_to_extend = dsp_to_extend self.dsp_change_trends_props_map = dsp_change_trends_props_map self.s3 = None self.interpolator = None self.csv = None self.pd = None self.current_date = None def pre_execute(self, context: Context): from common.utils.s3 import S3 from common.utils.interpolation import Interpolator from common.utils.csv import CSV from common.macros.generic import get_run_key import pandas as pd conn = S3Hook(aws_conn_id=self.aws_conn_id).get_conn() self.s3 = S3(conn, run_id=get_run_key(context["dag_run"]), bucket_name=S3_ARTIFACTS_BUCKET) self.interpolator = Interpolator self.csv = CSV( delimiter=CSV_DELIMITER, quote_char=CSV_QUOTE_CHAR, null_value=CSV_NULL_VALUE, normalizers=CSV_NORMALIZERS ) self.pd = pd self.current_date = get_max_chart_date(context["dag_run"]) # Prefixing s3 path self.raw_data_path = self.s3.get_prefixed_key(self.raw_data_path) self.agg_data_path = self.s3.get_prefixed_key(self.agg_data_path) self.index_data_path = self.s3.get_prefixed_key(self.index_data_path) def execute(self, context: Context): self.log.info("Processing search data") self.s3.wipe_by_prefix(self.agg_data_path) self.s3.wipe_by_prefix(self.index_data_path) raw_data_prefixes = [PurePosixPath(prefix) for prefix in self.s3.get_common_keys(self.raw_data_path)] gens = {str(prefix).split("/")[-1]: self._read_raw_data(prefix) for prefix in 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: dt.date | None = 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 @staticmethod def _dump_record(record: Record, id_: str) -> dict[str, str]: return {"id": id_, "value": 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.agg_data_path}/data_{index}.csv") self._write_records(index_records, header, f"{self.index_data_path}/data_{index}.csv") def _write_records(self, records: Records, header: Collection[str], key: S3Path): self.s3.write_object(key, self.csv.dump(records, header)) def _transform_raw_record(self, raw_record: Record, dsp: str, release_date: dt.date | None) -> 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) def _slugify_props(self, record: Record): for to_slugify in search_config.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]))) def _validate_country(self, record: Record): if "country" not in record: self.log.warning(f"Country field not in record: {record.get('id')}") return if record["country"] not in AVAILABLE_COUNTRIES: self.log.warning(f"{record['country']} is not a valid country, falling back to {N_A_SLUG}") record["country"] = N_A_SLUG 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) >= search_config.CHART_DAYS: continue if historical_date is not None and self.current_date - historical_date > dt.timedelta( search_config.MAX_HISTORICAL_DAYS_DIFF ): historical_date = None end_date = self.current_date start_date = end_date - dt.timedelta(days=search_config.CHART_DAYS - 1) record[chart] = self.interpolator.interpolate_series(start_date, end_date, dates, values, historical_date, historical_value) def change_trend_counter(self, values: list, release_date: dt.date | None, 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.current_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 def _extend_charts( self, record: Record, raw_record: Record, charts: Collection[str], dsp: str, release_date: dt.date | None ): 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.current_date # Fill "change" and "trend" if chart_base not in self.dsp_change_trends_props_map.get(dsp, ()): return change, trend, change_avg = self.change_trend_counter( values=values, release_date=release_date, days_gap=search_config.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=search_config.CHART_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: S3Path) -> Generator: for key in self.s3.get_keys(prefix): for row in self.csv.load(self.s3.read_object(key)): yield row yield None def _add_slugified_name(self, record: Record): if isinstance(record["name"], str): slugified_name = slugify(record["name"], separator="") else: self.log.warning( f"name field has type {type(record['name'])} which is incompatible. Falling back to None" ) slugified_name = None record["slugified_name"] = slugified_name