import datetime
import json
import pandas as pd
import numpy as np
import time
import tqdm
import matplotlib.pyplot as plt
import matplotlib.dates as mdates


def gather_data(engine, cutoff_date: str, filter_model_name: str = ""):
    """Master function to gather all data from the database."""

    print(f"Start gather data from {cutoff_date}")
    start_time = time.time()

    bots_action_query = f"""
    SELECT * FROM bots_generatedclipextra
    WHERE updated_at>='{cutoff_date}'
    """
    bots_action_df = pd.read_sql_query(bots_action_query, engine)
    print(f"Bots Action: {bots_action_df.shape[0]:,} rows")
    print(f" ---- Execution time: {time.time() - start_time:.2f} seconds")

    start_time = time.time()
    user_reaction_query = f"""
    SELECT * FROM bots_userreaction
    WHERE updated_at>='{cutoff_date}' AND play_count>0
    """
    reaction_df = pd.read_sql_query(user_reaction_query, engine)
    print(f"Reactions: {reaction_df.shape[0]:,} rows")
    print(f" ---- Execution time: {time.time() - start_time:.2f} seconds")

    start_time = time.time()
    generated_clip_query = f"""
    SELECT * FROM bots_generatedclip
    WHERE status='complete' AND created_at>='{cutoff_date}'
    """
    if filter_model_name:
        generated_clip_query += f" AND model_name::text='{filter_model_name}'"
    total_clip_df = pd.read_sql_query(generated_clip_query, engine)
    print(f"Total Clips: {total_clip_df.shape[0]:,} rows")
    print(f" ---- Execution time: {time.time() - start_time:.2f} seconds")

    start_time = time.time()
    playlist_clip_query = f"""
    SELECT * FROM bots_playlistclip
    WHERE updated_at>='{cutoff_date}'
    """
    playlist_clip_df = pd.read_sql_query(playlist_clip_query, engine)
    print(f"Playlist Clips: {playlist_clip_df.shape[0]:,} rows")
    print(f" ---- Execution time: {time.time() - start_time:.2f} seconds")

    start_time = time.time()
    auth_user_query = """
    SELECT *
    FROM auth_user_groups
    """
    auth_user_df = pd.read_sql_query(auth_user_query, engine)
    print(f"Authenticated Users: {auth_user_df.shape[0]:,} rows")
    print(f" ---- Execution time: {time.time() - start_time:.2f} seconds")

    start_time = time.time()
    discord_info_query = """
    SELECT * FROM bots_discordinfo
    WHERE subscription_status IN ('active', 'past_due')
    """
    discord_info_df = pd.read_sql_query(
        discord_info_query,
        engine,
    )
    # current active subscribers?
    print(f"Discord Info: {discord_info_df.shape[0]:,} rows")
    print(discord_info_df["subscription_status"].value_counts())
    print(f" ---- Execution time: {time.time() - start_time:.2f} seconds")

    return {
        "bots_action_df": bots_action_df,
        "reaction_df": reaction_df,
        "total_clip_df": total_clip_df,
        "playlist_clip_df": playlist_clip_df,
        "auth_user_df": auth_user_df,
        "discord_info_df": discord_info_df,
    }


def plot_clip_distribution(input_clip_df):
    # Create bins for date and hour
    input_clip_df["date"] = input_clip_df["created_at"].dt.date
    input_clip_df["hour"] = input_clip_df["created_at"].dt.hour

    # Group by date and hour, then count
    clip_counts = (
        input_clip_df.groupby(["date", "hour"]).size().reset_index(name="count")
    )

    # Create a datetime column for x-axis
    clip_counts["datetime"] = pd.to_datetime(clip_counts["date"]) + pd.to_timedelta(
        clip_counts["hour"], unit="h"
    )

    # Plot histogram
    plt.figure(figsize=(15, 8))
    plt.hist(
        clip_counts["datetime"],
        weights=clip_counts["count"],
        bins=len(clip_counts) * 2,
        edgecolor="black",
    )
    plt.xlabel("Date-Hour")
    plt.ylabel("Count")
    plt.title("Distribution of Clip Counts per Hour")
    # Format x-axis
    plt.gca().xaxis.set_major_formatter(mdates.DateFormatter("%Y-%m-%d %H:00"))
    plt.gca().xaxis.set_major_locator(mdates.DayLocator(interval=7))
    plt.xticks(rotation="vertical")  # Set x-axis tick labels to vertical

    plt.tight_layout()
    plt.show()

    # Print some statistics
    print(f"Total number of hours: {len(clip_counts)}")
    print(f"Average clips per hour: {clip_counts['count'].mean():.2f}")
    print(f"Max clips in an hour: {clip_counts['count'].max()}")
    print(f"Min clips in an hour: {clip_counts['count'].min()}")


def parse_parent_id(x):
    """Find out a clip's parent id."""
    if "history" not in x:
        return None
    out = x.get("history", [])
    if not isinstance(out, list) or len(out) == 0:
        return None
    # take the last one cause we continue off the children?
    out = out[-1]
    if isinstance(out, dict):
        # this is the continued info, which is a dict with id and continue_at
        return out["id"]
    else:
        return None


def parse_duration(x):
    """Find out a clip's duration."""
    if "duration" not in x:
        return None
    return x.get("duration")


def parse_source(x):
    """Find out a clip's source (web/ios)."""
    if "source" not in x:
        return None
    return x.get("source")


def parse_for_tag(x):
    if "tags" not in x:
        return ""
    out = x.get("tags", "")
    return out.lower() if out else ""


def parse_for_one_box(x):
    if "gpt_description_prompt" not in x:
        return False
    out = x.get("gpt_description_prompt", "")
    return out != ""


def parse_metadata_for_basics(x):
    """Parse the metadata for basics."""
    parent_id = parse_parent_id(x)
    duration = parse_duration(x)
    source = parse_source(x)
    return parent_id, duration, source


def get_concat_clip_ids(input_concat_df, input_clip_df, input_upload_clip_df):
    concat_clips_ids = {}
    history_error_counter = 0
    history_duplicate_error_counter = 0
    for _, row in tqdm.tqdm(input_concat_df.iterrows()):
        if concat_history_clips := row["metadata"].get("concat_history"):
            total_duration = row["metadata"].get("duration", 0)
            if total_duration == 0:
                print(row["metadata"], row["model_name"])
                continue
            start_s = 0
            for history_clip in concat_history_clips:
                if isinstance(history_clip, dict) and "id" in history_clip:
                    # the other key is `continue_at`
                    if history_clip["id"]:
                        # in case a clip ends up in multiple concats, we need to choose the optimal one
                        if history_clip["id"] in concat_clips_ids:
                            # keep the highest upvote clip
                            if (
                                row["reaction_upvote_count"]
                                < concat_clips_ids[history_clip["id"]]["concat_likes"]
                            ):
                                continue
                            # then keep the highest play count clip
                            if (
                                row["reaction_play_count"]
                                < concat_clips_ids[history_clip["id"]][
                                    "concat_play_counts"
                                ]
                            ):
                                continue
                            # multi-seed to concats
                            history_duplicate_error_counter += 1
                        concat_clips_ids[history_clip["id"]] = {
                            "total_start_s": start_s,
                            "total_clip_s": total_duration,
                            "concat_play_counts": row["reaction_play_count"],
                            "concat_in_playlist": row["is_in_playlist"],
                            "concat_likes": row["reaction_upvote_count"],
                            "concat_dislikes": row["reaction_dislike_count"],
                        }
                    else:
                        history_error_counter += 1
                    try:
                        # but we always update the start_s -- but keep the relative orders
                        if history_clip["continue_at"] is None:
                            # we need to go back and fetch the duration
                            if history_clip["id"] in input_clip_df["id"]:
                                start_s += input_clip_df[
                                    input_clip_df["id"] == history_clip["id"]
                                ]["duration"].iloc[0]
                            else:
                                # This is wrong but what can we do...
                                # this clip isn't kept in the query
                                start_s = 0
                        else:
                            start_s += history_clip["continue_at"]
                    except:
                        print(history_clip)
                        raise

    n_unique_uploads_in_concats = len(
        set(i for i in concat_clips_ids if i.startswith("m_"))
    )
    print(
        "total concat unique clips are:",
        len(concat_clips_ids),
        f"with error: {history_error_counter}, duplicate {history_duplicate_error_counter}",
        "\n",
        "uploads are in concats",
        n_unique_uploads_in_concats,
        "frac",
        f"{n_unique_uploads_in_concats / (input_upload_clip_df.shape[0] or 1):.3f}",
    )
    return concat_clips_ids


def validate_preference_data(input_clip_df):
    assert input_clip_df[input_clip_df["request_id"].isna()].shape[0] == 0
    check_df = input_clip_df.groupby("request_id")["id"].nunique()
    non_double_request_ids = check_df[check_df.values != 2].shape[0]
    assert non_double_request_ids == 0, non_double_request_ids
    print("Validation passed!")


def print_out_value_counts_nicely(input_df, input_column):
    source_counts = input_df[input_column].value_counts()
    total_count = source_counts.sum()

    # Print values and fractions
    for source, count in source_counts.items():
        fraction = count / total_count
        print(f"{source}: {count} ({fraction:.2%})")


def run_bot_detection(
    input_clip_df,
    input_reaction_df,
    inspection_date_cut="2024-05-10",
    min_generations_for_no_reaction=20,
    write_to_file=False,
):
    """Detect bots by looking at the play count and upvote count."""
    no_reaction_clip_df = input_clip_df[
        ~input_clip_df["id"].isin(input_reaction_df["clip_id"])
    ].copy()

    print("Shape of no_reaction_clip_df:")
    print(no_reaction_clip_df.shape)
    print("\nProportion of clips without reactions:")
    print(f"{no_reaction_clip_df.shape[0] / input_clip_df.shape[0]:.2%}")

    no_reaction_clip_df["no_reaction_count"] = no_reaction_clip_df.groupby("user_id")[
        "user_id"
    ].transform("count")
    no_reaction_clip_df["user_id"].nunique()
    bot_user_mask = (
        no_reaction_clip_df["no_reaction_count"] >= min_generations_for_no_reaction
    ) & (no_reaction_clip_df["created_at"] >= inspection_date_cut)
    sub_total_clip_df = input_clip_df[
        input_clip_df["user_id"].isin(
            no_reaction_clip_df[bot_user_mask]["user_id"].unique()
        )
    ].copy()
    ratio = no_reaction_clip_df[bot_user_mask].shape[0] / sub_total_clip_df.shape[0]
    print("Ratio of clips without reactions to total clips from the same users:")
    print(f"{ratio:.4f}")

    sub_total_clip_df["gen_count"] = sub_total_clip_df.groupby("user_id")[
        "user_id"
    ].transform("count")

    user_id_no_reaction_dict = no_reaction_clip_df.set_index("user_id")[
        "no_reaction_count"
    ].to_dict()
    user_id_total_dict = (
        sub_total_clip_df[~sub_total_clip_df["is_pro_user"]]
        .set_index("user_id")["gen_count"]
        .to_dict()
    )
    pro_user_id_total_dict = (
        sub_total_clip_df[sub_total_clip_df["is_pro_user"]]
        .set_index("user_id")["gen_count"]
        .to_dict()
    )

    user_ratio_dict = {}
    pro_user_ratio_dict = {}
    for user_id, total_gen in user_id_total_dict.items():
        user_ratio = user_id_no_reaction_dict.get(user_id, 0) / total_gen
        user_ratio_dict[user_id] = user_ratio
    for user_id, total_gen in pro_user_id_total_dict.items():
        user_ratio = user_id_no_reaction_dict.get(user_id, 0) / total_gen
        pro_user_ratio_dict[user_id] = user_ratio

    inspection_date_cut_timestamp = pd.to_datetime(inspection_date_cut)
    # Convert input_clip_df["created_at"] to datetime and ensure it's timezone-aware
    input_clip_min_date = pd.to_datetime(input_clip_df["created_at"].min()).tz_localize(
        inspection_date_cut_timestamp.tzinfo
    )

    cutoff_date = max(inspection_date_cut_timestamp, input_clip_min_date)
    plt.hist(
        user_ratio_dict.values(),
        bins=np.linspace(0, 1, 50),
        alpha=0.5,
        label="free user",
    )
    plt.hist(
        pro_user_ratio_dict.values(),
        bins=np.linspace(0, 1, 50),
        alpha=0.5,
        label="pro user",
    )
    plt.xlabel(
        f"fraction of generations (min {min_generations_for_no_reaction}) that have no actions"
    )
    plt.ylabel("number of users")
    plt.yscale("log")
    plt.legend()
    plt.title(f"Potential bots since {cutoff_date}")
    plt.show()
    super_bad_user_id = set()

    for user_id, user_ratio in user_ratio_dict.items():
        if user_ratio >= 0.99:
            super_bad_user_id.add(user_id)
    print("free", len(super_bad_user_id))
    super_bad_pro_user_id = set()
    for user_id, user_ratio in pro_user_ratio_dict.items():
        if user_ratio >= 0.99:
            super_bad_pro_user_id.add(user_id)
    print("pro", len(super_bad_pro_user_id))
    # Calculate the ratio of clips from super bad users
    bad_user_clip_ratio = (
        input_clip_df[input_clip_df["user_id"].isin(super_bad_user_id)].shape[0]
        / input_clip_df.shape[0]
    )

    # Print the result nicely
    print(f"Ratio of clips from super bad users: {bad_user_clip_ratio:.2%}")

    curr_date = datetime.datetime.today().strftime("%Y_%m_%d")
    if write_to_file:
        with open(f"/home/tony/Data/bots/bad_user_{curr_date}.json", "w") as fp:
            json.dump(list(super_bad_user_id), fp)
        with open(f"/home/tony/Data/bots/bad_pro_user_{curr_date}.json", "w") as fp:
            json.dump(list(super_bad_pro_user_id), fp)
    print("DONE")


def merge_concat_clips_with_reactions(input_concated_clip_df, input_reaction_df):
    """Merge the concat clips with the reactions."""
    concat_reaction_df = input_reaction_df[
        input_reaction_df["clip_id"].isin(set(input_concated_clip_df["id"]))
    ].copy()
    print(
        "concat reactions:",
        concat_reaction_df.shape[0],
        "unique concat clips:",
        concat_reaction_df["clip_id"].nunique(),
    )
    concat_reaction_df["upvote_count"] = concat_reaction_df["reaction_type"] == "L"
    concat_reaction_df["dislike_count"] = concat_reaction_df["reaction_type"] == "D"

    concat_total_reaction_df_sum = concat_reaction_df.groupby("clip_id")[
        ["play_count", "upvote_count", "dislike_count"]
    ].sum()
    concat_total_reaction_df_sum_df = concat_total_reaction_df_sum.reset_index().rename(
        columns={
            "clip_id": "id",
            "play_count": "reaction_play_count",
            "upvote_count": "reaction_upvote_count",
            "dislike_count": "reaction_dislike_count",
        }
    )
    concated_clips = input_concated_clip_df.merge(
        concat_total_reaction_df_sum_df, on="id", how="left"
    )

    print("total concats", concated_clips.shape)
    print(
        "check \n",
        concated_clips[
            ["reaction_play_count", "reaction_upvote_count", "reaction_dislike_count"]
        ].describe(),
    )
    return concated_clips


def plot_clip_basic_distributions(input_clip_df):
    fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(16, 12))

    ax1.hist(input_clip_df["dislike_count"], bins=np.linspace(0, 10, 30))
    ax1.set_yscale("log")
    ax1.set_xlabel("Number of dislike_count")
    ax1.set_ylabel("Number of clips")
    ax1.set_title("Dislike Count Distribution")

    ax2.hist(input_clip_df["upvote_count"], bins=np.linspace(0, 10, 30))
    ax2.set_yscale("log")
    ax2.set_xlabel("Number of like_count")
    ax2.set_ylabel("Number of clips")
    ax2.set_title("Upvote Count Distribution")

    ax3.hist(input_clip_df["is_public"], bins=np.linspace(0, 10, 30))
    ax3.set_yscale("log")
    ax3.set_xlabel("Number of is_public")
    ax3.set_ylabel("Number of clips")
    ax3.set_title("Public Clips Distribution")

    ax4.hist(input_clip_df["user_id"].value_counts(), bins=np.linspace(0, 1000, 100))
    ax4.set_yscale("log")
    ax4.set_xlabel("Number of preferences clips")
    ax4.set_ylabel("Number of users")
    ax4.set_title("User Preferences Distribution")

    plt.tight_layout()
    plt.show()
