import pandas as pd
import numpy as np
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import seaborn as sns
import os
from sklearn.decomposition import PCA

def perform_eda(csv_path):
    output_dir = 'EDA_Reports/figures'
    os.makedirs(output_dir, exist_ok=True)

    print("Loading dataset...")
    df = pd.read_csv(csv_path)
    
    if 'key' in df.columns:
        df = df.drop(columns=['key'])

    print(f"Dataset shape after dropping 'key': {df.shape}")

    sns.set_theme(style="whitegrid", palette="muted")
    plt.rcParams.update({'font.size': 12})

    print("Generating Correlation Heatmap...")
    plt.figure(figsize=(20, 16))
    corr = df.corr()
    
    mask = np.triu(np.ones_like(corr, dtype=bool))
    
    sns.heatmap(corr, mask=mask, annot=False, cmap='coolwarm', center=0,
                square=True, linewidths=.5, cbar_kws={"shrink": .5})
    plt.title('Feature Correlation Heatmap', fontsize=20, pad=20)
    plt.tight_layout()
    plt.savefig(f'{output_dir}/01_correlation_heatmap.png', dpi=300)
    plt.close()

    print("Generating Feature Value Distribution...")
    val_counts = {}
    for col in df.columns:
        counts = df[col].value_counts(normalize=True) * 100
        val_counts[col] = {
            '-1 (Phishy/Legitimate)': counts.get(-1, 0),
            '0 (Suspicious)': counts.get(0, 0),
            '1 (Legitimate/Phishy)': counts.get(1, 0)
        }
    
    dist_df = pd.DataFrame(val_counts).T
    
    plt.figure(figsize=(15, 12))
    dist_df.plot(kind='barh', stacked=True, color=['#ff9999', '#ffcc99', '#99ff99'], figsize=(16, 12), width=0.8)
    plt.title('Distribution of Values (-1, 0, 1) Across All Features', fontsize=18)
    plt.xlabel('Percentage (%)', fontsize=14)
    plt.ylabel('Features', fontsize=14)
    plt.legend(title='Category Value', bbox_to_anchor=(1.05, 1), loc='upper left')
    plt.tight_layout()
    plt.savefig(f'{output_dir}/02_value_distribution.png', dpi=300)
    plt.close()

    print("Generating PCA Visualization...")
    pca = PCA(n_components=2)
    pca_result = pca.fit_transform(df)
    
    plt.figure(figsize=(10, 8))
    plt.scatter(pca_result[:, 0], pca_result[:, 1], alpha=0.5, c='#3498db', edgecolors='w', s=50)
    plt.title('PCA: 2D Projection of the Dataset', fontsize=16)
    plt.xlabel(f'Principal Component 1 ({pca.explained_variance_ratio_[0]*100:.2f}%)', fontsize=12)
    plt.ylabel(f'Principal Component 2 ({pca.explained_variance_ratio_[1]*100:.2f}%)', fontsize=12)
    plt.tight_layout()
    plt.savefig(f'{output_dir}/03_pca_projection.png', dpi=300)
    plt.close()

    print("Generating Feature Variance Plot...")
    variances = df.var().sort_values(ascending=False).head(15)
    
    plt.figure(figsize=(12, 8))
    sns.barplot(x=variances.values, y=variances.index, hue=variances.index, palette='viridis', dodge=False)
    plt.title('Top 15 Features with Highest Variance', fontsize=16)
    plt.xlabel('Variance', fontsize=12)
    plt.ylabel('Features', fontsize=12)
    plt.tight_layout()
    plt.savefig(f'{output_dir}/04_top_feature_variances.png', dpi=300)
    plt.close()

    print("Generating Missing Values Plot...")
    missing_counts = df.isnull().sum()
    plt.figure(figsize=(12, 10))
    sns.barplot(x=missing_counts.values, y=missing_counts.index, palette='Reds_r', hue=missing_counts.index, legend=False)
    plt.title('Missing Values per Feature (Count = 0 for all)', fontsize=16)
    plt.xlabel('Number of Missing Values', fontsize=12)
    plt.ylabel('Features', fontsize=12)
    plt.tight_layout()
    plt.savefig(f'{output_dir}/05_missing_values.png', dpi=300)
    plt.close()

    print("Generating Class Distribution Plot...")
    target = 'Result' if 'Result' in df.columns else 'SSLfinal_State'
    
    plt.figure(figsize=(8, 6))
    ax = sns.countplot(x=target, data=df, palette='Set2', hue=target, legend=False)
    title_text = f'Distribution of {target}' + (' (Proxy for Class)' if target != 'Result' else '')
    plt.title(title_text, fontsize=16)
    plt.xlabel('Category (-1: Phishing, 0: Suspicious, 1: Legitimate)', fontsize=12)
    plt.ylabel('Count', fontsize=12)
    for p in ax.patches:
        ax.annotate(f'{int(p.get_height())}', (p.get_x() + p.get_width() / 2., p.get_height()),
                    ha='center', va='center', xytext=(0, 5), textcoords='offset points')
    plt.tight_layout()
    plt.savefig(f'{output_dir}/06_class_distribution.png', dpi=300)
    plt.close()

    print(f"EDA successfully completed. 6 figures saved in '{output_dir}'.")

if __name__ == '__main__':
    perform_eda('Phising_Testing_Dataset.csv')
