import json from datetime import datetime, timedelta from typing import Any, Dict, List, Tuple import pandas as pd import snowflake.snowpark as snowpark TODAY_COUNT_COLUMN_NAME = "TODAY_COUNT" YESTERDAY_COUNT_COLUMN_NAME = "YESTERDAY_COUNT" HOURLY_OCCURENCE_COLUMN_NAME = "HOURLY_OCCURENCE" TOTAL_COUNT_COLUMN_NAME = "TOTAL_COUNT" def post_to_slack( session: snowpark.Session, validation_results_so_far: List[Tuple[str, bool]], table_name: str, p_date: str, p_hour: int, ): failed_validations = [ validation_name for validation_name, validation_result in validation_results_so_far if validation_result is True ] if failed_validations: alert_message = ( f"{table_name} on {p_date} hour {p_hour}: " + f"Failed validations: {', '.join(failed_validations)}" ) escaped_message = alert_message.replace("'", "''") session.sql( f"SELECT SUNO_PROD.PROD.post_to_slack('#data-alerts', '{escaped_message}')" ).collect() def get_sql_result(session: snowpark.Session, schema: str, default_columns: List[str]): df = session.sql(schema) local_df = df.collect() if len(local_df) == 0: return pd.DataFrame(columns=default_columns) return pd.DataFrame(local_df) def get_data_counts( session: snowpark.Session, p_date: str, p_hour: int, table_name: str, selected_column: str | None = None, group_by_columns: str | Tuple[str, ...] | None = None, aggregation_type: str = "distinct_count", ) -> pd.DataFrame: yesterday_date = (datetime.strptime(p_date, "%Y-%m-%d") - timedelta(days=1)).strftime("%Y-%m-%d") # Handle group by columns (single string or tuple of strings) if group_by_columns is None: group_by_str = "" group_by_clause = "" group_by_column_list = [] elif isinstance(group_by_columns, str): group_by_str = f"{group_by_columns}, " group_by_clause = f"group by {group_by_columns} order by {group_by_columns}" group_by_column_list = [group_by_columns] else: # tuple of columns group_by_str = ", ".join(group_by_columns) + ", " group_by_clause = ( f"group by {', '.join(group_by_columns)} order by {', '.join(group_by_columns)}" ) group_by_column_list = list(group_by_columns) # Generate aggregation expressions based on type if selected_column is None: # Total row count today_expr = f"count_if(p_date = '{p_date}') as {TODAY_COUNT_COLUMN_NAME}" yesterday_expr = f"count_if(p_date = '{yesterday_date}') as {YESTERDAY_COUNT_COLUMN_NAME}" elif aggregation_type == "distinct_count": today_expr = f"count(distinct case when p_date = '{p_date}' then {selected_column} end) as {TODAY_COUNT_COLUMN_NAME}" yesterday_expr = f"count(distinct case when p_date = '{yesterday_date}' then {selected_column} end) as {YESTERDAY_COUNT_COLUMN_NAME}" elif aggregation_type == "sum": today_expr = f"sum(case when p_date = '{p_date}' then {selected_column} else 0 end) as {TODAY_COUNT_COLUMN_NAME}" yesterday_expr = f"sum(case when p_date = '{yesterday_date}' then {selected_column} else 0 end) as {YESTERDAY_COUNT_COLUMN_NAME}" else: raise ValueError(f"Unsupported aggregation_type: {aggregation_type}") query = f""" select {group_by_str} {today_expr}, {yesterday_expr} from {table_name} where p_date in ('{p_date}', '{yesterday_date}') and p_hour = {p_hour} {group_by_clause} """ default_columns = group_by_column_list + [TODAY_COUNT_COLUMN_NAME, YESTERDAY_COUNT_COLUMN_NAME] return get_sql_result(session, query, default_columns=default_columns) def fraction_difference(today_count, yesterday_count): # Handle None values and edge cases if today_count is None or yesterday_count is None or yesterday_count == 0: return 0.0 # No meaningful comparison possible else: return abs(today_count - yesterday_count) / yesterday_count def null_check(p_date: str, p_hour: int, column_name: str, table_name: str) -> str: return f""" select count_if({column_name} is null) as {HOURLY_OCCURENCE_COLUMN_NAME}, count(*) as {TOTAL_COUNT_COLUMN_NAME} from {table_name} where p_date = '{p_date}' and p_hour = {p_hour}; """ def length_check(p_date: str, p_hour: int, column_name: str, length: int, table_name: str) -> str: return f""" select count_if(length({column_name}) != {length}) as {HOURLY_OCCURENCE_COLUMN_NAME}, count(*) as {TOTAL_COUNT_COLUMN_NAME} from {table_name} where p_date = '{p_date}' and p_hour = {p_hour}; """ def negative_check(p_date: str, p_hour: int, column_name: str, table_name: str) -> str: return f""" select count_if({column_name} < 0) as {HOURLY_OCCURENCE_COLUMN_NAME}, count(*) as {TOTAL_COUNT_COLUMN_NAME} from {table_name} where p_date = '{p_date}' and p_hour = {p_hour}; """ def zero_check(p_date: str, p_hour: int, column_name: str, table_name: str) -> str: return f""" select count_if({column_name} = 0) as {HOURLY_OCCURENCE_COLUMN_NAME}, count(*) as {TOTAL_COUNT_COLUMN_NAME} from {table_name} where p_date = '{p_date}' and p_hour = {p_hour}; """ def less_than_check( p_date: str, p_hour: int, column_name_1: str, column_name_2: str, table_name: str ) -> str: return f""" select count_if({column_name_1} < {column_name_2}) as {HOURLY_OCCURENCE_COLUMN_NAME}, count(*) as {TOTAL_COUNT_COLUMN_NAME} from {table_name} where p_date = '{p_date}' and p_hour = {p_hour}; """ def less_than_or_equal_to_check( p_date: str, p_hour: int, column_name_1: str, column_name_2: str, table_name: str ) -> str: return f""" select count_if({column_name_1} <= {column_name_2}) as {HOURLY_OCCURENCE_COLUMN_NAME}, count(*) as {TOTAL_COUNT_COLUMN_NAME} from {table_name} where p_date = '{p_date}' and p_hour = {p_hour}; """ def greater_than_check( p_date: str, p_hour: int, column_name_1: str, column_name_2: str, table_name: str ) -> str: return f""" select count_if({column_name_1} > {column_name_2}) as {HOURLY_OCCURENCE_COLUMN_NAME}, count(*) as {TOTAL_COUNT_COLUMN_NAME} from {table_name} where p_date = '{p_date}' and p_hour = {p_hour}; """ def not_equal_to_check( p_date: str, p_hour: int, column_name_1: str, column_name_2: str, table_name: str ) -> str: return f""" select count_if({column_name_1} != {column_name_2}) as {HOURLY_OCCURENCE_COLUMN_NAME}, count(*) as {TOTAL_COUNT_COLUMN_NAME} from {table_name} where p_date = '{p_date}' and p_hour = {p_hour}; """ def top_count_events_check( p_date: str, p_hour: int, column_name: str, top_count_threshold: int, table_name: str ) -> str: return f""" WITH user_counts AS ( SELECT {column_name}, COUNT(*) as cnt FROM {table_name} WHERE p_date = '{p_date}' and p_hour = {p_hour} GROUP BY {column_name} ) SELECT (SELECT COUNT(*) FROM user_counts) as {TOTAL_COUNT_COLUMN_NAME}, COUNT(*) as {HOURLY_OCCURENCE_COLUMN_NAME} FROM user_counts WHERE cnt > {top_count_threshold}; """ def distinct_column_1_group_by_column_2_sum_check( p_date: str, p_hour: int, column_name_1: str, column_name_2: str, sum_threshold: int, table_name: str, ) -> str: return f""" WITH user_sums AS ( SELECT {column_name_2}, SUM({column_name_1}) as total_sum FROM {table_name} WHERE p_date = '{p_date}' AND p_hour = {p_hour} GROUP BY {column_name_2} ) SELECT COUNT(*) as {TOTAL_COUNT_COLUMN_NAME}, COUNT(CASE WHEN total_sum > {sum_threshold} THEN 1 END) as {HOURLY_OCCURENCE_COLUMN_NAME} FROM user_sums; """ def distinct_column_1_group_by_column_2_count_check( p_date: str, p_hour: int, column_name_1: str, column_name_2: str, distinct_count_threshold: int, table_name: str, ) -> str: return f""" WITH user_distinct_counts AS ( SELECT {column_name_2}, COUNT(DISTINCT {column_name_1}) as total_distinct_count FROM {table_name} WHERE p_date = '{p_date}' AND p_hour = {p_hour} GROUP BY {column_name_2} ) SELECT COUNT(*) as {TOTAL_COUNT_COLUMN_NAME}, COUNT(CASE WHEN total_distinct_count > {distinct_count_threshold} THEN 1 END) as {HOURLY_OCCURENCE_COLUMN_NAME} FROM user_distinct_counts; """ def duplicate_check(p_date: str, p_hour: int, column_names: List[str], table_name: str) -> str: columns_str = ", ".join(column_names) join_conditions = " AND ".join([f"t.{col} = d.{col}" for col in column_names]) return f""" WITH duplicate_groups AS ( SELECT {columns_str}, COUNT(*) as group_count FROM {table_name} WHERE p_date = '{p_date}' AND p_hour = {p_hour} GROUP BY {columns_str} HAVING COUNT(*) > 1 ), duplicate_rows AS ( SELECT t.* FROM {table_name} t INNER JOIN duplicate_groups d ON {join_conditions} WHERE t.p_date = '{p_date}' AND t.p_hour = {p_hour} ) SELECT COUNT(*) as {HOURLY_OCCURENCE_COLUMN_NAME}, (SELECT COUNT(*) FROM {table_name} WHERE p_date = '{p_date}' AND p_hour = {p_hour}) as {TOTAL_COUNT_COLUMN_NAME} FROM duplicate_rows; """ def false_check(p_date: str, p_hour: int, column_name: str, table_name: str) -> str: return f""" select count_if({column_name} = false) as {HOURLY_OCCURENCE_COLUMN_NAME}, count(*) as {TOTAL_COUNT_COLUMN_NAME} from {table_name} where p_date = '{p_date}' and p_hour = {p_hour}; """ def get_hourly_occurence( session: snowpark.Session, p_date: str, p_hour: int, column_names: List[str], check: Dict[str, Any], table_name: str, ) -> pd.DataFrame: if check["name"] == "length": query = length_check(p_date, p_hour, column_names[0], check["length"], table_name) elif check["name"] == "null": query = null_check(p_date, p_hour, column_names[0], table_name) elif check["name"] == "negative": query = negative_check(p_date, p_hour, column_names[0], table_name) elif check["name"] == "zero": query = zero_check(p_date, p_hour, column_names[0], table_name) elif check["name"] == "less_than": query = less_than_check(p_date, p_hour, column_names[0], column_names[1], table_name) elif check["name"] == "less_than_or_equal_to": query = less_than_or_equal_to_check(p_date, p_hour, column_names[0], column_names[1], table_name) elif check["name"] == "greater_than": query = greater_than_check(p_date, p_hour, column_names[0], column_names[1], table_name) elif check["name"] == "not_equal_to": query = not_equal_to_check(p_date, p_hour, column_names[0], column_names[1], table_name) elif check["name"] == "top_count_events": query = top_count_events_check( p_date, p_hour, column_names[0], check["top_count_threshold"], table_name ) elif check["name"] == "distinct_column_1_group_by_column_2_sum": query = distinct_column_1_group_by_column_2_sum_check( p_date, p_hour, column_names[0], column_names[1], check["sum_threshold"], table_name ) elif check["name"] == "distinct_column_1_group_by_column_2_count": query = distinct_column_1_group_by_column_2_count_check( p_date, p_hour, column_names[0], column_names[1], check["distinct_count_threshold"], table_name, ) elif check["name"] == "duplicate": query = duplicate_check(p_date, p_hour, column_names, table_name) elif check["name"] == "false": query = false_check(p_date, p_hour, column_names[0], table_name) else: raise ValueError(f"Invalid check name: {check['name']}") return get_sql_result( session, query, default_columns=[HOURLY_OCCURENCE_COLUMN_NAME, TOTAL_COUNT_COLUMN_NAME] ) def perform_data_volume_validate_diff( monitor_dataframe: pd.DataFrame, validation_results_so_far: List[Tuple[str, bool]], validation_name: str, p_date: str, p_hour: int, today_count: int, yesterday_count: int, diff_threshold: float, min_count_threshold: int, table_name: str, task_name: str, ) -> pd.DataFrame: diff = fraction_difference(today_count, yesterday_count) validation_result = diff >= diff_threshold and (yesterday_count - today_count > min_count_threshold) validation_results_so_far.append((validation_name, validation_result)) new_row = pd.DataFrame( { "table_name": table_name, "task_name": task_name, "p_date": p_date, "p_hour": p_hour, "validation_name": validation_name, "validation_value": json.dumps( { "today_count": today_count, "yesterday_count": yesterday_count, "fraction_difference": diff, } ), "validation_result": validation_result, }, index=[0], ) return pd.concat([monitor_dataframe, new_row], ignore_index=True) def perform_data_quality_check_validate_diff( monitor_dataframe: pd.DataFrame, validation_results_so_far: List[Tuple[str, bool]], validation_name: str, p_date: str, p_hour: int, hourly_occurence_value: int, total_count: int, threshold: float, table_name: str, task_name: str, ) -> pd.DataFrame: # Handle None values and edge cases if hourly_occurence_value is None or total_count is None or total_count == 0: # No data for this hour or no total count available - consider it as valid validation_result = True error_rate = 0.0 else: # Normal case - calculate the percentage validation_result = (hourly_occurence_value / total_count) > threshold error_rate = hourly_occurence_value / total_count validation_results_so_far.append((validation_name, validation_result)) new_row = pd.DataFrame( { "table_name": table_name, "task_name": task_name, "p_date": p_date, "p_hour": p_hour, "validation_name": validation_name, "validation_value": json.dumps( { "hourly_occurence_value": hourly_occurence_value, "total_count": total_count, "error_rate": error_rate, } ), "validation_result": validation_result, }, index=[0], ) return pd.concat([monitor_dataframe, new_row], ignore_index=True) def validate_data_volume_monitoring( session: snowpark.Session, p_date: str, p_hour: int, validation_results_so_far: List[Tuple[str, bool]], table_name: str, task_name: str, monitor_table_columns: List[str], diff_threshold: float, min_count_threshold: int, column_validation_dict: Dict[str, str], group_by_columns: List[str | Tuple[str, ...]], ): monitor_dataframe = pd.DataFrame(columns=monitor_table_columns) for validation_name, column_name in column_validation_dict.items(): # Determine aggregation type based on validation name aggregation_type = "sum" if validation_name.startswith("sum_") else "distinct_count" data_counts = get_data_counts( session, p_date, p_hour, table_name=table_name, selected_column=column_name, aggregation_type=aggregation_type, ) for today_value, yesterday_value in zip( data_counts[TODAY_COUNT_COLUMN_NAME].tolist(), data_counts[YESTERDAY_COUNT_COLUMN_NAME].tolist(), ): monitor_dataframe = perform_data_volume_validate_diff( monitor_dataframe, validation_results_so_far, f"{validation_name}_fraction_diff_above_{diff_threshold}", p_date, p_hour, today_value, yesterday_value, diff_threshold, min_count_threshold, table_name, task_name, ) if group_by_columns: for group_by_item in group_by_columns: for validation_name_prefix, column_name in column_validation_dict.items(): # Determine aggregation type based on validation name aggregation_type = ( "sum" if validation_name_prefix.startswith("sum_") else "distinct_count" ) data_counts = get_data_counts( session, p_date, p_hour, table_name=table_name, selected_column=column_name, group_by_columns=group_by_item, aggregation_type=aggregation_type, ) # Handle both single column and multiple column grouping if isinstance(group_by_item, str): # Single column grouping for today_value, yesterday_value, group_by_value in zip( data_counts[TODAY_COUNT_COLUMN_NAME].tolist(), data_counts[YESTERDAY_COUNT_COLUMN_NAME].tolist(), data_counts[group_by_item].tolist(), ): monitor_dataframe = perform_data_volume_validate_diff( monitor_dataframe, validation_results_so_far, f"{validation_name_prefix}_with_{group_by_item}_as_{group_by_value}_fraction_diff_above_{diff_threshold}", p_date, p_hour, today_value, yesterday_value, diff_threshold, min_count_threshold, table_name, task_name, ) else: # Multiple column grouping (tuple) group_by_values_lists = [data_counts[col].tolist() for col in group_by_item] for row_idx in range(len(data_counts)): today_value = data_counts[TODAY_COUNT_COLUMN_NAME].tolist()[row_idx] yesterday_value = data_counts[YESTERDAY_COUNT_COLUMN_NAME].tolist()[row_idx] group_by_values = [values_list[row_idx] for values_list in group_by_values_lists] group_by_description = "_".join( [f"{col}_{val}" for col, val in zip(group_by_item, group_by_values)] ) monitor_dataframe = perform_data_volume_validate_diff( monitor_dataframe, validation_results_so_far, f"{validation_name_prefix}_with_{group_by_description}_fraction_diff_above_{diff_threshold}", p_date, p_hour, today_value, yesterday_value, diff_threshold, min_count_threshold, table_name, task_name, ) session.create_dataframe(monitor_dataframe).write.save_as_table( "base_event_monitoring_results", mode="append" ) def validate_data_quality_check( session: snowpark.Session, p_date: str, p_hour: int, validation_results_so_far: List[Tuple[str, bool]], table_name: str, task_name: str, monitor_table_columns: List[str], column_validation_dict: Dict[str, Dict[str, Any]], ): monitor_dataframe = pd.DataFrame(columns=monitor_table_columns) for validation_name, dictionary in column_validation_dict.items(): for check in dictionary["checks"]: hourly_occurence = get_hourly_occurence( session, p_date, p_hour, dictionary["column_names"], check, table_name, ) for hourly_occurence_value, total_count in zip( hourly_occurence[HOURLY_OCCURENCE_COLUMN_NAME].tolist(), hourly_occurence[TOTAL_COUNT_COLUMN_NAME].tolist(), ): monitor_dataframe = perform_data_quality_check_validate_diff( monitor_dataframe, validation_results_so_far, f"{validation_name}_{check['name']}_error_rate_above_{check['threshold']}", p_date, p_hour, hourly_occurence_value, total_count, check["threshold"], table_name, task_name, ) session.create_dataframe(monitor_dataframe).write.save_as_table( "base_event_monitoring_results", mode="append" ) def validate_data( session: snowpark.Session, p_date: str, p_hour: int, table_name: str, task_name: str, monitor_table_columns: List[str], data_volume_checks: Dict[str, str], quality_checks: Dict[str, Dict[str, Any]], group_by_columns: List[str | Tuple[str, ...]], diff_threshold: float, min_count_threshold: int, ): validation_results_so_far = [] validate_data_volume_monitoring( session, p_date, p_hour, validation_results_so_far, table_name, task_name, monitor_table_columns, diff_threshold, min_count_threshold, data_volume_checks, group_by_columns, ) if quality_checks: validate_data_quality_check( session, p_date, p_hour, validation_results_so_far, table_name, task_name, monitor_table_columns, quality_checks, ) post_to_slack(session, validation_results_so_far, table_name, p_date, p_hour)