"""
Supplementary File S2: Monte Carlo Simulation of Fontan-Pathway Birth Incidence

Accompanies: "Beyond 70,000: Reassessing the Global Fontan Population Through
Citation-Chain Analysis and Tiered Epidemiologic Estimation"
Author: Marie-Josée Flora Herard
Journal: Pediatric Cardiology (submitted 2026)

This script produces the two-tier annual birth incidence estimates and Figure 1
reported in the manuscript. It requires only Python 3.8+ and NumPy/Matplotlib
(standard Anaconda or pip install).

Usage:
    python S2_fontan_birth_incidence_simulation.py

Outputs:
    - Console: composite prevalence rates, annual birth counts, 95% CIs,
      variance decomposition, and per-diagnosis contributions
    - Figure: monte_carlo_histogram.png (publication-ready, 300 dpi)

Seed is fixed (42) for exact reproducibility of reported results.
"""

import numpy as np
import matplotlib.pyplot as plt
import matplotlib.ticker as ticker

SEED = 42
N_ITERATIONS = 10_000
GLOBAL_BIRTHS = 140_000_000  # World Bank 2023 estimate

# ─────────────────────────────────────────────────────────────────────────────
# INPUT DATA: Nine Fontan-pathway cardiac diagnoses
#
# Each entry: (diagnosis_name, {
#     "prev":    (low, mid, high) birth prevalence per 1,000 live births,
#     "frac":    (low, mid, high) fraction with functionally univentricular
#                anatomy (Tier 1: anatomic candidates),
#     "surv":    (low, mid, high) survival-to-Fontan-completion rate
#                (Tier 2: anticipated completers),
#     "overlap": fixed overlap correction (subtracted from prevalence to
#                avoid double-counting across ICD-coded diagnoses)
# })
#
# Sources for each parameter are listed inline. All prevalence figures are
# drawn from population-based registries and meta-analyses; candidate
# fractions and survival rates from published cohort and registry outcomes.
# ─────────────────────────────────────────────────────────────────────────────

DIAGNOSES = [
    ("HLHS", {
        # Prevalence: Reller et al. 2008 (Atlanta, 0.23/1,000);
        # range 0.08-0.29 across registries (PMC10566614);
        # Liu et al. 2019 meta-analysis consistent
        "prev": (0.16, 0.23, 0.34),
        # Anatomic candidate fraction: HLHS is by definition univentricular
        # (includes AA-MA subtype, ~40-45% of HLHS)
        "frac": (0.95, 0.98, 1.00),
        # Survival to Fontan completion: population-level data lower than
        # tertiary-center SVR trial; 18% of HLHS neonates at US tertiary
        # centers not palliated (PMC11708278); Stage 1 mortality ~10-15%,
        # interstage ~10%, pre-Fontan attrition ~10% (PMC10296158)
        "surv": (0.45, 0.55, 0.65),
        "overlap": 0.0,
    }),
    ("Tricuspid atresia", {
        # Prevalence: Reller et al. 2008; StatPearls ~0.10/1,000;
        # multiple registries 0.05-0.12/1,000
        "prev": (0.05, 0.10, 0.11),
        # ~85-95% have functionally univentricular physiology
        "frac": (0.85, 0.90, 0.95),
        # Higher survival than HLHS due to systemic left ventricle;
        # 70-80% of TA survivors proceed to Fontan (StatPearls NBK558950)
        "surv": (0.60, 0.70, 0.80),
        "overlap": 0.0,
    }),
    ("Double inlet left ventricle", {
        # Prevalence: EUROCAT registries 0.05-0.10/1,000 (PMC3653524);
        # Bull 1999 (UK population-based)
        "prev": (0.05, 0.07, 0.10),
        # By definition univentricular
        "frac": (0.95, 0.98, 1.00),
        # Outcomes similar to tricuspid atresia
        "surv": (0.60, 0.70, 0.80),
        "overlap": 0.0,
    }),
    ("PA/IVS", {
        # Pulmonary atresia with intact ventricular septum
        # Prevalence: Bull 1999 (UK/Eire, 0.04-0.05/1,000); Reller 2008
        "prev": (0.04, 0.05, 0.07),
        # 25-35% have RV too small for biventricular repair and proceed
        # to single-ventricle pathway (PMC6571047); lower than the 40-60%
        # "small RV" fraction because some receive 1.5-ventricle repair
        "frac": (0.25, 0.30, 0.35),
        # Survival for the single-ventricle subset
        "surv": (0.55, 0.65, 0.75),
        "overlap": 0.0,
    }),
    ("Unbalanced AVSD", {
        # Subset of complete AVSD (~0.20/1,000); unbalanced ~10-15%
        # Derived prevalence of unbalanced subset
        "prev": (0.01, 0.02, 0.03),
        # Those identified as unbalanced are by definition SV candidates
        "frac": (0.85, 0.90, 0.95),
        # Trisomy 21 comorbidity affects outcomes in a subset;
        # Buratto et al. 2017 (multicenter uAVSD Fontan outcomes)
        "surv": (0.50, 0.60, 0.70),
        "overlap": 0.0,
    }),
    ("DORV (univentricular subset)", {
        # Double outlet right ventricle — total prevalence ~0.09/1,000
        # (range 0.06-0.13; Reller et al. 2008; Bjarke et al. 2023)
        # Only the subset with non-committed VSD or straddling AV valve
        "prev": (0.06, 0.09, 0.13),
        # ~20-30% of DORV is functionally univentricular
        "frac": (0.20, 0.25, 0.30),
        # Outcomes vary widely depending on specific anatomy
        "surv": (0.50, 0.60, 0.70),
        "overlap": 0.0,
    }),
    ("Heterotaxy syndrome", {
        # Laterality disorder frequently associated with complex SV anatomy
        # Prevalence: Lopez et al. 2015 (Texas, 0.118/1,000);
        # Western Australia 0.048/1,000;
        # low end reflects possible under-ascertainment in smaller registries
        "prev": (0.07, 0.10, 0.12),
        # ~55-80% have functionally univentricular cardiac anatomy
        # Kim et al. 2013: 59%; Bartz et al. 2006 supporting
        "frac": (0.55, 0.70, 0.80),
        # Lower survival due to associated anomalies (asplenia, etc.)
        "surv": (0.40, 0.50, 0.60),
        "overlap": 0.0,
    }),
    ("Severe Ebstein anomaly", {
        # Severe end of Ebstein spectrum — total Ebstein ~0.05/1,000
        # Prevalence: Reller et al. 2008
        "prev": (0.03, 0.05, 0.07),
        # 5-10% of Ebstein patients have severe neonatal presentation
        # requiring Starnes procedure and eventual Fontan pathway
        # (Knott-Craig et al. Circulation 2017)
        "frac": (0.05, 0.08, 0.10),
        # Neonatal presentation has lower survival
        "surv": (0.40, 0.50, 0.60),
        "overlap": 0.0,
    }),
    ("ccTGA (univentricular subset)", {
        # Congenitally corrected TGA — ~0.02-0.03/1,000 (~1:33,000)
        # Only the subset with associated lesions rendering biventricular
        # repair unfeasible
        "prev": (0.02, 0.03, 0.035),
        # STS Database 2010-2019: ~30% of ccTGA operations were SV
        # palliations (Hraska et al. 2025)
        "frac": (0.20, 0.30, 0.35),
        # Variable outcomes depending on associated lesions
        "surv": (0.50, 0.60, 0.70),
        "overlap": 0.0,
    }),
]

# Liu et al. 2019 narrow "single ventricle" (ICD-9 745.3 / ICD-10 Q20.4)
LIU_NARROW_SV = 0.093  # per 1,000 live births


def run_simulation():
    rng = np.random.default_rng(SEED)

    n_dx = len(DIAGNOSES)
    tier1_by_dx = np.zeros((N_ITERATIONS, n_dx))
    tier2_by_dx = np.zeros((N_ITERATIONS, n_dx))

    for i, (name, params) in enumerate(DIAGNOSES):
        prev_low, prev_mid, prev_high = params["prev"]
        frac_low, frac_mid, frac_high = params["frac"]
        surv_low, surv_mid, surv_high = params["surv"]

        prev_samples = rng.triangular(prev_low, prev_mid, prev_high, N_ITERATIONS)
        frac_samples = rng.triangular(frac_low, frac_mid, frac_high, N_ITERATIONS)
        surv_samples = rng.triangular(surv_low, surv_mid, surv_high, N_ITERATIONS)

        candidate_prev = (prev_samples - params["overlap"]) * frac_samples
        tier1_by_dx[:, i] = candidate_prev
        tier2_by_dx[:, i] = candidate_prev * surv_samples

    tier1_composite = tier1_by_dx.sum(axis=1)
    tier2_composite = tier2_by_dx.sum(axis=1)

    return tier1_composite, tier2_composite, tier1_by_dx, tier2_by_dx


def print_results(tier1, tier2, tier1_by_dx, tier2_by_dx):
    print("=" * 70)
    print("FONTAN-PATHWAY BIRTH INCIDENCE: TWO-TIER MONTE CARLO ESTIMATE")
    print(f"  Iterations: {N_ITERATIONS:,}  |  Seed: {SEED}  |  "
          f"Global births: {GLOBAL_BIRTHS:,}")
    print("=" * 70)

    # Tier 1: Anatomic candidates
    t1_med = np.median(tier1)
    t1_lo = np.percentile(tier1, 2.5)
    t1_hi = np.percentile(tier1, 97.5)
    t1_births_med = t1_med * GLOBAL_BIRTHS / 1000
    t1_births_lo = t1_lo * GLOBAL_BIRTHS / 1000
    t1_births_hi = t1_hi * GLOBAL_BIRTHS / 1000

    print(f"\nTIER 1 — Anatomic candidates (born with Fontan-pathway anatomy)")
    print(f"  Composite rate:  {t1_med:.4f}/1,000  "
          f"(95% CI: {t1_lo:.4f}–{t1_hi:.4f})")
    print(f"  Annual births:   ~{t1_births_med:,.0f}  "
          f"(95% CI: ~{t1_births_lo:,.0f}–{t1_births_hi:,.0f})")

    # Tier 2: Anticipated completers
    t2_med = np.median(tier2)
    t2_lo = np.percentile(tier2, 2.5)
    t2_hi = np.percentile(tier2, 97.5)
    t2_births_med = t2_med * GLOBAL_BIRTHS / 1000
    t2_births_lo = t2_lo * GLOBAL_BIRTHS / 1000
    t2_births_hi = t2_hi * GLOBAL_BIRTHS / 1000

    print(f"\nTIER 2 — Anticipated completers (surviving to completed Fontan)")
    print(f"  Composite rate:  {t2_med:.4f}/1,000  "
          f"(95% CI: {t2_lo:.4f}–{t2_hi:.4f})")
    print(f"  Annual births:   ~{t2_births_med:,.0f}  "
          f"(95% CI: ~{t2_births_lo:,.0f}–{t2_births_hi:,.0f})")

    # Liu et al. comparison
    liu_births = LIU_NARROW_SV * GLOBAL_BIRTHS / 1000
    prob_t2_exceeds = np.mean(tier2 > LIU_NARROW_SV) * 100
    print(f"\nLiu et al. (2019) narrow SV:  {LIU_NARROW_SV}/1,000  "
          f"(~{liu_births:,.0f}/yr)")
    print(f"  P(Tier 2 > Liu narrow SV) = {prob_t2_exceeds:.1f}%")

    # Per-diagnosis breakdown (Tier 1)
    print(f"\n{'─' * 70}")
    print("PER-DIAGNOSIS CONTRIBUTION TO TIER 1 (anatomic candidates)")
    print(f"{'─' * 70}")
    print(f"  {'Diagnosis':<35} {'Median/1k':>10} {'Annual':>10} {'Share':>8}")
    print(f"  {'─' * 63}")

    dx_medians = np.median(tier1_by_dx, axis=0)
    total_median = dx_medians.sum()
    for i, (name, _) in enumerate(DIAGNOSES):
        med = dx_medians[i]
        births = med * GLOBAL_BIRTHS / 1000
        share = med / total_median * 100
        print(f"  {name:<35} {med:>10.4f} {births:>10,.0f} {share:>7.1f}%")

    # Variance decomposition
    print(f"\n{'─' * 70}")
    print("VARIANCE DECOMPOSITION")
    print(f"{'─' * 70}")
    total_var = np.var(tier1)
    for i, (name, _) in enumerate(DIAGNOSES):
        dx_var = np.var(tier1_by_dx[:, i])
        pct = dx_var / total_var * 100
        print(f"  {name:<35} {pct:>6.1f}% of total variance")


def make_figure(tier1, tier2, tier1_by_dx):
    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 5.5),
                                    gridspec_kw={"width_ratios": [1.6, 1]})

    # Left panel: distribution histograms
    t1_births = tier1 * GLOBAL_BIRTHS / 1000
    t2_births = tier2 * GLOBAL_BIRTHS / 1000

    ax1.hist(t1_births, bins=80, alpha=0.55, color="#2166ac",
             label="Tier 1: Anatomic candidates", density=True)
    ax1.hist(t2_births, bins=80, alpha=0.55, color="#b2182b",
             label="Tier 2: Anticipated completers", density=True)

    # 95% CI shading
    for data, color in [(t1_births, "#2166ac"), (t2_births, "#b2182b")]:
        lo, hi = np.percentile(data, [2.5, 97.5])
        ax1.axvspan(lo, hi, alpha=0.08, color=color)

    # Median annotation lines
    t1_med = np.median(t1_births)
    t2_med = np.median(t2_births)
    ax1.axvline(t1_med, color="#2166ac", linestyle="-", linewidth=1.5,
                alpha=0.85)
    ax1.axvline(t2_med, color="#b2182b", linestyle="-", linewidth=1.5,
                alpha=0.85)

    ymax = ax1.get_ylim()[1]
    ax1.annotate(f"~{t1_med:,.0f}/yr",
                 xy=(t1_med, ymax * 0.78), fontsize=9.5, fontweight="bold",
                 color="#2166ac", ha="center",
                 bbox=dict(boxstyle="round,pad=0.25", fc="white",
                           ec="#2166ac", alpha=0.85))
    ax1.annotate(f"~{t2_med:,.0f}/yr",
                 xy=(t2_med, ymax * 0.78), fontsize=9.5, fontweight="bold",
                 color="#b2182b", ha="center",
                 bbox=dict(boxstyle="round,pad=0.25", fc="white",
                           ec="#b2182b", alpha=0.85))

    # Liu et al. reference line
    liu_line = LIU_NARROW_SV * GLOBAL_BIRTHS / 1000
    ax1.axvline(liu_line, color="gray", linestyle="--", linewidth=1.2,
                label=f"Liu et al. narrow SV (~{liu_line:,.0f}/yr)")

    ax1.set_xlabel("Annual global Fontan-pathway births", fontsize=11)
    ax1.set_ylabel("Density", fontsize=11)
    ax1.set_title("Distribution of Annual Birth Incidence Estimates",
                   fontsize=12, fontweight="bold")
    ax1.legend(fontsize=9, loc="upper right")
    ax1.xaxis.set_major_formatter(ticker.FuncFormatter(
        lambda x, _: f"{x/1000:.0f}k"))

    # Right panel: per-diagnosis bar chart (Tier 1)
    dx_names = [name for name, _ in DIAGNOSES]
    dx_medians = np.median(tier1_by_dx, axis=0) * GLOBAL_BIRTHS / 1000
    total_median = dx_medians.sum()
    sort_idx = np.argsort(dx_medians)

    bars = ax2.barh([dx_names[i] for i in sort_idx],
                    [dx_medians[i] for i in sort_idx],
                    color="#2166ac", alpha=0.7)

    for bar, idx in zip(bars, sort_idx):
        val = dx_medians[idx]
        pct = val / total_median * 100
        if pct >= 3:
            ax2.text(val + total_median * 0.01, bar.get_y() + bar.get_height() / 2,
                     f"{val:,.0f}  ({pct:.0f}%)",
                     va="center", ha="left", fontsize=8, color="#333333")
        else:
            ax2.text(val + total_median * 0.01, bar.get_y() + bar.get_height() / 2,
                     f"{val:,.0f}",
                     va="center", ha="left", fontsize=7.5, color="#666666")

    ax2.set_xlabel("Annual births (Tier 1)", fontsize=11)
    ax2.set_title("Per-Diagnosis Contribution",
                   fontsize=12, fontweight="bold")
    ax2.xaxis.set_major_formatter(ticker.FuncFormatter(
        lambda x, _: f"{x/1000:.0f}k"))
    ax2.set_xlim(right=total_median * 0.62)

    plt.tight_layout()
    plt.savefig("monte_carlo_histogram.png", dpi=300, bbox_inches="tight")
    print(f"\nFigure saved: monte_carlo_histogram.png")
    plt.close()


if __name__ == "__main__":
    tier1, tier2, tier1_by_dx, tier2_by_dx = run_simulation()
    print_results(tier1, tier2, tier1_by_dx, tier2_by_dx)
    make_figure(tier1, tier2, tier1_by_dx)
