#!/usr/bin/env python3
"""
Step 2: Data Cleaning and Exploratory Data Analysis (EDA)
GLP-1 Receptor Agonists Adverse Events Study

Objectives:
- Clean raw adverse event data from OpenFDA FAERS
- Remove duplicate records based on safetyreportid
- Flatten nested reactions for analysis
- Perform EDA: temporal trends, top adverse events, seriousness distribution
- Check for specific signals of interest (gastroparesis, thyroid neoplasm)
"""

import json
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import matplotlib.dates as mdates
import seaborn as sns
from pathlib import Path
from datetime import datetime
from collections import Counter
import warnings

warnings.filterwarnings('ignore')

# Set up paths
SESSION_DIR = Path("/app/sandbox/session_20260203_092044_b80cb683d2ee")
RAW_DATA_PATH = SESSION_DIR / "workflow/data/glp1_adverse_events.json"
CLEANED_DATA_PATH = SESSION_DIR / "workflow/data/glp1_cleaned.parquet"
FIGURES_DIR = SESSION_DIR / "figures"
RESULTS_DIR = SESSION_DIR / "results"

# Ensure directories exist
FIGURES_DIR.mkdir(exist_ok=True)
RESULTS_DIR.mkdir(exist_ok=True)

# Set plotting style
plt.rcParams['font.family'] = 'sans-serif'
plt.rcParams['font.size'] = 10
plt.rcParams['axes.linewidth'] = 0.5
plt.rcParams['figure.facecolor'] = 'white'
sns.set_palette("husl")

def load_raw_data():
    """Load raw JSON data with progress updates."""
    print(f"\n{'='*60}")
    print("STEP 1: Loading raw data")
    print(f"{'='*60}")
    print(f"Reading from: {RAW_DATA_PATH}")
    print(f"File size: {RAW_DATA_PATH.stat().st_size / (1024**3):.2f} GB")

    start_time = datetime.now()
    with open(RAW_DATA_PATH, 'r') as f:
        data = json.load(f)

    elapsed = (datetime.now() - start_time).total_seconds()
    print(f"Loaded {len(data)} records in {elapsed:.1f} seconds")
    return data


def deduplicate_records(data):
    """Remove duplicate records based on safetyreportid."""
    print(f"\n{'='*60}")
    print("STEP 2: Deduplication")
    print(f"{'='*60}")

    original_count = len(data)

    # Track unique reports
    seen_ids = set()
    unique_records = []
    duplicate_count = 0

    for i, record in enumerate(data):
        if i % 5000 == 0:
            print(f"  Processing record {i}/{original_count}...")

        report_id = record.get('safetyreportid')
        if report_id and report_id not in seen_ids:
            seen_ids.add(report_id)
            unique_records.append(record)
        else:
            duplicate_count += 1

    print(f"\n  Original records: {original_count:,}")
    print(f"  Unique records: {len(unique_records):,}")
    print(f"  Duplicates removed: {duplicate_count:,}")

    return unique_records


def flatten_records(data):
    """Flatten nested record structure for analysis."""
    print(f"\n{'='*60}")
    print("STEP 3: Flattening records")
    print(f"{'='*60}")

    flattened = []
    total = len(data)

    for i, record in enumerate(data):
        if i % 5000 == 0:
            print(f"  Flattening record {i}/{total}...")

        base_record = {
            'safetyreportid': record.get('safetyreportid'),
            'receivedate': record.get('receivedate'),
            'serious': record.get('serious'),
            'seriousnessdeath': record.get('seriousnessdeath'),
            'seriousnesslifethreatening': record.get('seriousnesslifethreatening'),
            'seriousnesshospitalization': record.get('seriousnesshospitalization'),
            'seriousnessdisabling': record.get('seriousnessdisabling'),
            'patient_age': record.get('patient_age'),
            'patient_sex': record.get('patient_sex'),
            'query_drug': record.get('query_drug'),
        }

        # Flatten reactions - create one row per reaction
        reactions = record.get('reactions', [])
        if reactions and len(reactions) > 0:
            for reaction in reactions:
                row = base_record.copy()
                if isinstance(reaction, dict):
                    row['reaction_meddraPT'] = reaction.get('reactionmeddrapt')
                    row['reaction_outcome'] = reaction.get('reactionoutcome')
                else:
                    row['reaction_meddraPT'] = reaction
                    row['reaction_outcome'] = None
                flattened.append(row)
        else:
            # Keep record even without reactions
            base_record['reaction_meddraPT'] = None
            base_record['reaction_outcome'] = None
            flattened.append(base_record)

    print(f"  Created {len(flattened):,} flattened rows")
    return flattened


def create_dataframe(flattened_data):
    """Create and process pandas DataFrame."""
    print(f"\n{'='*60}")
    print("STEP 4: Creating DataFrame")
    print(f"{'='*60}")

    df = pd.DataFrame(flattened_data)
    print(f"  Shape: {df.shape}")
    print(f"  Columns: {list(df.columns)}")

    # Convert receivedate to datetime
    print("\n  Converting dates...")
    df['receivedate'] = pd.to_datetime(df['receivedate'], format='%Y%m%d', errors='coerce')

    # Create year-quarter field
    df['year'] = df['receivedate'].dt.year
    df['quarter'] = df['receivedate'].dt.quarter
    df['year_quarter'] = df['receivedate'].dt.to_period('Q').astype(str)

    # Clean serious field
    df['is_serious'] = df['serious'].apply(lambda x: True if str(x) == '1' else False)

    # Clean patient sex
    sex_map = {'1': 'Male', '2': 'Female', 1: 'Male', 2: 'Female'}
    df['patient_sex_label'] = df['patient_sex'].map(sex_map).fillna('Unknown')

    # Normalize reaction terms (uppercase for matching)
    df['reaction_normalized'] = df['reaction_meddraPT'].str.upper().str.strip()

    print(f"\n  Date range: {df['receivedate'].min()} to {df['receivedate'].max()}")
    print(f"  Years covered: {sorted(df['year'].dropna().unique().astype(int))}")

    return df


def analyze_temporal_trends(df):
    """Analyze and visualize reporting trends over time."""
    print(f"\n{'='*60}")
    print("STEP 5: Temporal Analysis")
    print(f"{'='*60}")

    # Aggregate unique reports by drug and quarter
    reports_by_quarter = df.groupby(['query_drug', 'year_quarter'])['safetyreportid'].nunique().reset_index()
    reports_by_quarter.columns = ['drug', 'year_quarter', 'report_count']

    # Pivot for plotting
    pivot_data = reports_by_quarter.pivot(index='year_quarter', columns='drug', values='report_count').fillna(0)

    print(f"  Quarters analyzed: {len(pivot_data)}")
    print(f"  Drugs: {list(pivot_data.columns)}")

    # Create visualization
    fig, ax = plt.subplots(figsize=(14, 7))

    colors = {'semaglutide': '#1f77b4', 'liraglutide': '#ff7f0e',
              'exenatide': '#2ca02c', 'dulaglutide': '#d62728', 'tirzepatide': '#9467bd'}

    for drug in pivot_data.columns:
        ax.plot(pivot_data.index, pivot_data[drug], marker='o', markersize=4,
                label=drug.capitalize(), linewidth=2, color=colors.get(drug, '#333333'))

    ax.set_xlabel('Year-Quarter', fontsize=12)
    ax.set_ylabel('Number of Adverse Event Reports', fontsize=12)
    ax.set_title('GLP-1 Receptor Agonist Adverse Event Reports Over Time\n(FAERS Database)', fontsize=14)
    ax.legend(title='Drug', loc='upper left')
    ax.grid(True, alpha=0.3)

    # Rotate x-axis labels
    plt.xticks(rotation=45, ha='right')

    # Only show every 4th label to avoid crowding
    labels = ax.get_xticklabels()
    for i, label in enumerate(labels):
        if i % 4 != 0:
            label.set_visible(False)

    plt.tight_layout()
    plt.savefig(FIGURES_DIR / 'reports_over_time.png', dpi=300, bbox_inches='tight')
    plt.close()

    print(f"  Saved: {FIGURES_DIR / 'reports_over_time.png'}")

    return pivot_data


def analyze_top_adverse_events(df):
    """Identify and visualize top 20 most frequent adverse events."""
    print(f"\n{'='*60}")
    print("STEP 6: Top Adverse Events Analysis")
    print(f"{'='*60}")

    # Count adverse events (excluding nulls)
    reaction_counts = df[df['reaction_meddraPT'].notna()]['reaction_meddraPT'].value_counts()
    top_20 = reaction_counts.head(20)

    print(f"\n  Top 20 Adverse Events:")
    for i, (reaction, count) in enumerate(top_20.items(), 1):
        print(f"    {i:2}. {reaction}: {count:,}")

    # Create visualization
    fig, ax = plt.subplots(figsize=(12, 10))

    bars = ax.barh(range(len(top_20)), top_20.values, color=sns.color_palette("Blues_r", len(top_20)))
    ax.set_yticks(range(len(top_20)))
    ax.set_yticklabels(top_20.index)
    ax.invert_yaxis()  # Top item first

    ax.set_xlabel('Number of Reports', fontsize=12)
    ax.set_title('Top 20 Most Frequently Reported Adverse Events\nGLP-1 Receptor Agonists (FAERS Database)', fontsize=14)

    # Add count labels
    for i, (count, bar) in enumerate(zip(top_20.values, bars)):
        ax.text(count + max(top_20.values)*0.01, i, f'{count:,}', va='center', fontsize=9)

    ax.set_xlim(0, max(top_20.values) * 1.15)
    plt.tight_layout()
    plt.savefig(FIGURES_DIR / 'top_adverse_events.png', dpi=300, bbox_inches='tight')
    plt.close()

    print(f"\n  Saved: {FIGURES_DIR / 'top_adverse_events.png'}")

    return top_20


def analyze_seriousness(df):
    """Analyze distribution of serious vs non-serious events."""
    print(f"\n{'='*60}")
    print("STEP 7: Seriousness Analysis")
    print(f"{'='*60}")

    # Get unique reports only
    unique_reports = df.drop_duplicates(subset=['safetyreportid'])

    serious_counts = unique_reports['is_serious'].value_counts()
    serious_pct = serious_counts / serious_counts.sum() * 100

    print(f"\n  Seriousness Distribution (unique reports):")
    print(f"    Serious: {serious_counts.get(True, 0):,} ({serious_pct.get(True, 0):.1f}%)")
    print(f"    Non-serious: {serious_counts.get(False, 0):,} ({serious_pct.get(False, 0):.1f}%)")

    # Breakdown by specific seriousness criteria
    print("\n  Seriousness Breakdown:")
    print(f"    Deaths: {unique_reports['seriousnessdeath'].eq('1').sum():,}")
    print(f"    Life-threatening: {unique_reports['seriousnesslifethreatening'].eq('1').sum():,}")
    print(f"    Hospitalization: {unique_reports['seriousnesshospitalization'].eq('1').sum():,}")
    print(f"    Disabling: {unique_reports['seriousnessdisabling'].eq('1').sum():,}")

    # Seriousness by drug
    serious_by_drug = unique_reports.groupby('query_drug')['is_serious'].agg(['sum', 'count'])
    serious_by_drug['pct_serious'] = serious_by_drug['sum'] / serious_by_drug['count'] * 100
    serious_by_drug = serious_by_drug.sort_values('pct_serious', ascending=False)

    print("\n  Seriousness by Drug:")
    for drug, row in serious_by_drug.iterrows():
        print(f"    {drug.capitalize()}: {row['sum']:.0f}/{row['count']:.0f} ({row['pct_serious']:.1f}% serious)")

    return serious_counts, serious_by_drug


def check_signals_of_interest(df):
    """Check for presence of gastroparesis and thyroid-related events."""
    print(f"\n{'='*60}")
    print("STEP 8: Signal Detection - Target Adverse Events")
    print(f"{'='*60}")

    # Define search terms
    gastroparesis_terms = ['GASTROPARESIS', 'GASTRIC PARALYSIS', 'DELAYED GASTRIC EMPTYING',
                           'STOMACH PARALYSIS', 'GASTROPATHY']
    thyroid_terms = ['THYROID', 'THYROID CANCER', 'THYROID NEOPLASM', 'THYROID CARCINOMA',
                     'MEDULLARY THYROID', 'PAPILLARY THYROID', 'GOITER', 'GOITRE',
                     'THYROID NODULE', 'THYROID MASS', 'THYROID DISORDER', 'HYPERTHYROID',
                     'HYPOTHYROID', 'THYROIDITIS', 'THYROID C-CELL']

    results = {}

    # Check gastroparesis
    print("\n  GASTROPARESIS AND RELATED:")
    gastro_mask = df['reaction_normalized'].str.contains('|'.join(gastroparesis_terms), na=False, regex=True)
    gastro_events = df[gastro_mask]

    if len(gastro_events) > 0:
        gastro_unique = gastro_events.drop_duplicates('safetyreportid')
        print(f"    FOUND: {len(gastro_unique)} unique reports")
        gastro_by_drug = gastro_unique['query_drug'].value_counts()
        for drug, count in gastro_by_drug.items():
            print(f"      - {drug.capitalize()}: {count}")

        # Specific terms found
        terms_found = gastro_events['reaction_meddraPT'].value_counts()
        print(f"    Terms found:")
        for term, count in terms_found.head(10).items():
            print(f"      - {term}: {count}")
        results['gastroparesis'] = {'count': len(gastro_unique), 'by_drug': gastro_by_drug.to_dict()}
    else:
        print("    NOT FOUND in dataset")
        results['gastroparesis'] = {'count': 0, 'by_drug': {}}

    # Check thyroid events
    print("\n  THYROID NEOPLASM AND RELATED:")
    thyroid_mask = df['reaction_normalized'].str.contains('|'.join(thyroid_terms), na=False, regex=True)
    thyroid_events = df[thyroid_mask]

    if len(thyroid_events) > 0:
        thyroid_unique = thyroid_events.drop_duplicates('safetyreportid')
        print(f"    FOUND: {len(thyroid_unique)} unique reports")
        thyroid_by_drug = thyroid_unique['query_drug'].value_counts()
        for drug, count in thyroid_by_drug.items():
            print(f"      - {drug.capitalize()}: {count}")

        # Specific terms found
        terms_found = thyroid_events['reaction_meddraPT'].value_counts()
        print(f"    Terms found:")
        for term, count in terms_found.head(15).items():
            print(f"      - {term}: {count}")
        results['thyroid'] = {'count': len(thyroid_unique), 'by_drug': thyroid_by_drug.to_dict()}
    else:
        print("    NOT FOUND in dataset")
        results['thyroid'] = {'count': 0, 'by_drug': {}}

    return results


def save_cleaned_data(df):
    """Save cleaned dataset to parquet format."""
    print(f"\n{'='*60}")
    print("STEP 9: Saving Cleaned Data")
    print(f"{'='*60}")

    # Select columns for output
    output_cols = [
        'safetyreportid', 'receivedate', 'year', 'quarter', 'year_quarter',
        'serious', 'is_serious', 'seriousnessdeath', 'seriousnesslifethreatening',
        'seriousnesshospitalization', 'seriousnessdisabling',
        'patient_age', 'patient_sex', 'patient_sex_label',
        'query_drug', 'reaction_meddraPT', 'reaction_normalized', 'reaction_outcome'
    ]

    df_output = df[output_cols].copy()

    # Save to parquet
    df_output.to_parquet(CLEANED_DATA_PATH, index=False)
    print(f"  Saved cleaned data: {CLEANED_DATA_PATH}")
    print(f"  Shape: {df_output.shape}")
    print(f"  File size: {CLEANED_DATA_PATH.stat().st_size / (1024**2):.2f} MB")

    # Also save CSV for inspection
    csv_path = SESSION_DIR / "workflow/data/glp1_cleaned.csv"
    df_output.to_csv(csv_path, index=False)
    print(f"  Also saved CSV: {csv_path}")

    return df_output


def generate_summary_report(df, pivot_data, top_20, serious_stats, signal_results):
    """Generate comprehensive EDA summary report."""
    print(f"\n{'='*60}")
    print("STEP 10: Generating Summary Report")
    print(f"{'='*60}")

    unique_reports = df.drop_duplicates(subset=['safetyreportid'])
    serious_counts, serious_by_drug = serious_stats

    report_lines = [
        "=" * 70,
        "GLP-1 RECEPTOR AGONISTS ADVERSE EVENTS - EDA SUMMARY REPORT",
        "=" * 70,
        f"Generated: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}",
        "",
        "DATA OVERVIEW",
        "-" * 40,
        f"Total unique adverse event reports: {len(unique_reports):,}",
        f"Total reaction-level observations: {len(df):,}",
        f"Date range: {df['receivedate'].min().strftime('%Y-%m-%d')} to {df['receivedate'].max().strftime('%Y-%m-%d')}",
        f"Years covered: {int(df['year'].min())} - {int(df['year'].max())}",
        "",
        "REPORTS BY DRUG",
        "-" * 40,
    ]

    drug_counts = unique_reports['query_drug'].value_counts()
    for drug, count in drug_counts.items():
        pct = count / len(unique_reports) * 100
        report_lines.append(f"  {drug.capitalize():15} {count:>6,} ({pct:>5.1f}%)")

    report_lines.extend([
        "",
        "SERIOUSNESS DISTRIBUTION",
        "-" * 40,
        f"  Serious events: {serious_counts.get(True, 0):,} ({serious_counts.get(True, 0)/len(unique_reports)*100:.1f}%)",
        f"  Non-serious:    {serious_counts.get(False, 0):,} ({serious_counts.get(False, 0)/len(unique_reports)*100:.1f}%)",
        "",
        "  Breakdown by outcome:",
        f"    Deaths:            {unique_reports['seriousnessdeath'].eq('1').sum():,}",
        f"    Life-threatening:  {unique_reports['seriousnesslifethreatening'].eq('1').sum():,}",
        f"    Hospitalizations:  {unique_reports['seriousnesshospitalization'].eq('1').sum():,}",
        f"    Disabling:         {unique_reports['seriousnessdisabling'].eq('1').sum():,}",
        "",
        "TOP 20 MOST FREQUENT ADVERSE EVENTS",
        "-" * 40,
    ])

    for i, (event, count) in enumerate(top_20.items(), 1):
        report_lines.append(f"  {i:2}. {event:40} {count:>6,}")

    report_lines.extend([
        "",
        "SIGNALS OF INTEREST",
        "-" * 40,
        "",
        "GASTROPARESIS/DELAYED GASTRIC EMPTYING:",
    ])

    if signal_results['gastroparesis']['count'] > 0:
        report_lines.append(f"  PRESENT - {signal_results['gastroparesis']['count']} unique reports")
        for drug, count in signal_results['gastroparesis']['by_drug'].items():
            report_lines.append(f"    {drug.capitalize()}: {count}")
    else:
        report_lines.append("  NOT FOUND")

    report_lines.extend([
        "",
        "THYROID NEOPLASM/DISORDERS:",
    ])

    if signal_results['thyroid']['count'] > 0:
        report_lines.append(f"  PRESENT - {signal_results['thyroid']['count']} unique reports")
        for drug, count in signal_results['thyroid']['by_drug'].items():
            report_lines.append(f"    {drug.capitalize()}: {count}")
    else:
        report_lines.append("  NOT FOUND")

    report_lines.extend([
        "",
        "PATIENT DEMOGRAPHICS",
        "-" * 40,
    ])

    sex_dist = unique_reports['patient_sex_label'].value_counts()
    for sex, count in sex_dist.items():
        report_lines.append(f"  {sex}: {count:,} ({count/len(unique_reports)*100:.1f}%)")

    # Age statistics
    ages = pd.to_numeric(unique_reports['patient_age'], errors='coerce')
    valid_ages = ages[(ages >= 0) & (ages <= 120)]
    if len(valid_ages) > 0:
        report_lines.extend([
            "",
            f"  Age statistics (n={len(valid_ages):,}):",
            f"    Mean: {valid_ages.mean():.1f} years",
            f"    Median: {valid_ages.median():.1f} years",
            f"    Range: {valid_ages.min():.0f} - {valid_ages.max():.0f} years",
        ])

    report_lines.extend([
        "",
        "OUTPUT FILES",
        "-" * 40,
        f"  Cleaned data: workflow/data/glp1_cleaned.parquet",
        f"  Reports over time: figures/reports_over_time.png",
        f"  Top adverse events: figures/top_adverse_events.png",
        f"  This summary: results/eda_summary.txt",
        "",
        "=" * 70,
        "END OF REPORT",
        "=" * 70,
    ])

    report_text = "\n".join(report_lines)

    # Save report
    report_path = RESULTS_DIR / "eda_summary.txt"
    with open(report_path, 'w') as f:
        f.write(report_text)

    print(f"  Saved: {report_path}")
    print("\n" + report_text)

    return report_text


def main():
    """Main execution function."""
    print("\n" + "=" * 70)
    print("GLP-1 ADVERSE EVENTS - STEP 2: DATA CLEANING AND EDA")
    print("=" * 70)
    print(f"Session: {SESSION_DIR}")
    print(f"Started: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")

    # Step 1: Load raw data
    data = load_raw_data()

    # Step 2: Deduplicate
    unique_data = deduplicate_records(data)

    # Step 3: Flatten records
    flattened_data = flatten_records(unique_data)

    # Step 4: Create DataFrame
    df = create_dataframe(flattened_data)

    # Step 5: Temporal analysis
    pivot_data = analyze_temporal_trends(df)

    # Step 6: Top adverse events
    top_20 = analyze_top_adverse_events(df)

    # Step 7: Seriousness analysis
    serious_stats = analyze_seriousness(df)

    # Step 8: Signal detection
    signal_results = check_signals_of_interest(df)

    # Step 9: Save cleaned data
    save_cleaned_data(df)

    # Step 10: Generate summary report
    generate_summary_report(df, pivot_data, top_20, serious_stats, signal_results)

    print(f"\n{'='*70}")
    print("STEP 2 COMPLETE")
    print(f"{'='*70}")
    print(f"Finished: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")


if __name__ == "__main__":
    main()
