"""
Feature Engineering Module
=========================
Contains all feature creation functions for user clustering.
"""

import pandas as pd
import numpy as np
import json
import logging
from typing import List, Optional
from utils import parse_metadata_field

logger = logging.getLogger(__name__)


def create_temporal_features(total_clip_df: pd.DataFrame) -> pd.DataFrame:
    """Create temporal-based features for each user."""
    logger.info("Creating temporal features")

    # Ensure created_at is datetime
    total_clip_df["created_at"] = pd.to_datetime(total_clip_df["created_at"])

    # Group by user
    user_groups = total_clip_df.groupby("user_id")

    # Initialize features dataframe with user_id as a column (not index)
    features = pd.DataFrame()
    features["user_id"] = user_groups.size().index

    # Basic counts
    features["total_clips_created"] = user_groups.size().values

    # Date features
    features["first_creation_date"] = user_groups["created_at"].min().values
    features["last_creation_date"] = user_groups["created_at"].max().values
    features["account_age_days"] = (
        features["last_creation_date"] - features["first_creation_date"]
    ).dt.days

    # Activity features - vectorized
    # Create a date column for grouping
    total_clip_df["creation_date"] = total_clip_df["created_at"].dt.date
    days_active = total_clip_df.groupby("user_id")["creation_date"].nunique()
    features["days_active"] = days_active.reindex(features["user_id"]).fillna(0).values

    # Creation frequency
    features["creation_frequency"] = features["total_clips_created"] / (
        features["account_age_days"] + 1
    )

    # Creation consistency
    daily_counts = total_clip_df.groupby(["user_id", "creation_date"]).size()
    consistency = daily_counts.groupby(level=0).std().fillna(0)
    features["creation_consistency"] = (
        consistency.reindex(features["user_id"]).fillna(0).values
    )

    # Peak hour - vectorized
    if "hour" in total_clip_df.columns:
        # Use value_counts and take the index of the max
        hour_counts = total_clip_df.groupby("user_id")["hour"].value_counts()
        peak_hours = hour_counts.groupby(level=0).idxmax().str[-1].astype(int)
        features["peak_hour_creation"] = (
            peak_hours.reindex(features["user_id"]).fillna(0).values
        )
    else:
        total_clip_df["hour"] = total_clip_df["created_at"].dt.hour
        hour_counts = total_clip_df.groupby("user_id")["hour"].value_counts()
        peak_hours = hour_counts.groupby(level=0).idxmax().str[-1].astype(int)
        features["peak_hour_creation"] = (
            peak_hours.reindex(features["user_id"]).fillna(0).values
        )

    # Weekend vs weekday
    total_clip_df["is_weekend"] = total_clip_df["created_at"].dt.dayofweek.isin([5, 6])
    weekend_ratio = user_groups["is_weekend"].mean()
    features["weekend_vs_weekday_ratio"] = weekend_ratio.values

    # Creation burst (max clips in a day)
    max_daily = daily_counts.groupby(level=0).max()
    features["creation_burst_score"] = (
        max_daily.reindex(features["user_id"]).fillna(0).values
    )

    # Time between creations - vectorized approach
    # Sort by user and time, then calculate time differences
    sorted_df = total_clip_df.sort_values(["user_id", "created_at"])
    sorted_df["time_diff"] = sorted_df.groupby("user_id")["created_at"].diff()
    sorted_df["time_diff_seconds"] = sorted_df["time_diff"].dt.total_seconds()

    # Calculate std of time differences per user
    time_std = sorted_df.groupby("user_id")["time_diff_seconds"].std().fillna(0)
    features["time_between_creations_std"] = (
        time_std.reindex(features["user_id"]).fillna(0).values
    )

    # Clean up temporary columns
    total_clip_df.drop(
        ["creation_date", "is_weekend", "hour"], axis=1, errors="ignore", inplace=True
    )

    return features


def create_bot_detection_features(
    total_clip_df: pd.DataFrame, reaction_df: pd.DataFrame
) -> pd.DataFrame:
    """Create features to help detect bot behavior."""
    logger.info("Creating bot detection features")

    features = pd.DataFrame()
    unique_users = total_clip_df["user_id"].unique()
    features["user_id"] = unique_users

    # Count clips per user
    clips_per_user = total_clip_df.groupby("user_id").size()
    features["clips_created"] = (
        clips_per_user.reindex(features["user_id"]).fillna(0).values
    )

    # Count reactions given by each user
    reactions_given = reaction_df.groupby("user_id")["clip_id"].count()
    features["reactions_given"] = (
        reactions_given.reindex(features["user_id"]).fillna(0).values
    )

    # Has any reactions
    features["has_any_reactions"] = features["reactions_given"] > 0

    # Consumption ratio
    features["consumption_ratio"] = features["reactions_given"] / (
        features["clips_created"] + 1
    )

    # Creation to reaction ratio
    features["creation_to_reaction_ratio"] = features["clips_created"] / (
        features["reactions_given"] + 1
    )

    # Time between creations - vectorized approach
    # Sort by user and time
    sorted_df = total_clip_df.sort_values(["user_id", "created_at"])
    sorted_df["time_diff"] = sorted_df.groupby("user_id")["created_at"].diff()
    sorted_df["time_diff_seconds"] = sorted_df["time_diff"].dt.total_seconds()

    # Calculate avg and min time between creations per user
    time_stats = sorted_df.groupby("user_id")["time_diff_seconds"].agg(["mean", "min"])

    features["avg_time_between_creations"] = (
        time_stats["mean"].reindex(features["user_id"]).fillna(3600).values
    )  # Default 1 hour
    features["min_time_between_creations"] = (
        time_stats["min"].reindex(features["user_id"]).fillna(3600).values
    )

    # Drop temporary column
    features = features.drop("clips_created", axis=1)

    return features


def create_content_quality_features(total_clip_df: pd.DataFrame) -> pd.DataFrame:
    """Create content quality related features."""
    logger.info("Creating content quality features")

    user_groups = total_clip_df.groupby("user_id")

    features = pd.DataFrame()
    features["user_id"] = user_groups.size().index

    # Duration features
    if "duration" in total_clip_df.columns:
        features["avg_clip_duration"] = user_groups["duration"].mean().values
    else:
        features["avg_clip_duration"] = 30  # Default duration

    # Public/private ratio
    if "is_public" in total_clip_df.columns:
        features["public_clip_ratio"] = user_groups["is_public"].mean().values
    else:
        features["public_clip_ratio"] = 0.5

    # Deletion rate
    if "is_deleted" in total_clip_df.columns:
        features["deletion_rate"] = user_groups["is_deleted"].mean().values
    else:
        features["deletion_rate"] = 0

    # Pro user ratio
    if "is_pro_user" in total_clip_df.columns:
        features["pro_user_ratio"] = user_groups["is_pro_user"].mean().values
    else:
        features["pro_user_ratio"] = 0

    # Continued clips - vectorized approach
    if "continued_parent" in total_clip_df.columns:
        # Calculate the ratio of clips with continued_parent not null
        total_clip_df["has_continued_parent"] = total_clip_df[
            "continued_parent"
        ].notna()
        continued_ratio = total_clip_df.groupby("user_id")[
            "has_continued_parent"
        ].mean()
        features["continued_clip_ratio"] = continued_ratio.values
        # Clean up
        total_clip_df.drop("has_continued_parent", axis=1, inplace=True)
    else:
        features["continued_clip_ratio"] = 0

    # Lyrics generation detection - vectorized approach
    if "metadata" in total_clip_df.columns:
        # Use the parse_metadata_field function that handles both dict and string

        # Apply the parsing function to each metadata entry
        total_clip_df["has_gpt_desc"] = total_clip_df["metadata"].apply(
            parse_metadata_field
        )

        # Aggregate at user level
        gpt_stats = total_clip_df.groupby("user_id")["has_gpt_desc"].agg(
            ["any", "mean"]
        )
        features["has_lyrics_generation"] = gpt_stats["any"].values
        features["lyrics_generation_ratio"] = gpt_stats["mean"].fillna(0).values

        # Clean up
        total_clip_df.drop("has_gpt_desc", axis=1, inplace=True)
    else:
        features["has_lyrics_generation"] = False
        features["lyrics_generation_ratio"] = 0

    return features


def create_engagement_features(
    total_clip_df: pd.DataFrame, boosts_df: pd.DataFrame, reaction_df: pd.DataFrame
) -> pd.DataFrame:
    """Create engagement-based features."""
    logger.info("Creating engagement features")

    # Merge total_clip_df with boosts data
    clip_boosts = pd.merge(
        total_clip_df[["id", "user_id"]],
        boosts_df,
        left_on="id",
        right_on="clip_id",
        how="left",
    )

    user_groups = total_clip_df.groupby("user_id")

    features = pd.DataFrame()
    features["user_id"] = user_groups.size().index

    # Play count
    if "play_count" in total_clip_df.columns:
        features["avg_play_count_per_clip"] = user_groups["play_count"].mean().fillna(0)
    else:
        features["avg_play_count_per_clip"] = 0

    # Upvote count
    if "upvote_count" in total_clip_df.columns:
        features["avg_upvote_count"] = user_groups["upvote_count"].mean().fillna(0)
    else:
        features["avg_upvote_count"] = 0

    # Download counts from boosts
    if not clip_boosts.empty:
        boost_groups = clip_boosts.groupby("user_id")

        if (
            "download_audio_count" in clip_boosts.columns
            and "download_video_count" in clip_boosts.columns
        ):
            clip_boosts["total_downloads"] = clip_boosts["download_audio_count"].fillna(
                0
            ) + clip_boosts["download_video_count"].fillna(0)
            features["avg_download_count"] = (
                boost_groups["total_downloads"].mean().fillna(0)
            )
        else:
            features["avg_download_count"] = 0

        if "share_count" in clip_boosts.columns:
            features["share_rate"] = boost_groups["share_count"].mean().fillna(0)
        else:
            features["share_rate"] = 0
    else:
        features["avg_download_count"] = 0
        features["share_rate"] = 0

    # Fill missing values
    features = features.fillna(0)

    # Engagement score (composite)
    features["engagement_score"] = (
        features["avg_play_count_per_clip"] * 0.3
        + features["avg_upvote_count"] * 0.3
        + features["avg_download_count"] * 0.2
        + features["share_rate"] * 0.2
    )

    # Self engagement ratio
    if not reaction_df.empty:
        # Find reactions where user reacted to their own clips
        user_clips = total_clip_df[["id", "user_id"]].rename(columns={"id": "clip_id"})
        reactions_with_clip_owner = pd.merge(
            reaction_df[["clip_id", "user_id"]],
            user_clips,
            on="clip_id",
            suffixes=("_reactor", "_owner"),
        )

        # Count self reactions and total reactions - vectorized
        self_reactions_mask = (
            reactions_with_clip_owner["user_id_reactor"]
            == reactions_with_clip_owner["user_id_owner"]
        )

        # Group by reactor and count self reactions
        self_reaction_counts = (
            reactions_with_clip_owner[self_reactions_mask]
            .groupby("user_id_reactor")
            .size()
        )

        # Count total reactions per user
        total_reaction_counts = reaction_df.groupby("user_id").size()

        # Calculate ratio
        self_ratio = (self_reaction_counts / total_reaction_counts).fillna(0)
        features["self_engagement_ratio"] = self_ratio.reindex(
            features["user_id"]
        ).fillna(0)
    else:
        features["self_engagement_ratio"] = 0

    return features


def create_creation_source_features(total_clip_df: pd.DataFrame) -> pd.DataFrame:
    """Create features based on creation source."""
    logger.info("Creating creation source features")

    user_groups = total_clip_df.groupby("user_id")

    features = pd.DataFrame()
    features["user_id"] = user_groups.size().index

    if "creation_source" in total_clip_df.columns:
        # Source diversity
        features["creation_source_diversity"] = user_groups["creation_source"].nunique()

        # Primary source - use value_counts approach
        source_counts = total_clip_df.groupby("user_id")[
            "creation_source"
        ].value_counts()
        primary_sources = source_counts.groupby(level=0).idxmax().str[-1]
        features["primary_creation_source"] = primary_sources.fillna("unknown")

        # Source distribution - vectorized approach using pivot
        # Get top 3 creation sources overall
        top_sources = total_clip_df["creation_source"].value_counts().head(3).index

        # Create a filtered dataframe with only top sources
        filtered_df = total_clip_df[total_clip_df["creation_source"].isin(top_sources)]

        # Use crosstab for efficient percentage calculation
        if len(filtered_df) > 0:
            source_crosstab = pd.crosstab(
                filtered_df["user_id"],
                filtered_df["creation_source"],
                normalize="index",
            )

            # Rename columns with proper prefix
            for source in source_crosstab.columns:
                col_name = f'source_pct_{str(source).replace(" ", "_").lower()}'
                features[col_name] = (
                    source_crosstab[source].reindex(features["user_id"]).fillna(0)
                )
    else:
        features["creation_source_diversity"] = 1
        features["primary_creation_source"] = "unknown"

    return features


def create_model_usage_features(total_clip_df: pd.DataFrame) -> pd.DataFrame:
    """Create features based on model usage patterns."""
    logger.info("Creating model usage features")

    user_groups = total_clip_df.groupby("user_id")

    features = pd.DataFrame()
    features["user_id"] = user_groups.size().index

    if "model_name" in total_clip_df.columns:
        # Model diversity
        features["model_diversity"] = user_groups["model_name"].nunique()

        # Primary model - use value_counts approach
        model_counts = total_clip_df.groupby("user_id")["model_name"].value_counts()
        primary_models = model_counts.groupby(level=0).idxmax().str[-1]
        features["primary_model"] = primary_models.fillna("unknown")

        # Model switch rate - vectorized approach
        # Sort by user and created_at, then check when model changes
        sorted_df = total_clip_df[["user_id", "created_at", "model_name"]].sort_values(
            ["user_id", "created_at"]
        )
        sorted_df["model_changed"] = (
            sorted_df.groupby("user_id")["model_name"].shift()
            != sorted_df["model_name"]
        )

        # Calculate switch rate per user
        switch_counts = (
            sorted_df.groupby("user_id")["model_changed"].sum() - 1
        )  # Subtract 1 for first entry
        user_counts = user_groups.size() - 1  # Total transitions
        switch_rate = (switch_counts / user_counts).fillna(0)
        features["model_switch_rate"] = switch_rate.reindex(features["user_id"]).fillna(
            0
        )

        # Latest model adoption - simplified approach
        # Assume models with 'v' followed by numbers are versioned, higher is newer
        model_versions = (
            total_clip_df["model_name"]
            .str.extract(r"v(\d+)", expand=False)
            .astype(float)
        )
        if model_versions.notna().any():
            total_clip_df["model_version"] = model_versions.fillna(0)
            median_version = total_clip_df["model_version"].median()

            # Calculate percentage of clips using newer models per user
            newer_model_usage = total_clip_df.groupby("user_id")["model_version"].apply(
                lambda x: (x > median_version).mean()
            )
            features["latest_model_adoption"] = newer_model_usage.fillna(0.5)

            # Clean up
            total_clip_df.drop("model_version", axis=1, inplace=True)
        else:
            features["latest_model_adoption"] = 0.5
    else:
        features["model_diversity"] = 1
        features["primary_model"] = "unknown"
        features["model_switch_rate"] = 0
        features["latest_model_adoption"] = 0.5

    return features


def create_task_features(total_clip_df: pd.DataFrame) -> pd.DataFrame:
    """Create features based on task types."""
    logger.info("Creating task features")

    user_groups = total_clip_df.groupby("user_id")

    features = pd.DataFrame()
    features["user_id"] = user_groups.size().index

    if "task" in total_clip_df.columns:
        # Task diversity
        features["task_diversity"] = user_groups["task"].nunique()

        # Primary task - use value_counts approach
        task_counts = total_clip_df.groupby("user_id")["task"].value_counts()
        primary_tasks = task_counts.groupby(level=0).idxmax().str[-1]
        features["primary_task"] = primary_tasks.fillna("unknown")

        # Task specialization score (inverse of diversity relative to total tasks)
        total_tasks = total_clip_df["task"].nunique()
        features["task_specialization_score"] = 1 - (features["task_diversity"] - 1) / (
            max(total_tasks - 1, 1)  # Avoid division by zero
        )
        features["task_specialization_score"] = features[
            "task_specialization_score"
        ].clip(0, 1)

        # Task distribution for top tasks - vectorized approach
        # Get top 3 tasks
        top_tasks = total_clip_df["task"].value_counts().head(3).index

        # Filter to only top tasks and create crosstab
        filtered_df = total_clip_df[total_clip_df["task"].isin(top_tasks)]

        if len(filtered_df) > 0:
            task_crosstab = pd.crosstab(
                filtered_df["user_id"], filtered_df["task"], normalize="index"
            )

            # Add percentage columns
            for task in task_crosstab.columns:
                col_name = f'task_pct_{str(task).replace(" ", "_").lower()}'
                features[col_name] = (
                    task_crosstab[task].reindex(features["user_id"]).fillna(0)
                )
    else:
        features["task_diversity"] = 1
        features["primary_task"] = "unknown"
        features["task_specialization_score"] = 1

    return features


def create_platform_features(total_clip_df: pd.DataFrame) -> pd.DataFrame:
    """Create features based on platform/source usage."""
    logger.info("Creating platform features")

    user_groups = total_clip_df.groupby("user_id")

    features = pd.DataFrame()
    features["user_id"] = user_groups.size().index

    if "source" in total_clip_df.columns:
        # Platform diversity
        features["platform_diversity"] = user_groups["source"].nunique()

        # Primary platform - use value_counts approach
        platform_counts = total_clip_df.groupby("user_id")["source"].value_counts()
        primary_platforms = platform_counts.groupby(level=0).idxmax().str[-1]
        features["primary_platform"] = primary_platforms.fillna("web")

        # Platform distribution - vectorized approach
        # Create crosstab for all platforms
        platform_crosstab = pd.crosstab(
            total_clip_df["user_id"], total_clip_df["source"], normalize="index"
        )

        # Add percentage columns for each platform
        for platform in ["web", "ios", "android"]:
            if platform in platform_crosstab.columns:
                features[f"{platform}_usage_pct"] = (
                    platform_crosstab[platform].reindex(features["user_id"]).fillna(0)
                )
            else:
                features[f"{platform}_usage_pct"] = 0

        # Mobile vs web ratio
        features["mobile_vs_web_ratio"] = (
            features["ios_usage_pct"] + features["android_usage_pct"]
        ) / (features["web_usage_pct"] + 0.001)  # Avoid division by zero

        # iOS preference score (among mobile users)
        mobile_total = features["ios_usage_pct"] + features["android_usage_pct"]
        features["ios_preference_score"] = np.where(
            mobile_total > 0, features["ios_usage_pct"] / mobile_total, 0
        )

        # Cross-platform user
        features["cross_platform_user"] = features["platform_diversity"] > 1
    else:
        features["platform_diversity"] = 1
        features["primary_platform"] = "web"
        features["web_usage_pct"] = 1
        features["ios_usage_pct"] = 0
        features["android_usage_pct"] = 0
        features["mobile_vs_web_ratio"] = 0
        features["ios_preference_score"] = 0
        features["cross_platform_user"] = False

    return features


def create_prompt_features(
    total_clip_df: pd.DataFrame, boosts_df: pd.DataFrame
) -> pd.DataFrame:
    """Create features based on prompt behavior."""
    logger.info("Creating prompt features")

    user_groups = total_clip_df.groupby("user_id")

    features = pd.DataFrame()
    features["user_id"] = user_groups.size().index

    # Prompt length - vectorized
    if "prompt_text" in total_clip_df.columns:
        # Calculate prompt lengths
        total_clip_df["prompt_length"] = (
            total_clip_df["prompt_text"].fillna("").str.len()
        )
        features["avg_prompt_length"] = user_groups["prompt_length"].mean().fillna(0)

        # Clean up
        total_clip_df.drop("prompt_length", axis=1, inplace=True)

        # Unique prompts ratio - vectorized
        total_prompts = user_groups.size()
        unique_prompts = user_groups["prompt_text"].nunique()
        features["unique_prompts_ratio"] = (unique_prompts / total_prompts).fillna(0)
        features["prompt_diversity_score"] = features["unique_prompts_ratio"]
    else:
        features["avg_prompt_length"] = 50  # Default
        features["unique_prompts_ratio"] = 0.5
        features["prompt_diversity_score"] = 0.5

    # Prompt reuse from boosts
    if not boosts_df.empty and "reuse_prompt_count" in boosts_df.columns:
        # Merge to get reuse counts per user
        clip_reuse = pd.merge(
            total_clip_df[["id", "user_id"]],
            boosts_df[["clip_id", "reuse_prompt_count"]],
            left_on="id",
            right_on="clip_id",
            how="left",
        )

        # Calculate average reuse rate
        reuse_stats = clip_reuse.groupby("user_id")["reuse_prompt_count"].mean()
        features["prompt_reuse_rate"] = reuse_stats.reindex(features["user_id"]).fillna(
            0
        )
    else:
        features["prompt_reuse_rate"] = 0

    return features


def create_playlist_features(
    total_clip_df: pd.DataFrame, playlist_df: pd.DataFrame
) -> pd.DataFrame:
    """Create features based on playlist participation."""
    logger.info("Creating playlist features")

    # Join playlist data with total_clip_df to get user_id
    playlist_with_users = pd.merge(
        playlist_df,
        total_clip_df[["id", "user_id"]],
        left_on="clip_id",
        right_on="id",
        how="inner",
    )

    # Drop the redundant 'id' column if it exists to avoid confusion
    if "id" in playlist_with_users.columns:
        playlist_with_users.drop("id", axis=1, inplace=True)

    # Get all users
    all_users = total_clip_df["user_id"].unique()

    features = pd.DataFrame()
    features["user_id"] = all_users

    if not playlist_with_users.empty and "user_id" in playlist_with_users.columns:
        # Clips in playlists per user
        clips_in_playlists = playlist_with_users.groupby("user_id")["clip_id"].nunique()
        features["clips_in_playlists_count"] = (
            clips_in_playlists.reindex(features["user_id"]).fillna(0).values
        )

        # Total clips per user
        total_clips = total_clip_df.groupby("user_id").size()
        features["total_clips"] = (
            total_clips.reindex(features["user_id"]).fillna(0).values
        )

        # Participation rate
        features["playlist_participation_rate"] = (
            features["clips_in_playlists_count"]
            / (features["total_clips"] + 1)  # Add 1 to avoid division by zero
        ).fillna(0)

        # Average playlist position
        if "relative_index" in playlist_with_users.columns:
            avg_position = playlist_with_users.groupby("user_id")[
                "relative_index"
            ].mean()
            features["avg_playlist_position"] = (
                avg_position.reindex(features["user_id"]).fillna(0).values
            )
        else:
            features["avg_playlist_position"] = 0

        # Unique playlists count
        unique_playlists = playlist_with_users.groupby("user_id")[
            "playlist_id"
        ].nunique()
        features["unique_playlists_count"] = (
            unique_playlists.reindex(features["user_id"]).fillna(0).values
        )
    else:
        features["clips_in_playlists_count"] = 0
        # Calculate total clips per user even when no playlist data
        total_clips = total_clip_df.groupby("user_id").size()
        features["total_clips"] = (
            total_clips.reindex(features["user_id"]).fillna(0).values
        )
        features["playlist_participation_rate"] = 0
        features["avg_playlist_position"] = 0
        features["unique_playlists_count"] = 0

    return features


def aggregate_all_features(
    feature_dfs: List[pd.DataFrame], users: np.ndarray
) -> pd.DataFrame:
    """Aggregate all feature dataframes into a single dataframe."""
    logger.info("Aggregating all features")

    # Start with base dataframe of all users
    result = pd.DataFrame({"user_id": users})

    # Merge all feature dataframes
    for df in feature_dfs:
        if not df.empty:
            result = pd.merge(result, df, on="user_id", how="left")

    # Fill any remaining NaN values with appropriate defaults
    numeric_columns = result.select_dtypes(include=[np.number]).columns
    result[numeric_columns] = result[numeric_columns].fillna(0)

    # Fill categorical columns
    categorical_columns = result.select_dtypes(include=["object"]).columns
    if "user_id" in categorical_columns:
        categorical_columns = categorical_columns.drop("user_id")  # Don't fill user_id
    result[categorical_columns] = result[categorical_columns].fillna("unknown")

    # Fill boolean columns
    bool_columns = result.select_dtypes(include=["bool"]).columns
    result[bool_columns] = result[bool_columns].fillna(False)

    logger.info(f"Aggregated features shape: {result.shape}")

    return result
