"""
Reaction feature extraction for user clustering.
"""

import pandas as pd
import numpy as np
from typing import Dict, List, Optional, Tuple, Any
import gc


class ReactionFeatureExtractor:
    """Extract features from user reactions data."""

    def __init__(self, eps: float = 1e-8):
        self.eps = eps
        self.base_features = [
            "user_id",
            "reaction_like_ratio",
            "reaction_dislike_ratio",
            "reaction_frequency",
            "creator_feedback_ratio",
            "community_interaction_score",
            "controversy_score",
            "engagement_ratio",
        ]
        self.optional_features = [
            "total_plays",
            "avg_play_count",
            "play_frequency",
            "play_engagement_ratio",
            "flag_rate",
        ]

    def extract_features(
        self, reaction_df: pd.DataFrame, user_ids: pd.Series, verbose: bool = True
    ) -> Tuple[pd.DataFrame, Dict[str, Any]]:
        """
        Extract reaction features for specified users.

        Args:
            reaction_df: DataFrame containing reaction data
            user_ids: Series of user IDs to process
            verbose: Whether to print progress messages

        Returns:
            Tuple of (reaction_stats DataFrame, summary dict)
        """
        if verbose:
            print("   Processing reactions for filtered users...")

        # Filter reactions for specified users
        user_reactions = reaction_df[reaction_df["user_id"].isin(user_ids)].copy()

        # Check available columns
        available_cols = set(user_reactions.columns)
        if verbose:
            print(f"   Available columns: {sorted(available_cols)}")

        # Build aggregation dict
        agg_dict = self._build_aggregation_dict(user_reactions, available_cols)

        if not agg_dict or len(user_reactions) == 0:
            # Return empty dataframe with expected columns
            return self._create_empty_stats(), {"users_with_reactions": 0}

        # Perform aggregation
        reaction_stats = user_reactions.groupby("user_id", as_index=False).agg(agg_dict)

        # Process aggregated data
        reaction_stats = self._process_aggregated_stats(reaction_stats)

        # Calculate derived metrics
        reaction_stats = self._calculate_derived_metrics(reaction_stats)

        # Create summary
        summary = self._create_summary(reaction_stats)

        return reaction_stats, summary

    def _build_aggregation_dict(
        self, user_reactions: pd.DataFrame, available_cols: set
    ) -> Dict[str, Any]:
        """Build aggregation dictionary based on available columns."""
        agg_dict = {}

        # Process reaction types if available
        if "reaction_type" in available_cols:
            # Create dummy columns for reaction types
            user_reactions["is_like"] = (user_reactions["reaction_type"] == "L").astype(
                int
            )
            user_reactions["is_dislike"] = (
                user_reactions["reaction_type"] == "D"
            ).astype(int)

            agg_dict["reaction_type"] = "count"  # total reactions
            agg_dict["is_like"] = "sum"  # total likes
            agg_dict["is_dislike"] = "sum"  # total dislikes

        # Optional numeric columns
        optional_numeric_cols = [
            "play_count",
            "view_count",
            "share_count",
            "save_count",
        ]
        for col in optional_numeric_cols:
            if col in available_cols:
                agg_dict[col] = ["sum", "mean", "max"]

        # Boolean columns
        if "flagged" in available_cols:
            agg_dict["flagged"] = ["sum", "mean"]

        return agg_dict

    def _process_aggregated_stats(self, reaction_stats: pd.DataFrame) -> pd.DataFrame:
        """Flatten and rename aggregated columns."""
        # Flatten multi-level column names
        new_cols = []
        for col in reaction_stats.columns:
            if isinstance(col, tuple) and col[1]:
                new_cols.append(f"{col[0]}_{col[1]}")
            elif isinstance(col, tuple):
                new_cols.append(col[0])
            else:
                new_cols.append(col)
        reaction_stats.columns = new_cols

        # Rename core columns
        rename_dict = {
            "reaction_type_count": "total_reactions",
            "is_like_sum": "likes",
            "is_dislike_sum": "dislikes",
        }

        # Add optional column renames
        optional_numeric_cols = [
            "play_count",
            "view_count",
            "share_count",
            "save_count",
        ]
        for col in optional_numeric_cols:
            if f"{col}_sum" in reaction_stats.columns:
                rename_dict[f"{col}_sum"] = f"total_{col.replace('_count', 's')}"
                rename_dict[f"{col}_mean"] = f"avg_{col}"
                rename_dict[f"{col}_max"] = f"max_{col}"

        if "flagged_sum" in reaction_stats.columns:
            rename_dict["flagged_sum"] = "total_flags"
            rename_dict["flagged_mean"] = "flag_rate"

        reaction_stats.rename(columns=rename_dict, inplace=True)

        # Ensure all required columns exist
        for col in ["total_reactions", "likes", "dislikes", "total_flags"]:
            if col not in reaction_stats.columns:
                reaction_stats[col] = 0

        return reaction_stats

    def _calculate_derived_metrics(self, reaction_stats: pd.DataFrame) -> pd.DataFrame:
        """Calculate derived metrics from base statistics."""
        # Reaction ratios - use np.where to avoid NaN
        reaction_stats["reaction_like_ratio"] = np.where(
            reaction_stats["total_reactions"] > 0,
            reaction_stats["likes"] / reaction_stats["total_reactions"],
            0,
        ).round(3)

        reaction_stats["reaction_dislike_ratio"] = np.where(
            reaction_stats["total_reactions"] > 0,
            reaction_stats["dislikes"] / reaction_stats["total_reactions"],
            0,
        ).round(3)

        # Frequency (per 30 days)
        reaction_stats["reaction_frequency"] = (
            reaction_stats["total_reactions"] / 30
        ).round(2)

        # Creator-focused metrics - use np.where to avoid NaN
        reaction_stats["creator_feedback_ratio"] = np.where(
            reaction_stats["total_reactions"] > 0,
            reaction_stats["likes"] / reaction_stats["total_reactions"],
            0,
        ).round(3)

        reaction_stats["community_interaction_score"] = (
            reaction_stats["total_reactions"] * reaction_stats["creator_feedback_ratio"]
        ).round(2)

        reaction_stats["controversy_score"] = np.where(
            reaction_stats["total_reactions"] > 0,
            (reaction_stats["dislikes"] + reaction_stats["total_flags"])
            / reaction_stats["total_reactions"],
            0,
        ).round(3)

        # Handle optional metrics
        if "total_plays" not in reaction_stats.columns:
            reaction_stats["total_plays"] = 0
            reaction_stats["avg_play_count"] = 0
            reaction_stats["play_frequency"] = 0
            reaction_stats["play_engagement_ratio"] = 0
        else:
            reaction_stats["play_frequency"] = (
                reaction_stats["total_plays"] / 30
            ).round(2)
            reaction_stats["play_engagement_ratio"] = np.where(
                reaction_stats["total_reactions"] > 0,
                reaction_stats["total_plays"] / reaction_stats["total_reactions"],
                0,
            ).round(3)

        if "flag_rate" not in reaction_stats.columns:
            reaction_stats["flag_rate"] = 0

        # General engagement ratio - use np.where to avoid NaN
        reaction_stats["engagement_ratio"] = np.where(
            reaction_stats["total_reactions"] > 0,
            reaction_stats["likes"] / reaction_stats["total_reactions"],
            0,
        ).round(3)

        return reaction_stats

    def _create_empty_stats(self) -> pd.DataFrame:
        """Create empty stats dataframe with expected columns."""
        columns = {
            "user_id": [],
            "total_reactions": [],
            "likes": [],
            "dislikes": [],
            "reaction_like_ratio": [],
            "reaction_dislike_ratio": [],
            "reaction_frequency": [],
            "total_plays": [],
            "avg_play_count": [],
            "play_frequency": [],
            "play_engagement_ratio": [],
            "engagement_ratio": [],
            "controversy_score": [],
            "flag_rate": [],
            "total_flags": [],
            "creator_feedback_ratio": [],
            "community_interaction_score": [],
        }
        return pd.DataFrame(columns)

    def _create_summary(self, reaction_stats: pd.DataFrame) -> Dict[str, Any]:
        """Create summary statistics."""
        summary = {
            "users_with_reactions": len(reaction_stats),
            "avg_reactions": reaction_stats["total_reactions"].mean()
            if len(reaction_stats) > 0
            else 0,
            "avg_engagement_ratio": reaction_stats["engagement_ratio"].mean()
            if len(reaction_stats) > 0
            else 0,
        }

        if (
            "total_plays" in reaction_stats.columns
            and reaction_stats["total_plays"].sum() > 0
        ):
            summary["avg_plays"] = reaction_stats["total_plays"].mean()

        return summary

    def get_features_to_merge(self, reaction_stats: pd.DataFrame) -> List[str]:
        """Get list of features to merge into main dataframe."""
        features_to_merge = self.base_features.copy()

        # Add optional features that exist
        for feature in self.optional_features:
            if feature in reaction_stats.columns:
                features_to_merge.append(feature)

        # Always include creator feedback metrics
        if "creator_feedback_ratio" not in features_to_merge:
            features_to_merge.append("creator_feedback_ratio")
        if "community_interaction_score" not in features_to_merge:
            features_to_merge.append("community_interaction_score")

        return features_to_merge

    def merge_features(
        self,
        features_df: pd.DataFrame,
        reaction_stats: pd.DataFrame,
        check_existing: bool = True,
    ) -> pd.DataFrame:
        """
        Merge reaction features into main features dataframe.

        Args:
            features_df: Main features dataframe
            reaction_stats: Reaction statistics dataframe
            check_existing: Whether to check for existing features to avoid duplicates

        Returns:
            Updated features dataframe
        """
        features_to_merge = self.get_features_to_merge(reaction_stats)

        # Check for existing features if requested
        if check_existing:
            # Remove features that already exist in features_df
            features_to_merge = [
                f
                for f in features_to_merge
                if f != "user_id" and f not in features_df.columns
            ]

            if not features_to_merge:
                print("   ⚠️ All reaction features already exist, skipping merge")
                return features_df

        # Perform merge
        features_df = features_df.merge(
            reaction_stats[["user_id"] + features_to_merge], on="user_id", how="left"
        )

        # Fill NaN values
        numeric_cols = features_df.select_dtypes(include=[np.number]).columns
        numeric_cols = [col for col in numeric_cols if col != "user_id"]
        features_df[numeric_cols] = features_df[numeric_cols].fillna(0)

        return features_df
