import asyncio import re from functools import partial from time import time from typing import Callable from pandas import DataFrame from ... import logger from ...constants import SnowFlakeColumns as SFCols from ...typings import AuditTypeShort from ...utils.pandas_helpers import to_set from .io.snowflake import LookTask, SnowflakeRequest from .looks import at, mv, sr, sr_ugc logger = logger.new_logger(__name__) async def fetch( label_id: int, *, include_sr: bool = True, include_mv: bool = False, include_at: bool = False, ) -> dict[AuditTypeShort, DataFrame]: """ Fetches audit data from Snowflake based on the underlying queries of the audit looks (SR and SR_UGC). Optionally, MV and AT data can also be fetched. Args: label_id (int): Label ID. include_sr (bool, optional): Whether to fetch SR data. Defaults to True. include_mv (bool, optional): Whether to fetch MV data. Defaults to False. include_at (bool, optional): Whether to fetch AT data. Defaults to False. """ start_time = time() if not any([include_sr, include_mv, include_at]): raise ValueError( "At least one of 'include_sr', 'include_mv', or 'include_at' must be True." ) tasks = [] if include_sr: tasks.extend( [ fetch_sr, fetch_sr_ugc1, fetch_sr_ugc2, fetch_sr_ugc3, ] ) if include_mv: tasks.append(fetch_mv) if include_at: tasks.append(fetch_at) results = await asyncio.gather(*map(lambda x: x(label_id), tasks)) logger.info(f"Data fetched in {time() - start_time:.2f}s.") callable_prefix: str = "fetch_" try: result = dict( zip((t.__name__.split(callable_prefix)[1] for t in tasks), results) ) except IndexError as ex: raise ValueError( f"Invalid fetch function name. Must start with '{callable_prefix}'." ) from ex return result async def _fetch_sr_ugc(label_id: int, look_name: str, query: Callable) -> DataFrame: results = await _snowflake_request_factory().execute( [LookTask(look_name=look_name, query=query(label_id))] ) return results[0] async def fetch_sr_ugc1(label_id: int) -> DataFrame: return await _fetch_sr_ugc(label_id, look_name="SR_UGC1", query=sr_ugc.sr_ugc1) async def fetch_sr_ugc2(label_id: int) -> DataFrame: return await _fetch_sr_ugc(label_id, look_name="SR_UGC2", query=sr_ugc.sr_ugc2) async def fetch_sr_ugc3(label_id: int) -> DataFrame: return await _fetch_sr_ugc(label_id, look_name="SR_UGC3", query=sr_ugc.sr_ugc3) async def fetch_sr(label_id: int) -> DataFrame: def sr3(df) -> DataFrame: """Preprocess SR3.""" _col_preprocessor(df) df.rename( columns={SFCols.TERRITORY_LIST: SFCols.REGISTRY_TERRITORIES}, inplace=True, ) return df def sr6(df) -> DataFrame: """Preprocess SR6.""" _col_preprocessor(df) df = ( df.groupby(SFCols.UPC)[SFCols.ABBRIVATION] .apply(lambda x: list(x.unique())) .reset_index() ) return df def sr8(df) -> DataFrame: """Preprocess SR8.""" _col_preprocessor(df) df.rename( columns={SFCols.TERRITORY_LIST: SFCols.CONFLICTING_TERRITORIES}, inplace=True, ) return df look_tasks: list[LookTask] = [ LookTask( look_name="SR_UGC2", query=sr_ugc.sr_ugc2(label_id), preprocessor=_col_preprocessor, ), LookTask(look_name="SR1", query=sr.sr1(label_id)), LookTask(look_name="SR2", query=sr.sr2(label_id)), LookTask(look_name="SR3", query=sr.sr3(label_id), preprocessor=sr3), LookTask(look_name="SR4", query=sr.sr4(label_id)), LookTask(look_name="SR5", query=sr.sr5(label_id)), LookTask(look_name="SR6", query=sr.sr6(label_id), preprocessor=sr6), LookTask(look_name="SR7", query=sr.sr7(label_id)), LookTask(look_name="SR8", query=sr.sr8(label_id), preprocessor=sr8), ] worker = _snowflake_request_factory() results = await worker.execute(look_tasks) ( df_sr_ugc2, df_sr1, df_sr2, df_sr3, df_sr4, df_sr5, df_sr6, df_sr7, sr8, ) = results # noqa # Table merging operations for df in [df_sr3, df_sr5, df_sr7, sr8]: df_sr1 = df_sr1.merge(df, how="left", on=SFCols.ISRC) df_sr1 = df_sr1.merge(df_sr6, how="left", on=SFCols.UPC) df_sr1[SFCols.IS_CARVED_OUT] = df_sr1[SFCols.UPC].isin(df_sr2[SFCols.UPC]) df_sr1[SFCols.IS_LOCKED] = df_sr1[SFCols.ISRC].isin(df_sr4[SFCols.ISRC]) df_sr1 = df_sr1.merge(df_sr4, how="left", on=SFCols.ISRC) df_sr1 = df_sr1.merge(df_sr_ugc2, how="left", on=SFCols.ASSET_ID) return df_sr1 async def fetch_mv(label_id: int) -> DataFrame: look_tasks: list[LookTask] = [ LookTask(look_name="MV1", query=mv.mv1(label_id)), LookTask(look_name="MV2", query=mv.mv2(label_id)), ] df_mv1, df_mv2 = await _snowflake_request_factory().execute(look_tasks) df_mv = df_mv1.merge(df_mv2, how="left", on=SFCols.ASSET_ID) # Cast any empty ISRCs to None. There seems to be bad data in Snowflake # where ISRCs are empty strings instead of NULL. This is already fixed # in the ingestion query using NULLIF function, but just in case, enforce # it here as well. df_mv[SFCols.ISRC] = df_mv[SFCols.ISRC].apply(lambda x: x or None) return df_mv def _preprocess_at3(df) -> DataFrame: """Preprocess AT3.""" _col_preprocessor(df) # Ensure there's no duplicate UPCs (they must be grouped at the source). Raise # an exception if there are duplicates. if df[SFCols.UPC].duplicated().any(): raise ValueError("Duplicate UPCs found in AT3.") def unique_carved_out_countries(row): """Merge all the territories from VTR, RTR, and STR columns into a single column and return a unique set of territories for which each UPC is carved out. VTR: Vendor Territory Restrictions RTR: Release Territory Restrictions STR: Subaccount Territory Restrictions """ def handler(x): return set((x or "").split(",")) vtr_countries = handler(row[SFCols.VTR_COUNTRIES]) rtr_countries = handler(row[SFCols.RTR_COUNTRIES]) str_countries = handler(row[SFCols.STR_COUNTRIES]) merged_countries = vtr_countries | rtr_countries | str_countries abbrivation = {v.strip() for v in merged_countries if v.strip()} return abbrivation # Create a new column with the unique carved out countries for each UPC, # merging the VTR, RTR, and STR columns. To prevent this crashing if the # input df is empty, check if it's empty first or return an empty series. df[SFCols.ABBRIVATION] = ( df.apply(unique_carved_out_countries, axis=1) if not df.empty else [] ) return df async def fetch_at(label_id: int) -> DataFrame: look_tasks: list[LookTask] = [ LookTask(look_name="AT1", query=at.at1(label_id)), LookTask(look_name="AT2", query=at.at2(label_id)), LookTask(look_name="AT3", query=at.at3(label_id), preprocessor=_preprocess_at3), LookTask(look_name="AT4", query=at.at4(label_id)), LookTask(look_name="AT5", query=at.at5(label_id)), LookTask(look_name="AT6", query=at.at6(label_id)), ] worker = _snowflake_request_factory() results = await worker.execute(look_tasks) df_at1, df_at2, df_at3, df_at4, df_at5, df_at6 = results # noqa # Table merging operations df_at1[SFCols.IS_CARVED_OUT] = df_at1[SFCols.UPC].isin(df_at2[SFCols.UPC]) df_at1[SFCols.IS_DELIVERED] = df_at1[SFCols.DELIVERED_STATUS].notna() & ( df_at1[SFCols.DELIVERED_STATUS].str.lower() == "delivered" ) df_at1 = df_at1.merge(df_at3, how="left", on=SFCols.UPC) df_at1 = df_at1.merge(df_at4, how="left", on=[SFCols.UPC, SFCols.ISRC]) for df in [df_at5, df_at6]: df_at1 = df_at1.merge(df, how="left", on=SFCols.ASSET_ID) return df_at1 def _col_preprocessor(df) -> DataFrame: """Remove the table name from the column names and force lowercase for more readable and consistent column names. """ df.columns = df.columns.str.split(".").str[-1].str.lower() return df # Set a default preprocessor for all looks LookTask = partial(LookTask, preprocessor=_col_preprocessor) class Formatters: """Formatters for SnowflakeRequest.""" @staticmethod def formatter_1(x) -> set | None: """Convert to set using comma or semicolon as separator. Empty strings are converted to None. """ if not x: return None return to_set(x, separator_regex=re.compile(r"\s*[,;]\s*")) @staticmethod def formatter_2(x) -> set | None: """Convert to set using whitespace as separator. Empty strings are converted to None. """ if not x: return None return to_set(x.strip() if x else x, separator_regex=re.compile(r"\s+")) def _snowflake_request_factory() -> SnowflakeRequest: formatter_1 = Formatters.formatter_1 formatter_2 = Formatters.formatter_2 formatters = { SFCols.REGISTRY_TERRITORIES: formatter_1, SFCols.CONFLICTING_TERRITORIES: formatter_1, SFCols.LIST_CONFLICTING_TERRITORIES: formatter_1, SFCols.ACTIVE_REFERENCE_IDS: formatter_2, SFCols.INACTIVE_REFERENCE_IDS: formatter_2, SFCols.VTR_COUNTRIES: formatter_1, SFCols.RTR_COUNTRIES: formatter_1, SFCols.STR_COUNTRIES: formatter_1, SFCols.OWNERSHIP: formatter_1, } return SnowflakeRequest(formatters=formatters)