#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Analisi psicometrica per il Test dei Miti.

Requisiti:
    pip install -r requirements.txt

Uso:
    python analisi_psicometrica.py --input ../data/dati_raccolti.csv --output ../report/

Il file CSV deve contenere:
- colonne anagrafiche: eta, genere, istruzione, terapia_in_corso, diagnosi
- colonne item: <mito>_1, <mito>_2 per ognuno dei 25 miti
- opzionale: colonne test_retest_<mito> per l'analisi test-retest
- opzionale: colonne di altri strumenti (es. bdi_total, stai_total) per validità convergente
"""

import argparse
import os
import json
import warnings
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
from scipy import stats
from sklearn.decomposition import FactorAnalysis
from factor_analyzer import calculate_kmo

warnings.filterwarnings('ignore')


MITI = [
    "adamo", "satan", "crono", "edipo", "sacrificio",
    "mostro", "cannibale", "zombie", "eroe", "icaro",
    "ordalia", "gige", "distruttore", "rivoluzionario", "stregone",
    "sacerdote", "medusa", "profeta", "re", "guardiano",
    "alchimista", "selvaggio", "narciso", "folle", "apocalisse"
]


def cronbach_alpha(items_df):
    """Calcola l'alfa di Cronbach per una scala."""
    items = items_df.values
    n_items = items.shape[1]
    if n_items < 2:
        return np.nan
    variances = items.var(axis=0, ddof=1)
    total_variance = items.sum(axis=1).var(ddof=1)
    if total_variance == 0:
        return np.nan
    return (n_items / (n_items - 1)) * (1 - variances.sum() / total_variance)


def scale_scores(df):
    """Calcola i punteggi standardizzati 0-100 per ogni mito."""
    scores = {}
    for mito in MITI:
        cols = [f"{mito}_1", f"{mito}_2"]
        if all(c in df.columns for c in cols):
            raw = df[cols].sum(axis=1)
            scores[mito] = (raw / 10) * 100
    return pd.DataFrame(scores)


def descriptive_stats(df, output_dir):
    """Statistiche descrittive per ogni item e per ogni scala."""
    report = []
    report.append("=" * 70)
    report.append("STATISTICHE DESCRITTIVE")
    report.append("=" * 70)

    # Per scala
    scores = scale_scores(df)
    desc = scores.describe().T
    desc['median'] = scores.median()
    desc = desc[['count', 'mean', 'std', 'median', 'min', 'max']]
    report.append("\nPunteggi per scala (0-100):\n")
    report.append(desc.round(2).to_string())

    # Per item
    item_cols = [f"{m}_{i}" for m in MITI for i in [1, 2] if f"{m}_{i}" in df.columns]
    item_desc = df[item_cols].describe().T[['mean', 'std', 'min', 'max']]
    report.append("\n\nStatistiche per item:\n")
    report.append(item_desc.round(2).to_string())

    # Salvataggio
    desc.to_csv(os.path.join(output_dir, 'statistiche_scala.csv'))
    item_desc.to_csv(os.path.join(output_dir, 'statistiche_item.csv'))

    # Grafico distribuzione punteggi
    plt.figure(figsize=(14, 8))
    scores.boxplot(rot=45)
    plt.title('Distribuzione dei punteggi per scala')
    plt.ylabel('Punteggio (0-100)')
    plt.tight_layout()
    plt.savefig(os.path.join(output_dir, 'distribuzione_scala.png'), dpi=150)
    plt.close()

    return "\n".join(report), scores


def reliability_analysis(df, output_dir):
    """Calcola l'affidabilità interna per ogni scala."""
    report = []
    report.append("\n" + "=" * 70)
    report.append("AFFIDABILITÀ INTERNA (Alfa di Cronbach)")
    report.append("=" * 70)

    alphas = {}
    for mito in MITI:
        cols = [f"{mito}_1", f"{mito}_2"]
        if all(c in df.columns for c in cols):
            alpha = cronbach_alpha(df[cols])
            alphas[mito] = alpha
            status = "OK" if alpha >= 0.7 else ("accettabile" if alpha >= 0.6 else "debole")
            report.append(f"{mito:15} α = {alpha:.3f}  ({status})")

    alpha_df = pd.DataFrame(list(alphas.items()), columns=['scala', 'cronbach_alpha'])
    alpha_df.to_csv(os.path.join(output_dir, 'cronbach_alpha.csv'), index=False)

    # Grafico
    plt.figure(figsize=(12, 6))
    colors = ['#5b9a6d' if a >= 0.7 else ('#c9a227' if a >= 0.6 else '#c94f4f') for a in alphas.values()]
    plt.barh(list(alphas.keys()), list(alphas.values()), color=colors)
    plt.axvline(0.7, color='white', linestyle='--', alpha=0.7, label='soglia accettabile (0.70)')
    plt.xlabel("Alfa di Cronbach")
    plt.title("Affidabilità interna per scala")
    plt.legend()
    plt.tight_layout()
    plt.savefig(os.path.join(output_dir, 'cronbach_alpha.png'), dpi=150)
    plt.close()

    return "\n".join(report)


def factor_analysis(df, output_dir):
    """Analisi fattoriale esploratoria su tutti gli item."""
    report = []
    report.append("\n" + "=" * 70)
    report.append("ANALISI FATTORIALE ESPLORATORIA")
    report.append("=" * 70)

    item_cols = [f"{m}_{i}" for m in MITI for i in [1, 2] if f"{m}_{i}" in df.columns]
    X = df[item_cols].dropna()

    if len(X) < 100:
        report.append("\nCampione troppo piccolo per un'analisi fattoriale affidabile (N < 100).")
        return "\n".join(report)

    # KMO
    try:
        kmo_all, kmo_model = calculate_kmo(X)
        report.append(f"\nKMO complessivo: {kmo_model:.3f}")
        if kmo_model < 0.6:
            report.append("Attenzione: KMO inferiore a 0.6. La struttura fattoriale potrebbe essere debole.")
    except Exception as e:
        report.append(f"\nImpossibile calcolare KMO: {e}")

    # Factor Analysis con 25 fattori (modello teorico)
    n_factors = min(25, len(item_cols) - 1)
    fa = FactorAnalysis(n_components=n_factors, random_state=42, max_iter=1000)
    fa.fit(X)

    # Varianza spiegata
    variance = fa.noise_variance_
    explained = 1 - variance
    report.append(f"\nVarianza spiegata dai {n_factors} fattori: {explained.sum():.2f}")
    report.append(f"Varianza spiegata media per item: {explained.mean():.3f}")

    # Loadings
    loadings = pd.DataFrame(
        fa.components_.T,
        index=item_cols,
        columns=[f"Fattore {i+1}" for i in range(n_factors)]
    )
    loadings.to_csv(os.path.join(output_dir, 'factor_loadings.csv'))

    # Heatmap loadings (primi 10 fattori)
    plt.figure(figsize=(14, 20))
    sns.heatmap(loadings.iloc[:, :10], cmap='RdBu_r', center=0, vmin=-1, vmax=1)
    plt.title('Loadings fattoriali (primi 10 fattori)')
    plt.tight_layout()
    plt.savefig(os.path.join(output_dir, 'factor_loadings.png'), dpi=150)
    plt.close()

    return "\n".join(report)


def test_retest_analysis(df, output_dir):
    """Analisi test-retest se presenti colonne retest."""
    report = []
    report.append("\n" + "=" * 70)
    report.append("AFFIDABILITÀ TEST-RETEST")
    report.append("=" * 70)

    retest_cols = [c for c in df.columns if c.startswith('retest_')]
    if not retest_cols:
        report.append("\nNessuna colonna 'retest_<mito>' trovata. Salto l'analisi test-retest.")
        return "\n".join(report)

    scores_t1 = scale_scores(df)
    scores_t2 = pd.DataFrame()
    for c in retest_cols:
        mito = c.replace('retest_', '')
        if mito in MITI:
            scores_t2[mito] = df[c]

    common = [m for m in MITI if m in scores_t1.columns and m in scores_t2.columns]
    correlations = {}
    for m in common:
        mask = scores_t1[m].notna() & scores_t2[m].notna()
        if mask.sum() < 5:
            continue
        r, p = stats.pearsonr(scores_t1.loc[mask, m], scores_t2.loc[mask, m])
        correlations[m] = {'r': r, 'p': p, 'n': mask.sum()}
        report.append(f"{m:15} r = {r:.3f}, p = {p:.3f}, N = {mask.sum()}")

    if correlations:
        retest_df = pd.DataFrame(correlations).T
        retest_df.to_csv(os.path.join(output_dir, 'test_retest.csv'))

    return "\n".join(report)


def convergent_validity(df, output_dir):
    """Correlazioni con altri strumenti, se presenti."""
    report = []
    report.append("\n" + "=" * 70)
    report.append("VALIDITÀ CONVERGENTE E DISCRIMINANTE")
    report.append("=" * 70)

    external_cols = [c for c in df.columns if c.startswith('ext_')]
    if not external_cols:
        report.append("\nNessuna colonna esterna 'ext_<nome_test>' trovata. Salto l'analisi.")
        report.append("Per usarla, aggiungi colonne come 'ext_bdi_total', 'ext_stai_total', ecc.")
        return "\n".join(report)

    scores = scale_scores(df)
    corr_results = []
    for ext in external_cols:
        for mito in MITI:
            if mito not in scores.columns:
                continue
            mask = df[ext].notna() & scores[mito].notna()
            if mask.sum() < 5:
                continue
            r, p = stats.pearsonr(df.loc[mask, ext], scores.loc[mask, mito])
            corr_results.append({
                'strumento_esterno': ext,
                'mito': mito,
                'r': r,
                'p': p,
                'n': mask.sum()
            })

    corr_df = pd.DataFrame(corr_results)
    corr_df.to_csv(os.path.join(output_dir, 'correlazioni_esterne.csv'), index=False)
    report.append(f"\nCorrelazioni calcolate: {len(corr_df)}")
    report.append("\nPrime 10 correlazioni significative (p < 0.05):\n")
    sig = corr_df[corr_df['p'] < 0.05].sort_values('p').head(10)
    if len(sig) > 0:
        report.append(sig.round(3).to_string(index=False))
    else:
        report.append("Nessuna correlazione significativa trovata.")

    # Heatmap
    if len(external_cols) > 0:
        pivot = corr_df.pivot(index='mito', columns='strumento_esterno', values='r')
        plt.figure(figsize=(10, 14))
        sns.heatmap(pivot, cmap='RdBu_r', center=0, vmin=-1, vmax=1, annot=True, fmt='.2f')
        plt.title('Correlazioni tra miti e strumenti esterni')
        plt.tight_layout()
        plt.savefig(os.path.join(output_dir, 'correlazioni_esterne.png'), dpi=150)
        plt.close()

    return "\n".join(report)


def group_comparison(df, output_dir):
    """Confronto tra gruppi clinici e non clinici, se presente la colonna 'gruppo'."""
    report = []
    report.append("\n" + "=" * 70)
    report.append("CONFRONTO TRA GRUPPI")
    report.append("=" * 70)

    if 'gruppo' not in df.columns:
        report.append("\nColonna 'gruppo' non trovata. Salto il confronto.")
        report.append("Usa 'controllo', 'terapia', 'diagnosi' o altre etichette per confrontare gruppi.")
        return "\n".join(report)

    scores = scale_scores(df)
    scores['gruppo'] = df['gruppo']
    groups = scores.groupby('gruppo')[MITI].mean()
    groups.to_csv(os.path.join(output_dir, 'media_per_gruppo.csv'))

    report.append("\nMedia per gruppo:\n")
    report.append(groups.round(2).to_string())

    # Test t tra i due gruppi più numerosi
    group_names = scores['gruppo'].value_counts().index[:2]
    if len(group_names) == 2:
        g1, g2 = group_names
        report.append(f"\n\nConfronto {g1} vs {g2} (test t di Student):\n")
        for mito in MITI:
            if mito not in scores.columns:
                continue
            a = scores[scores['gruppo'] == g1][mito].dropna()
            b = scores[scores['gruppo'] == g2][mito].dropna()
            if len(a) < 3 or len(b) < 3:
                continue
            t, p = stats.ttest_ind(a, b)
            report.append(f"{mito:15} t = {t:7.3f}, p = {p:.3f}")

    return "\n".join(report)


def correlation_matrix(scores, output_dir):
    """Matrice di correlazione tra le scale."""
    corr = scores.corr()
    corr.to_csv(os.path.join(output_dir, 'matrice_correlazioni_scala.csv'))

    plt.figure(figsize=(14, 12))
    sns.heatmap(corr, cmap='RdBu_r', center=0, vmin=-1, vmax=1, annot=True, fmt='.2f')
    plt.title('Matrice di correlazione tra le scale mitiche')
    plt.tight_layout()
    plt.savefig(os.path.join(output_dir, 'matrice_correlazioni.png'), dpi=150)
    plt.close()


def main():
    parser = argparse.ArgumentParser(description='Analisi psicometrica Test dei Miti')
    parser.add_argument('--input', required=True, help='Percorso al CSV con i dati raccolti')
    parser.add_argument('--output', default='report', help='Cartella di output per i risultati')
    args = parser.parse_args()

    os.makedirs(args.output, exist_ok=True)

    df = pd.read_csv(args.input)
    print(f"Dataset caricato: {len(df)} partecipanti, {len(df.columns)} colonne")

    report_parts = []
    report_parts.append(f"ANALISI PSICOMETRICA - TEST DEI MITI")
    report_parts.append(f"Partecipanti: {len(df)}")
    report_parts.append(f"Item analizzati: 50 (2 per mito)")
    report_parts.append(f"Miti: {len(MITI)}")

    desc_report, scores = descriptive_stats(df, args.output)
    report_parts.append(desc_report)
    report_parts.append(reliability_analysis(df, args.output))
    report_parts.append(factor_analysis(df, args.output))
    report_parts.append(test_retest_analysis(df, args.output))
    report_parts.append(convergent_validity(df, args.output))
    report_parts.append(group_comparison(df, args.output))

    correlation_matrix(scores, args.output)

    final_report = "\n".join(report_parts)
    with open(os.path.join(args.output, 'report_psicometrico.txt'), 'w', encoding='utf-8') as f:
        f.write(final_report)

    print(final_report)
    print(f"\nReport salvato in: {args.output}")


if __name__ == '__main__':
    main()
