"""
Clustering operations for user segmentation.
"""

import pandas as pd
import numpy as np
from typing import Dict, List, Optional, Tuple, Any
from sklearn.preprocessing import StandardScaler
from sklearn.cluster import MiniBatchKMeans
from sklearn.metrics import silhouette_score, silhouette_samples
import gc
import warnings

warnings.filterwarnings("ignore", category=FutureWarning)


class UserClusterer:
    """Handle user clustering operations with memory efficiency."""

    def __init__(
        self,
        target_clusters: range = range(4, 7),
        min_silhouette: float = 0.10,
        sample_size: int = 5000,  # Reduced default for stability
        random_state: int = 42,
        max_memory_gb: float = 4.0,  # Maximum memory to use
    ):
        """
        Initialize UserClusterer with memory-efficient settings.

        Args:
            target_clusters: Range of cluster numbers to test
            min_silhouette: Minimum acceptable silhouette score
            sample_size: Size of sample for silhouette calculation (reduced for stability)
            random_state: Random seed for reproducibility
            max_memory_gb: Maximum memory to use in GB
        """
        self.target_clusters = target_clusters
        self.min_silhouette = min_silhouette
        self.sample_size = sample_size
        self.random_state = random_state
        self.max_memory_gb = max_memory_gb
        self.scaler = StandardScaler()
        self.best_model = None
        self.best_k = None
        self.best_silhouette = None

    def prepare_features(
        self,
        features_df: pd.DataFrame,
        feature_columns: List[str],
        verbose: bool = True,
    ) -> Tuple[pd.DataFrame, List[str]]:
        """
        Prepare features for clustering with memory management.

        Args:
            features_df: DataFrame with all features
            feature_columns: List of desired feature columns
            verbose: Whether to print progress messages

        Returns:
            Tuple of (prepared DataFrame, list of available features)
        """
        # Check which features actually exist
        available_features = [
            col for col in feature_columns if col in features_df.columns
        ]
        missing_features = [
            col for col in feature_columns if col not in features_df.columns
        ]

        if verbose:
            print(f"📊 Feature availability:")
            print(f"   Total defined: {len(feature_columns)}")
            print(f"   Available: {len(available_features)}")
            print(f"   Missing: {len(missing_features)}")
            if missing_features:
                print(
                    f"   Missing features: {', '.join(missing_features[:5])}{'...' if len(missing_features) > 5 else ''}"
                )

        # Prepare data with available features
        X = features_df[available_features].fillna(0)

        # More conservative sampling for large datasets
        n_samples = len(X)
        if n_samples > 100000:
            # For very large datasets, use even smaller samples
            actual_sample_size = min(self.sample_size, n_samples // 20)
            if verbose:
                print(f"   ⚠️ Large dataset detected ({n_samples:,} users)")
                print(f"   Using reduced sample size: {actual_sample_size:,}")
        else:
            actual_sample_size = min(self.sample_size, n_samples)

        # Sample with memory efficiency
        np.random.seed(self.random_state)
        sample_idx = np.random.choice(n_samples, actual_sample_size, replace=False)

        # Create sample more efficiently
        X_sample = X.iloc[sample_idx].copy()
        features_sample = features_df.iloc[sample_idx].copy()

        # Force garbage collection
        gc.collect()

        if verbose:
            print(
                f"\n📊 Clustering data: {len(available_features)} features | {len(X_sample):,} samples"
            )

        return features_sample, available_features

    def find_optimal_clusters(
        self, X: pd.DataFrame, verbose: bool = True
    ) -> Tuple[int, float, Any]:
        """
        Find optimal number of clusters using silhouette score with memory management.

        Args:
            X: Feature matrix
            verbose: Whether to print progress messages

        Returns:
            Tuple of (best k, best silhouette score, best model)
        """
        # Standardize features
        X_scaled = self.scaler.fit_transform(X)

        # Dynamic batch size based on data size
        n_samples = len(X)
        if n_samples > 50000:
            batch_size = 5000  # Larger batch for large datasets
        elif n_samples > 10000:
            batch_size = 2000
        else:
            batch_size = min(1000, n_samples // 2)

        # Test different k values
        if verbose:
            print(
                f"\n🔍 Testing k={list(self.target_clusters)} with memory-efficient settings..."
            )

        results = []
        for k in self.target_clusters:
            try:
                if verbose:
                    print(f"   Testing k={k}...", end="", flush=True)

                # Create model with conservative settings
                model = MiniBatchKMeans(
                    n_clusters=k,
                    random_state=self.random_state,
                    batch_size=batch_size,
                    n_init=5,  # Reduced from 10 for memory efficiency
                    max_iter=200,  # Reduced from 300
                    reassignment_ratio=0.01,  # More conservative reassignment
                    verbose=0,
                )

                # Fit the model
                clusters = model.fit_predict(X_scaled)

                # Calculate silhouette score with even more conservative sampling
                if n_samples > 50000:
                    # For very large datasets, use subsample
                    sil_sample_size = min(1000, n_samples // 50)
                    sil_sample_idx = np.random.choice(
                        n_samples, sil_sample_size, replace=False
                    )
                    score = silhouette_score(
                        X_scaled[sil_sample_idx], clusters[sil_sample_idx]
                    )
                else:
                    # For smaller datasets, use standard sampling
                    sil_sample_size = min(self.sample_size // 2, n_samples)
                    score = silhouette_score(
                        X_scaled, clusters, sample_size=sil_sample_size
                    )

                results.append((k, score, model))
                if verbose:
                    print(f" silhouette={score:.3f}")

                # Force garbage collection after each iteration
                gc.collect()

            except Exception as e:
                if verbose:
                    print(f" FAILED: {str(e)}")
                # Add a poor score for failed clustering
                results.append((k, -1.0, None))

        # Select best model (excluding failed ones)
        valid_results = [(k, s, m) for k, s, m in results if s > -1.0]

        if not valid_results:
            raise RuntimeError(
                "All clustering attempts failed. Try reducing sample_size or target_clusters."
            )

        self.best_k, self.best_silhouette, self.best_model = max(
            valid_results, key=lambda x: x[1]
        )

        if verbose:
            success_indicator = (
                "✅" if self.best_silhouette > self.min_silhouette else "⚠️"
            )
            print(
                f"\n🎯 Best: k={self.best_k}, silhouette={self.best_silhouette:.3f} ({success_indicator} target: {self.min_silhouette})"
            )
            print(
                f"   Batch size: {batch_size}, Sample size: {sil_sample_size if 'sil_sample_size' in locals() else 'adaptive'}"
            )
            if self.best_silhouette < self.min_silhouette:
                print(
                    f"   Note: Silhouette score below target, but proceeding with best available result"
                )

        return self.best_k, self.best_silhouette, self.best_model

    def apply_clustering(
        self, features_sample: pd.DataFrame, feature_columns: List[str]
    ) -> pd.DataFrame:
        """
        Apply clustering to the feature sample with memory efficiency.

        Args:
            features_sample: DataFrame with features
            feature_columns: List of feature columns used

        Returns:
            DataFrame with cluster assignments
        """
        if self.best_model is None:
            raise ValueError(
                "No model has been fitted yet. Run find_optimal_clusters first."
            )

        X = features_sample[feature_columns]
        X_scaled = self.scaler.transform(X)

        # Predict in batches for large datasets
        n_samples = len(X_scaled)
        if n_samples > 100000:
            # Process in chunks
            chunk_size = 50000
            predictions = []

            for i in range(0, n_samples, chunk_size):
                end_idx = min(i + chunk_size, n_samples)
                chunk_pred = self.best_model.predict(X_scaled[i:end_idx])
                predictions.extend(chunk_pred)
                gc.collect()  # Clean up after each chunk

            features_sample["cluster"] = predictions
        else:
            # Standard prediction for smaller datasets
            features_sample["cluster"] = self.best_model.predict(X_scaled)

        return features_sample

    def analyze_clusters(
        self, features_sample: pd.DataFrame, verbose: bool = True
    ) -> pd.DataFrame:
        """
        Analyze cluster characteristics with memory efficiency.

        Args:
            features_sample: DataFrame with cluster assignments
            verbose: Whether to print analysis

        Returns:
            DataFrame with cluster statistics
        """
        cluster_stats = []

        for cluster_id in range(self.best_k):
            cluster_data = features_sample[features_sample["cluster"] == cluster_id]

            # Basic stats with error handling
            stats = {
                "cluster_id": cluster_id,
                "size": len(cluster_data),
                "size_pct": len(cluster_data) / len(features_sample) * 100,
            }

            # Add means with safe handling
            numeric_cols = [
                ("engagement_score", "engagement_score_mean"),
                ("total_clips_created", "total_clips_created_mean"),
                ("reaction_frequency", "reaction_frequency_mean"),
                ("total_downloads", "total_downloads_mean"),
                ("share_count", "share_count_mean"),
                ("advanced_model_ratio", "advanced_model_ratio_mean"),
                ("v4p5_ratio", "v4p5_ratio_mean"),
                ("public_clip_ratio", "public_clip_ratio_mean"),
            ]

            for col_name, stat_name in numeric_cols:
                if col_name in cluster_data.columns and len(cluster_data) > 0:
                    stats[stat_name] = cluster_data[col_name].mean()
                else:
                    stats[stat_name] = 0

            cluster_stats.append(stats)

        cluster_stats_df = pd.DataFrame(cluster_stats)

        if verbose:
            self._print_cluster_analysis(features_sample, cluster_stats_df)

        return cluster_stats_df

    def _print_cluster_analysis(
        self, features_sample: pd.DataFrame, cluster_stats_df: pd.DataFrame
    ) -> None:
        """Print detailed cluster analysis."""
        print("\n📊 DETAILED CLUSTER ANALYSIS:")
        print("=" * 80)

        # Sort clusters by engagement score
        sorted_clusters = cluster_stats_df.sort_values(
            "engagement_score_mean", ascending=False
        )

        for _, row in sorted_clusters.iterrows():
            cluster_id = int(row["cluster_id"])
            print(f"\n{'='*80}")
            print(
                f"CLUSTER {cluster_id} - {row['size']:,} users ({row['size_pct']:.1f}%)"
            )
            print(f"{'='*80}")

            # Key metrics
            print(f"\n📊 Key Metrics:")
            print(f"  • Engagement Score: {row['engagement_score_mean']:.1f}")
            print(f"  • Avg Clips Created: {row['total_clips_created_mean']:.1f}")
            print(f"  • Reaction Frequency: {row['reaction_frequency_mean']:.2f}/day")
            print(f"  • Downloads: {row['total_downloads_mean']:.1f}")
            print(f"  • Shares: {row['share_count_mean']:.1f}")

            # Advanced features if available
            if "advanced_model_ratio_mean" in row:
                print(f"\n🤖 Model Usage:")
                print(
                    f"  • Advanced Model Usage: {row.get('advanced_model_ratio_mean', 0):.1%}"
                )
                print(f"  • v4p5 Usage: {row.get('v4p5_ratio_mean', 0):.1%}")

            if "public_clip_ratio_mean" in row:
                print(f"\n📢 Visibility:")
                print(f"  • Public Content: {row.get('public_clip_ratio_mean', 0):.1%}")

            # Interpret cluster
            interpretation = self._interpret_cluster(row)
            print(f"\n💡 Interpretation: {interpretation}")

    def _interpret_cluster(self, cluster_stats: pd.Series) -> str:
        """Generate interpretation for a cluster based on its statistics."""
        interpretations = []

        # Engagement level
        if cluster_stats["engagement_score_mean"] > 1000:
            interpretations.append("Elite engagement tier")
        elif cluster_stats["engagement_score_mean"] > 400:
            interpretations.append("Premium engagement tier")
        elif cluster_stats["engagement_score_mean"] > 200:
            interpretations.append("High engagement tier")
        elif cluster_stats["engagement_score_mean"] > 100:
            interpretations.append("Moderate engagement tier")
        else:
            interpretations.append("Growing engagement tier")

        # Creation volume
        if cluster_stats["total_clips_created_mean"] > 500:
            interpretations.append("mega content producers")
        elif cluster_stats["total_clips_created_mean"] > 100:
            interpretations.append("prolific creators")
        elif cluster_stats["total_clips_created_mean"] > 50:
            interpretations.append("active creators")
        elif cluster_stats["total_clips_created_mean"] > 20:
            interpretations.append("regular creators")
        else:
            interpretations.append("casual creators")

        # Advanced model usage
        if (
            "v4p5_ratio_mean" in cluster_stats
            and cluster_stats["v4p5_ratio_mean"] > 0.3
        ):
            interpretations.append("v4p5 power users")
        elif (
            "advanced_model_ratio_mean" in cluster_stats
            and cluster_stats["advanced_model_ratio_mean"] > 0.5
        ):
            interpretations.append("advanced model users")

        # Influencer indicators
        if (
            "public_clip_ratio_mean" in cluster_stats
            and cluster_stats["public_clip_ratio_mean"] > 0.5
        ):
            if cluster_stats["share_count_mean"] > 10:
                interpretations.append("music influencers")

        return " | ".join(interpretations[:3])  # Limit to top 3 characteristics

    def generate_cluster_names(
        self, features_sample: pd.DataFrame, cluster_stats_df: pd.DataFrame
    ) -> Dict[int, str]:
        """Generate descriptive names for clusters."""
        cluster_names = {}

        for i in range(self.best_k):
            cluster_data = features_sample[features_sample["cluster"] == i]
            row = cluster_stats_df[cluster_stats_df["cluster_id"] == i].iloc[0]

            name_parts = []

            # Engagement level
            if row["engagement_score_mean"] > 1000:
                name_parts.append("Elite")
            elif row["engagement_score_mean"] > 400:
                name_parts.append("VeryHigh")
            elif row["engagement_score_mean"] > 200:
                name_parts.append("High")
            elif row["engagement_score_mean"] > 100:
                name_parts.append("Moderate")
            else:
                name_parts.append("Low")

            # Creation volume
            if row["total_clips_created_mean"] > 500:
                name_parts.append("MegaCreators")
            elif row["total_clips_created_mean"] > 100:
                name_parts.append("Prolific")
            elif row["total_clips_created_mean"] > 50:
                name_parts.append("Active")
            elif row["total_clips_created_mean"] > 20:
                name_parts.append("Regular")
            else:
                name_parts.append("Casual")

            # Special characteristics
            if "v4p5_ratio_mean" in row and row["v4p5_ratio_mean"] > 0.3:
                name_parts.append("ProUsers")
            elif row["reaction_frequency_mean"] > 2:
                name_parts.append("Interactive")
            elif row["total_downloads_mean"] > 50:
                name_parts.append("Collectors")

            cluster_names[i] = "_".join(name_parts)

        return cluster_names
