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)
