"""
Output Generator Module
======================
Handles creation and saving of clustering outputs.
"""

import pandas as pd
import numpy as np
import joblib
import logging
from pathlib import Path
from typing import Dict, List, Any, Tuple

logger = logging.getLogger(__name__)


def create_cluster_assignments(
    user_ids: pd.Series, cluster_labels: np.ndarray, kmeans_model: Any
) -> pd.DataFrame:
    """
    Create dataframe with user cluster assignments and distances.

    Args:
        user_ids: Series of user IDs
        cluster_labels: Cluster assignments
        kmeans_model: Fitted k-means model

    Returns:
        DataFrame with user_id, cluster_label, and cluster_distance
    """
    logger.info("Creating cluster assignments")

    # Create base assignments
    assignments = pd.DataFrame({"user_id": user_ids, "cluster_label": cluster_labels})

    # Calculate distance to assigned cluster center
    # Note: This requires the original scaled features
    # For now, we'll skip distance calculation in this function
    # It should be calculated in the main pipeline where we have access to X_scaled

    logger.info(f"Created assignments for {len(assignments)} users")

    return assignments


def create_cluster_profiles(
    user_features: pd.DataFrame, cluster_labels: np.ndarray
) -> Tuple[pd.DataFrame, Dict[str, List[str]]]:
    """
    Create profiles showing mean feature values for each cluster.

    Args:
        user_features: DataFrame with all user features
        cluster_labels: Cluster assignments

    Returns:
        Tuple of (cluster_profiles DataFrame, top_features_per_cluster dict)
    """
    logger.info("Creating cluster profiles")

    # Add cluster labels to features
    features_with_clusters = user_features.copy()
    features_with_clusters["cluster_label"] = cluster_labels

    # Group by cluster and calculate statistics
    numeric_columns = features_with_clusters.select_dtypes(
        include=[np.number]
    ).columns.tolist()
    if "user_id" in numeric_columns:
        numeric_columns.remove("user_id")  # Remove user_id from aggregation

    # Calculate mean for numeric features
    cluster_means = features_with_clusters.groupby("cluster_label")[
        numeric_columns
    ].mean()

    # Calculate mode for categorical features - vectorized approach
    categorical_columns = features_with_clusters.select_dtypes(
        include=["object", "category"]
    ).columns.tolist()
    if "user_id" in categorical_columns:
        categorical_columns.remove("user_id")

    cluster_modes = pd.DataFrame(index=cluster_means.index)

    if categorical_columns:
        # Use value_counts to find modes efficiently
        for col in categorical_columns:
            mode_series = (
                features_with_clusters.groupby("cluster_label")[col]
                .value_counts()
                .groupby(level=0)
                .idxmax()
                .str[-1]
            )
            cluster_modes[col] = mode_series

    # Calculate boolean feature percentages
    bool_columns = features_with_clusters.select_dtypes(
        include=["bool"]
    ).columns.tolist()
    cluster_bool_pcts = (
        features_with_clusters.groupby("cluster_label")[bool_columns].mean() * 100
    )

    # Combine all profiles
    cluster_profiles = pd.concat(
        [cluster_means, cluster_modes, cluster_bool_pcts], axis=1
    )

    # Add cluster sizes
    cluster_sizes = features_with_clusters["cluster_label"].value_counts().sort_index()
    cluster_profiles.insert(0, "cluster_size", cluster_sizes)
    cluster_profiles.insert(
        1, "cluster_size_pct", cluster_sizes / len(features_with_clusters) * 100
    )

    # Identify top distinguishing features for each cluster
    # Calculate z-scores for each feature in each cluster
    overall_means = cluster_means.mean()
    overall_stds = cluster_means.std()

    z_scores = (cluster_means - overall_means) / (overall_stds + 1e-8)

    # Get top 5 distinguishing features for each cluster
    top_features_per_cluster = {}
    for cluster in z_scores.index:
        cluster_z = z_scores.loc[cluster].abs().sort_values(ascending=False)
        top_features_per_cluster[f"cluster_{cluster}_top_features"] = cluster_z.head(
            5
        ).index.tolist()

    logger.info(f"Created profiles for {len(cluster_profiles)} clusters")

    return cluster_profiles, top_features_per_cluster


def map_clusters_to_types(cluster_profiles: pd.DataFrame) -> Dict[int, str]:
    """
    Map clusters to expected user types based on their profiles.

    Enhanced to recognize more detailed user segments:
    - bot: Automated/suspicious accounts
    - inactive: Very low activity users
    - casual: Occasional users, mostly lyrics
    - explorer: Trying different features/models
    - creator: Active content creators
    - power_user: Heavy users with high engagement
    - pro_serious: Professional users with focused usage

    Args:
        cluster_profiles: DataFrame with cluster statistics

    Returns:
        Dictionary mapping cluster labels to user types
    """
    logger.info("Mapping clusters to expected types")

    cluster_mapping = {}

    for cluster_idx in cluster_profiles.index:
        profile = cluster_profiles.loc[cluster_idx]

        # Helper function to get numeric value
        def get_numeric(key, default=0):
            val = profile.get(key, default)
            if isinstance(val, str):
                try:
                    return float(val)
                except:
                    return default
            return val if pd.notna(val) else default

        # Bot detection criteria - automated or suspicious behavior
        if (
            get_numeric("has_any_reactions", 100) < 20  # Less than 20% have reactions
            and get_numeric("consumption_ratio", 1) < 0.2  # Very low consumption
            and get_numeric("total_clips_created", 0) > 50  # High creation volume
            and get_numeric("creation_burst_score", 0) > 20  # High burst creation
        ):
            cluster_mapping[cluster_idx] = "bot"

        # Inactive users - very low activity
        elif (
            get_numeric("total_clips_created", 0) < 3  # Less than 3 clips
            and get_numeric("days_active", 0) <= 1  # Active 1 day or less
            and get_numeric("engagement_score", 0) < 0.1  # Low engagement
        ):
            cluster_mapping[cluster_idx] = "inactive"

        # Explorer users - trying different features
        elif (
            get_numeric("model_diversity", 0) > 1.5  # Uses multiple models
            or get_numeric("task_diversity", 0) > 1.5  # Tries different tasks
            or (
                get_numeric("platform_diversity", 0) > 0.5  # Uses multiple platforms
                and get_numeric("total_clips_created", 0) < 100
            )  # Not extreme creator
        ):
            cluster_mapping[cluster_idx] = "explorer"

        # Power users - heavy usage across the board
        elif (
            get_numeric("total_clips_created", 0) > 100  # High volume
            and get_numeric("engagement_score", 0) > 0.5  # High engagement
            and get_numeric("days_active", 0) > 5  # Consistently active
            and get_numeric("has_any_reactions", 0) > 80  # Social user
        ):
            cluster_mapping[cluster_idx] = "power_user"

        # Content creators - focused on creation
        elif (
            get_numeric("total_clips_created", 0) > 20  # Active creator
            and (
                get_numeric("creation_consistency", 0) > 0.5  # Regular creation
                or get_numeric("creation_burst_score", 0) > 15
            )  # Or burst creator
            and get_numeric("days_active", 0) > 2  # Active multiple days
        ):
            cluster_mapping[cluster_idx] = "creator"

        # Pro serious users - professional/focused usage
        elif (
            get_numeric("pro_user_ratio", 0) > 0.5  # At least 50% pro
            and (
                get_numeric("task_specialization_score", 0) > 0.3  # Some task focus
                or get_numeric("total_clips_created", 0) > 30
            )  # Or high volume
            and get_numeric("deletion_rate", 0) < 0.3  # Doesn't delete most content
        ):
            cluster_mapping[cluster_idx] = "pro_serious"

        # Casual users - occasional, often lyrics-focused
        elif (
            get_numeric("has_lyrics_generation", 0) > 30  # Uses lyrics
            or get_numeric("lyrics_generation_ratio", 0) > 0.2  # 20%+ lyrics
            or (
                get_numeric("total_clips_created", 0) < 20  # Low volume
                and get_numeric("days_active", 0) < 5
            )  # Infrequent
        ):
            cluster_mapping[cluster_idx] = "casual"

        # Default to other for unclassified patterns
        else:
            cluster_mapping[cluster_idx] = "other"

    # Log mapping with more details
    for cluster, user_type in cluster_mapping.items():
        size = cluster_profiles.loc[cluster, "cluster_size"]
        size_pct = cluster_profiles.loc[cluster, "cluster_size_pct"]
        logger.info(
            f"Cluster {cluster} ({size} users, {size_pct:.1f}%) mapped to: {user_type}"
        )

    return cluster_mapping


def save_all_outputs(
    assignments: pd.DataFrame,
    profiles: Tuple[pd.DataFrame, Dict[str, List[str]]],
    cluster_mapping: Dict[int, str],
    kmeans_model: Any,
    user_features: pd.DataFrame,
    output_dir: str,
) -> None:
    """
    Save all clustering outputs to files.

    Args:
        assignments: User cluster assignments
        profiles: Tuple of (cluster profiles DataFrame, top features dict)
        cluster_mapping: Mapping of clusters to user types
        kmeans_model: Trained k-means model
        user_features: Original user features
        output_dir: Directory to save outputs
    """
    logger.info(f"Saving outputs to {output_dir}")

    output_path = Path(output_dir)
    output_path.mkdir(parents=True, exist_ok=True)

    # Unpack profiles tuple
    profiles_df, top_features = profiles

    # Add user type to assignments
    assignments["user_type"] = assignments["cluster_label"].map(cluster_mapping)

    # Save user assignments
    assignments_file = output_path / "user_cluster_assignments.csv"
    assignments.to_csv(assignments_file, index=False)
    logger.info(f"Saved user assignments to {assignments_file}")

    # Save cluster profiles
    profiles_file = output_path / "cluster_profiles.csv"
    profiles_df.to_csv(profiles_file)
    logger.info(f"Saved cluster profiles to {profiles_file}")

    # Save top features per cluster
    top_features_file = output_path / "cluster_top_features.json"
    import json

    with open(top_features_file, "w") as f:
        json.dump(top_features, f, indent=2)
    logger.info(f"Saved top features to {top_features_file}")

    # Save cluster mapping
    mapping_file = output_path / "cluster_type_mapping.json"
    with open(mapping_file, "w") as f:
        json.dump(cluster_mapping, f, indent=2)
    logger.info(f"Saved cluster mapping to {mapping_file}")

    # Save k-means model
    model_file = output_path / "kmeans_model.pkl"
    joblib.dump(kmeans_model, model_file)
    logger.info(f"Saved k-means model to {model_file}")

    # Save feature matrix (for future predictions)
    features_file = output_path / "user_features.parquet"
    user_features.to_parquet(features_file, index=False)
    logger.info(f"Saved user features to {features_file}")

    # Create summary report
    create_summary_report(assignments, profiles_df, cluster_mapping, output_path)

    logger.info("All outputs saved successfully")


def create_summary_report(
    assignments: pd.DataFrame,
    profiles: pd.DataFrame,
    cluster_mapping: Dict[int, str],
    output_path: Path,
) -> None:
    """
    Create a human-readable summary report of clustering results.

    Args:
        assignments: User cluster assignments
        profiles: Cluster profiles
        cluster_mapping: Mapping of clusters to user types
        output_path: Path to save report
    """
    logger.info("Creating summary report")

    report_lines = []
    report_lines.append("USER CLUSTERING SUMMARY REPORT")
    report_lines.append("=" * 50)
    report_lines.append("")

    # Overall statistics
    report_lines.append("Overall Statistics:")
    report_lines.append(f"- Total users clustered: {len(assignments):,}")
    report_lines.append(f"- Number of clusters: {len(profiles)}")
    report_lines.append("")

    # Cluster breakdown
    report_lines.append("Cluster Breakdown:")
    for cluster_idx in sorted(profiles.index):
        cluster_type = cluster_mapping.get(cluster_idx, "unknown")
        size = profiles.loc[cluster_idx, "cluster_size"]
        size_pct = profiles.loc[cluster_idx, "cluster_size_pct"]

        report_lines.append(f"\nCluster {cluster_idx} ({cluster_type}):")
        report_lines.append(f"  - Size: {size:,} users ({size_pct:.1f}%)")

        # Key characteristics
        report_lines.append("  - Key characteristics:")

        # Select key features based on cluster type
        if cluster_type == "bot":
            key_features = [
                "consumption_ratio",
                "creation_burst_score",
                "min_time_between_creations",
                "has_any_reactions",
            ]
        elif cluster_type == "casual":
            key_features = [
                "has_lyrics_generation",
                "lyrics_generation_ratio",
                "total_clips_created",
                "pro_user_ratio",
            ]
        elif cluster_type == "pro_serious":
            key_features = [
                "pro_user_ratio",
                "model_diversity",
                "engagement_score",
                "platform_diversity",
            ]
        else:
            key_features = [
                "total_clips_created",
                "creation_frequency",
                "engagement_score",
                "model_diversity",
            ]

        for feature in key_features:
            if feature in profiles.columns:
                value = profiles.loc[cluster_idx, feature]
                if isinstance(value, (int, float)):
                    if feature.endswith("_ratio") or feature.endswith("_pct"):
                        report_lines.append(f"    * {feature}: {value:.1%}")
                    else:
                        report_lines.append(f"    * {feature}: {value:.2f}")
                else:
                    report_lines.append(f"    * {feature}: {value}")

    # User type distribution
    report_lines.append("\n" + "=" * 50)
    report_lines.append("User Type Distribution:")
    type_counts = assignments["user_type"].value_counts()
    for user_type, count in type_counts.items():
        pct = count / len(assignments) * 100
        report_lines.append(f"- {user_type}: {count:,} users ({pct:.1f}%)")

    # Save report
    report_file = output_path / "clustering_summary_report.txt"
    with open(report_file, "w") as f:
        f.write("\n".join(report_lines))

    logger.info(f"Saved summary report to {report_file}")
