"""
User Clustering Pipeline - Main Orchestration
============================================
This module orchestrates the entire user clustering pipeline based on creation patterns.
"""

import pandas as pd
import numpy as np
from pathlib import Path
import logging
import json
from typing import Dict, Tuple, Optional, List

# Import our custom modules
from data_loader import load_all_dataframes, validate_data
from feature_engineering import (
    create_temporal_features,
    create_bot_detection_features,
    create_content_quality_features,
    create_engagement_features,
    create_creation_source_features,
    create_model_usage_features,
    create_task_features,
    create_platform_features,
    create_prompt_features,
    create_playlist_features,
    aggregate_all_features,
)
from clustering import (
    prepare_features_for_clustering,
    find_optimal_clusters,
    perform_clustering,
    evaluate_cluster_stability,
    plot_cluster_visualization,
)
from output_generator import (
    create_cluster_assignments,
    create_cluster_profiles,
    map_clusters_to_types,
    save_all_outputs,
)
from utils import setup_logging, parse_metadata_field

# Configure logging
logger = logging.getLogger(__name__)


def load_data(input_dir: str) -> Dict[str, pd.DataFrame]:
    """Load data from input directory."""
    logger.info(f"Loading data from {input_dir}")

    # Try pickle files first (for real data)
    files_to_load_pkl = {
        "boosts_action_df": "bots_action.pkl",
        "reaction_df": "reaction.pkl",
        "total_clip_df": "total_clip.pkl",
        "playlist_clip_df": "playlist_clip.pkl",
    }

    # Try loading pickle files
    dfs = load_all_dataframes(input_dir, files_to_load_pkl)

    # If pickle files not found, try CSV files (for test data)
    if len(dfs) < 4:
        files_to_load_csv = {
            "boosts_action_df": "boosts_action_df.csv",
            "reaction_df": "reaction_df.csv",
            "total_clip_df": "total_clip_df.csv",
            "playlist_clip_df": "playlist_clip_df.csv",
        }

        # Try loading CSV files
        dfs_csv = load_all_dataframes(input_dir, files_to_load_csv)

        # Merge with any pickle files that were found
        for key, df in dfs_csv.items():
            if key not in dfs:
                dfs[key] = df

    validate_data(dfs)

    return dfs


def preprocess_data(dfs: Dict[str, pd.DataFrame]) -> Dict[str, pd.DataFrame]:
    """
    Perform data type conversions and preprocessing.

    Args:
        dfs: Dictionary of raw dataframes

    Returns:
        Dictionary of preprocessed dataframes
    """
    logger.info("Preprocessing data")

    # Convert timestamps
    timestamp_columns = {
        "boosts_action_df": ["created_at", "updated_at", "first_published_at"],
        "reaction_df": ["updated_at"],
        "total_clip_df": ["created_at", "updated_at"],
        "playlist_clip_df": ["updated_at"],
    }

    for df_name, columns in timestamp_columns.items():
        if df_name in dfs:
            for col in columns:
                if col in dfs[df_name].columns:
                    # Handle mixed timezone-aware and timezone-naive timestamps
                    dfs[df_name][col] = pd.to_datetime(
                        dfs[df_name][col], errors="coerce", utc=True
                    ).dt.tz_localize(None)  # Convert to UTC then remove timezone

    # Note: Metadata parsing is now handled directly in feature_engineering.py
    # via the parse_metadata_field function

    return dfs


def engineer_features(
    total_clip_df: pd.DataFrame,
    boosts_df: pd.DataFrame,
    reaction_df: pd.DataFrame,
    playlist_df: pd.DataFrame,
) -> pd.DataFrame:
    """
    Create all user-level features from the input dataframes.

    Args:
        total_clip_df: Main clips dataframe
        boosts_df: Boosts/actions dataframe
        reaction_df: User reactions dataframe
        playlist_df: Playlist clips dataframe

    Returns:
        DataFrame with user_id and all engineered features
    """
    logger.info("Starting feature engineering")

    # Get unique users
    users = total_clip_df["user_id"].unique()
    logger.info(f"Found {len(users)} unique users")

    # Create feature groups
    feature_dfs = []

    # 1. Temporal features
    temporal_features = create_temporal_features(total_clip_df)
    feature_dfs.append(temporal_features)

    # 2. Bot detection features
    bot_features = create_bot_detection_features(total_clip_df, reaction_df)
    feature_dfs.append(bot_features)

    # 3. Content quality features
    quality_features = create_content_quality_features(total_clip_df)
    feature_dfs.append(quality_features)

    # 4. Engagement features
    engagement_features = create_engagement_features(
        total_clip_df, boosts_df, reaction_df
    )
    feature_dfs.append(engagement_features)

    # 5. Creation source features
    source_features = create_creation_source_features(total_clip_df)
    feature_dfs.append(source_features)

    # 6. Model usage features
    model_features = create_model_usage_features(total_clip_df)
    feature_dfs.append(model_features)

    # 7. Task features
    task_features = create_task_features(total_clip_df)
    feature_dfs.append(task_features)

    # 8. Platform features
    platform_features = create_platform_features(total_clip_df)
    feature_dfs.append(platform_features)

    # 9. Prompt features
    prompt_features = create_prompt_features(total_clip_df, boosts_df)
    feature_dfs.append(prompt_features)

    # 10. Playlist features
    playlist_features = create_playlist_features(total_clip_df, playlist_df)
    feature_dfs.append(playlist_features)

    # Aggregate all features
    user_features = aggregate_all_features(feature_dfs, users)

    logger.info(
        f"Created {len(user_features.columns) - 1} features for {len(user_features)} users"
    )

    return user_features


def generate_outputs(
    user_features: pd.DataFrame,
    cluster_labels: np.ndarray,
    kmeans_model: object,
    output_dir: str,
    X_scaled: Optional[np.ndarray] = None,
) -> None:
    """
    Create and save all output files.

    Args:
        user_features: Original user features dataframe
        cluster_labels: Cluster assignments
        kmeans_model: Trained k-means model
        output_dir: Directory to save outputs
        X_scaled: Scaled feature matrix for visualization (optional)
    """
    logger.info("Generating outputs")

    # Create cluster assignments
    assignments = create_cluster_assignments(
        user_features["user_id"], cluster_labels, kmeans_model
    )

    # Create cluster profiles (returns tuple of profiles and top features)
    profiles_data = create_cluster_profiles(user_features, cluster_labels)

    # Map clusters to expected types
    cluster_mapping = map_clusters_to_types(profiles_data[0])  # Pass the DataFrame part

    # Save all outputs
    save_all_outputs(
        assignments,
        profiles_data,
        cluster_mapping,
        kmeans_model,
        user_features,
        output_dir,
    )

    # Create visualization if X_scaled is provided
    if X_scaled is not None:
        plot_cluster_visualization(X_scaled, cluster_labels, output_dir)

    logger.info(f"All outputs saved to {output_dir}")


def main(input_dir: str, output_dir: str, sample_size: Optional[int] = None):
    """
    Orchestrate the full user clustering pipeline.

    Args:
        input_dir: Directory containing input CSV files
        output_dir: Directory to save outputs
        sample_size: Optional sample size for memory efficiency
    """
    # Create output directory
    Path(output_dir).mkdir(parents=True, exist_ok=True)

    # Setup logging with output directory
    setup_logging(output_dir=output_dir)
    logger.info("Starting user clustering pipeline")

    try:
        # Step 1: Load data
        dfs = load_data(input_dir)

        # Step 2: Preprocess data
        dfs = preprocess_data(dfs)

        # Step 3: Engineer features
        user_features = engineer_features(
            dfs["total_clip_df"],
            dfs["boosts_action_df"],
            dfs["reaction_df"],
            dfs["playlist_clip_df"],
        )

        # Step 4: Prepare features for clustering
        sampled_df, X_scaled, scaler = prepare_features_for_clustering(
            user_features, sample_size
        )

        # Step 5: Find optimal number of clusters
        min_cluster_size_ratio = 0.0025  # Reduced to 0.25% (50 users for 20k sample) to allow more specialized groups
        # This allows us to capture more nuanced user segments while still avoiding tiny clusters
        optimal_k = find_optimal_clusters(
            X_scaled,
            k_range=(5, 15),  # Further expanded range to explore more clusters
            output_dir=output_dir,
            min_cluster_size_ratio=min_cluster_size_ratio,
        )
        logger.info(f"Optimal number of clusters: {optimal_k}")

        # Step 6: Perform clustering
        kmeans_model, cluster_labels = perform_clustering(
            X_scaled, optimal_k, min_cluster_size_ratio=min_cluster_size_ratio
        )

        # Step 7: Evaluate stability
        stability_score = evaluate_cluster_stability(X_scaled, optimal_k)
        logger.info(f"Cluster stability score: {stability_score:.3f}")

        # Step 8: Generate outputs
        # If we sampled, we need to align the features back
        if sample_size and len(sampled_df) < len(user_features):
            output_features = sampled_df
        else:
            output_features = user_features

        generate_outputs(
            output_features, cluster_labels, kmeans_model, output_dir, X_scaled
        )

        logger.info("Pipeline completed successfully!")

    except Exception as e:
        logger.error(f"Pipeline failed: {str(e)}")
        raise


if __name__ == "__main__":
    import argparse

    parser = argparse.ArgumentParser(description="User Clustering Pipeline")
    parser.add_argument(
        "--input-dir", required=True, help="Input directory with CSV files"
    )
    parser.add_argument(
        "--output-dir", required=True, help="Output directory for results"
    )
    parser.add_argument(
        "--sample-size", type=int, help="Optional sample size for memory efficiency"
    )

    args = parser.parse_args()

    main(args.input_dir, args.output_dir, args.sample_size)
