"""
Content creation feature extraction for user clustering.
"""

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


class ContentFeatureExtractor:
    """Extract features from user content creation data."""

    def __init__(self, analysis_date: pd.Timestamp = None):
        self.analysis_date = analysis_date or pd.to_datetime("2025-06-17")
        self.days_30 = timedelta(days=30)

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

        Args:
            total_clip_df: DataFrame containing all clips data
            user_ids: Series of user IDs to process
            verbose: Whether to print progress messages

        Returns:
            Tuple of (DataFrame with content creation features, summary dict)
        """
        if verbose:
            print("   Processing clips from filtered users...")

        # Filter clips for specified users
        user_clips = total_clip_df[total_clip_df["user_id"].isin(user_ids)].copy()

        # Check available columns
        available_cols = set(user_clips.columns)
        if verbose:
            print(f"   Available columns: {sorted(list(available_cols))[:10]}...")

        # Basic clip counts - includes non-creators
        clip_counts = self._get_basic_clip_counts(user_clips)

        # IMPORTANT: Add all users, including those with zero clips (non-creators)
        all_users_df = pd.DataFrame({"user_id": user_ids})
        clip_counts = all_users_df.merge(clip_counts, on="user_id", how="left")
        clip_counts["total_clips_created"] = (
            clip_counts["total_clips_created"].fillna(0).astype(int)
        )

        # Add temporal features if available
        if "created_at" in available_cols and len(user_clips) > 0:
            clip_counts = self._add_temporal_features(user_clips, clip_counts, verbose)
        else:
            # Set default temporal features for all users
            clip_counts["days_creating"] = 0
            clip_counts["is_recent_creator"] = 0
            clip_counts["clip_creation_rate"] = 0

        # Extract categorical features
        categorical_features = []

        # Model features
        if "model_name" in available_cols and len(user_clips) > 0:
            model_features = self._extract_model_features(user_clips, verbose)
            categorical_features.append(model_features)

        # Task features
        if "task" in available_cols and len(user_clips) > 0:
            task_features = self._extract_categorical_features(
                user_clips, "user_id", "task", "task", top_n=5
            )
            categorical_features.append(task_features)

        # Source features
        if "source" in available_cols and len(user_clips) > 0:
            source_features = self._extract_categorical_features(
                user_clips, "user_id", "source", "source", top_n=3
            )
            categorical_features.append(source_features)

        # Type features (public/private)
        if "type" in available_cols and len(user_clips) > 0:
            type_features = self._extract_type_features(user_clips)
            categorical_features.append(type_features)

        # Merge all features
        all_features = self._merge_categorical_features(
            clip_counts, categorical_features
        )

        # Add generation pattern features
        all_features = self._add_generation_patterns(all_features)

        # Create summary statistics
        summary = self._create_summary(all_features, user_clips, verbose)

        return all_features, summary

    def _get_basic_clip_counts(self, user_clips: pd.DataFrame) -> pd.DataFrame:
        """Get basic clip count statistics per user."""
        clip_counts = (
            user_clips.groupby("user_id")
            .size()
            .to_frame("total_clips_created")
            .reset_index()
        )
        return clip_counts

    def _add_temporal_features(
        self, user_clips: pd.DataFrame, clip_counts: pd.DataFrame, verbose: bool
    ) -> pd.DataFrame:
        """Add temporal features based on creation dates."""
        try:
            # Convert to datetime
            user_clips["created_at"] = pd.to_datetime(
                user_clips["created_at"], errors="coerce"
            )

            # Remove timezone if present
            if (
                hasattr(user_clips["created_at"].dt, "tz")
                and user_clips["created_at"].dt.tz is not None
            ):
                user_clips["created_at"] = user_clips["created_at"].dt.tz_localize(None)

            # Calculate date statistics
            date_stats = (
                user_clips.groupby("user_id")["created_at"]
                .agg(["min", "max"])
                .reset_index()
            )
            date_stats.columns = ["user_id", "first_clip_date", "last_clip_date"]

            # Calculate days creating
            date_stats["days_creating"] = (
                self.analysis_date - date_stats["first_clip_date"]
            ).dt.days.clip(lower=1)

            # Check recent activity
            date_stats["is_recent_creator"] = (
                (self.analysis_date - date_stats["last_clip_date"]).dt.days <= 30
            ).astype(int)

            # Merge with clip counts
            clip_counts = clip_counts.merge(
                date_stats[["user_id", "days_creating", "is_recent_creator"]],
                on="user_id",
                how="left",
            )

            # Fill NaN values immediately after merge
            clip_counts["days_creating"] = clip_counts["days_creating"].fillna(30)
            clip_counts["is_recent_creator"] = clip_counts["is_recent_creator"].fillna(
                0
            )

            # Calculate creation rate - use np.where to avoid NaN
            clip_counts["clip_creation_rate"] = np.where(
                clip_counts["days_creating"] > 0,
                clip_counts["total_clips_created"] / clip_counts["days_creating"],
                0,
            ).round(3)

        except Exception as e:
            if verbose:
                print(f"   Warning: Could not process creation dates: {e}")
            clip_counts["clip_creation_rate"] = 0
            clip_counts["is_recent_creator"] = 0
            clip_counts["days_creating"] = 30

        return clip_counts

    def _extract_model_features(
        self, user_clips: pd.DataFrame, verbose: bool
    ) -> pd.DataFrame:
        """Extract model usage features."""
        # Identify advanced model usage
        user_clips["is_advanced_model"] = (
            user_clips["model_name"]
            .str.contains(
                r"v4p5|v4\.5|v4(?!\.0)|premium|pro", case=False, regex=True, na=False
            )
            .astype(int)
        )

        # Calculate advanced model stats
        advanced_model_stats = (
            user_clips.groupby("user_id")
            .agg({"is_advanced_model": ["sum", "mean"]})
            .reset_index()
        )
        advanced_model_stats.columns = [
            "user_id",
            "advanced_model_count",
            "advanced_model_ratio",
        ]

        # Extract model diversity
        model_features = self._extract_categorical_features(
            user_clips,
            "user_id",
            "model_name",
            "model",
            top_n=3,
            include_diversity=True,
        )

        # Merge advanced model stats
        model_features = model_features.merge(
            advanced_model_stats, on="user_id", how="left"
        )

        # Fill NaN values from merge
        model_features["advanced_model_count"] = model_features[
            "advanced_model_count"
        ].fillna(0)
        model_features["advanced_model_ratio"] = model_features[
            "advanced_model_ratio"
        ].fillna(0)

        # Check for v4p5 specifically
        try:
            user_clips["uses_v4p5"] = (
                user_clips["model_name"]
                .str.contains(r"v4p5|v4\.5", case=False, na=False)
                .astype(int)
            )

            v4p5_usage = (
                user_clips.groupby("user_id")["uses_v4p5"]
                .agg(["sum", "mean"])
                .reset_index()
            )
            v4p5_usage.columns = ["user_id", "v4p5_count", "v4p5_ratio"]
            model_features = model_features.merge(v4p5_usage, on="user_id", how="left")
            # Fill NaN values from merge
            model_features["v4p5_count"] = model_features["v4p5_count"].fillna(0)
            model_features["v4p5_ratio"] = model_features["v4p5_ratio"].fillna(0)
        except:
            pass

        return model_features

    def _extract_categorical_features(
        self,
        df: pd.DataFrame,
        user_col: str,
        cat_col: str,
        prefix: str,
        top_n: int = 3,
        include_diversity: bool = True,
    ) -> pd.DataFrame:
        """Extract features from categorical columns efficiently."""
        if cat_col not in df.columns:
            return pd.DataFrame()

        # Create pivot table
        pivot = pd.crosstab(df[user_col], df[cat_col], normalize="index")

        # Get top categories
        top_cats = df[cat_col].value_counts().head(top_n).index

        # Create feature dict
        features = {}

        # Add top category fractions
        for i, cat in enumerate(top_cats):
            if cat in pivot.columns:
                features[f"{prefix}_top{i+1}_{str(cat)[:20]}"] = pivot[cat]

        if include_diversity:
            # Calculate diversity (entropy)
            pivot_safe = pivot + 1e-10
            entropy = -(pivot_safe * np.log2(pivot_safe)).sum(axis=1)
            features[f"{prefix}_diversity"] = entropy

            # Count of unique values
            features[f"{prefix}_unique_count"] = (pivot > 0).sum(axis=1)

            # Concentration (Herfindahl index)
            features[f"{prefix}_concentration"] = (pivot**2).sum(axis=1)

            # Dominant category encoded
            if len(pivot.columns) > 0:
                dominant = pivot.idxmax(axis=1)
                features[f"{prefix}_dominant_encoded"] = dominant.apply(
                    lambda x: hash(str(x)) % 1000
                )

        return pd.DataFrame(features).reset_index()

    def _extract_type_features(self, user_clips: pd.DataFrame) -> pd.DataFrame:
        """Extract features related to clip types (public/private)."""
        type_pivot = pd.crosstab(
            user_clips["user_id"], user_clips["type"], normalize="index"
        )
        type_features = pd.DataFrame(index=type_pivot.index)

        # Check for public content
        if "public" in type_pivot.columns:
            type_features["public_clip_ratio"] = type_pivot["public"]

            # Get absolute count of public clips
            public_counts = (
                user_clips[user_clips["type"] == "public"].groupby("user_id").size()
            )
            type_features["public_clip_count"] = public_counts.reindex(
                type_features.index, fill_value=0
            )

        type_features["type_diversity"] = (type_pivot > 0).sum(axis=1)
        return type_features.reset_index()

    def _merge_categorical_features(
        self, clip_counts: pd.DataFrame, categorical_features: List[pd.DataFrame]
    ) -> pd.DataFrame:
        """Merge all categorical features efficiently."""
        all_features = clip_counts

        # Merge each feature set
        for feat_df in categorical_features:
            if len(feat_df) > 0:
                all_features = all_features.merge(feat_df, on="user_id", how="left")

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

        return all_features

    def _add_generation_patterns(self, features_df: pd.DataFrame) -> pd.DataFrame:
        """Add features related to generation patterns."""
        # Daily generation rate to identify free tier users
        if "days_creating" in features_df.columns:
            # Use np.where to avoid division by zero and NaN
            features_df["daily_generation_rate"] = np.where(
                features_df["days_creating"] > 0,
                features_df["total_clips_created"] / features_df["days_creating"],
                0,
            ).round(2)

            # Identify likely free tier users (≤20 generations per day)
            features_df["is_free_tier_pattern"] = (
                features_df["daily_generation_rate"] <= 20
            ).astype(int)
        else:
            features_df["daily_generation_rate"] = 0
            features_df["is_free_tier_pattern"] = 0

        return features_df

    def merge_features(
        self, features_df: pd.DataFrame, content_features: pd.DataFrame
    ) -> pd.DataFrame:
        """Merge content features into main features dataframe."""
        # Get columns to merge (exclude user_id)
        merge_cols = [col for col in content_features.columns if col != "user_id"]

        # Check for existing features
        new_cols = [col for col in merge_cols if col not in features_df.columns]

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

        # Perform merge
        features_df = features_df.merge(
            content_features[["user_id"] + new_cols], 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

    def _create_summary(
        self, all_features: pd.DataFrame, user_clips: pd.DataFrame, verbose: bool
    ) -> Dict[str, Any]:
        """Create summary statistics for content features."""
        summary = {
            "total_creators": len(all_features),
            "total_clips": len(user_clips),
            "active_creators": (all_features["total_clips_created"] > 0).sum()
            if "total_clips_created" in all_features.columns
            else 0,
            "recent_creators": all_features["is_recent_creator"].sum()
            if "is_recent_creator" in all_features.columns
            else 0,
        }

        # Add creation rate statistics
        if "clip_creation_rate" in all_features.columns:
            summary["avg_creation_rate"] = all_features["clip_creation_rate"].mean()

        # Add model diversity stats
        if "model_diversity" in all_features.columns:
            summary["avg_model_diversity"] = all_features["model_diversity"].mean()

        # Add advanced model usage stats
        if "advanced_model_ratio" in all_features.columns:
            summary["advanced_model_users"] = (
                all_features["advanced_model_ratio"] > 0
            ).sum()
            summary["avg_advanced_model_ratio"] = all_features[
                "advanced_model_ratio"
            ].mean()

        # Add v4p5 usage stats
        if "v4p5_ratio" in all_features.columns:
            summary["v4p5_users"] = (all_features["v4p5_ratio"] > 0).sum()
            summary["avg_v4p5_ratio"] = all_features["v4p5_ratio"].mean()

        # Add public content stats
        if "public_clip_ratio" in all_features.columns:
            summary["public_content_creators"] = (
                all_features["public_clip_ratio"] > 0
            ).sum()
            summary["avg_public_ratio"] = all_features["public_clip_ratio"].mean()

        return summary
