#!/usr/bin/env python3
"""Adapted population_allele_freq analysis — ABCA4 ClinVar pathogenic vs benign
variant frequencies in 1000 Genomes super-populations.

Question: For ABCA4, do ClinVar-pathogenic variants have lower 1000 Genomes
allele frequency than ClinVar-benign variants?

Inputs:
  --clinvar  CSV from fetch_clinvar.py (clinvar_uid, clinical_significance, ...)
  --kg       CSV from fetch_1000g.py (chrom, pos, id, ref, alt, af, eas_af, ...)

Outputs:
  figure.png  — box+strip of log10(AF) for pathogenic vs benign + per-pop AF bars
  stats.json  — Mann-Whitney U + chi-square aggregated allele counts
"""
from __future__ import annotations

import argparse
import json
import sys

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from scipy import stats

PATHOGENIC = ("pathogenic", "likely pathogenic", "likely_pathogenic")
BENIGN = ("benign", "likely benign", "likely_benign")
POPS = ["afr", "amr", "eas", "eur", "sas"]
SUPERPOP_2N = {"afr": 1322, "amr": 694, "eas": 1008, "eur": 1006, "sas": 978}


def classify(sig: str) -> str | None:
    s = (sig or "").strip().lower()
    if not s:
        return None
    if any(p in s for p in PATHOGENIC) and "conflicting" not in s:
        return "pathogenic"
    if any(b in s for b in BENIGN) and "conflicting" not in s:
        return "benign"
    return None


def main() -> int:
    ap = argparse.ArgumentParser(description=__doc__,
                                 formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("--clinvar", required=True)
    ap.add_argument("--kg", required=True)
    ap.add_argument("--seed", type=int, default=1234)
    ap.add_argument("--outdir", default=".")
    args = ap.parse_args()

    rng = np.random.default_rng(args.seed)

    clinvar = pd.read_csv(args.clinvar, dtype=str).fillna("")
    kg = pd.read_csv(args.kg, dtype=str).fillna("")

    clinvar["call"] = clinvar["clinical_significance"].map(classify)
    clinvar = clinvar.dropna(subset=["call"])

    for p in POPS:
        kg[f"{p}_af"] = pd.to_numeric(kg.get(f"{p}_af"), errors="coerce")
    kg["af"] = pd.to_numeric(kg.get("af"), errors="coerce")

    clinvar["pos_int"] = pd.to_numeric(clinvar["pos"], errors="coerce")
    clinvar = clinvar.dropna(subset=["pos_int"])
    clinvar["pos_int"] = clinvar["pos_int"].astype(int)

    kg["pos_int"] = pd.to_numeric(kg["pos"], errors="coerce")
    kg = kg.dropna(subset=["pos_int"])
    kg["pos_int"] = kg["pos_int"].astype(int)

    kg_dedup = kg.dropna(subset=["af"])
    kg_dedup = kg_dedup[kg_dedup["af"] > 0].copy()
    kg_dedup = kg_dedup.sort_values("af", ascending=False).drop_duplicates(subset=["pos_int"])

    merged_rows = []
    for _, cv_row in clinvar.iterrows():
        pos = cv_row["pos_int"]
        nearby = kg_dedup[(kg_dedup["pos_int"] >= pos - 1) & (kg_dedup["pos_int"] <= pos + 1)]
        if len(nearby) > 0:
            kg_row = nearby.iloc[0]
            merged_rows.append({
                "locus": f"1:{pos}",
                "call": cv_row["call"],
                "af": kg_row["af"],
                **{f"{p}_af": kg_row[f"{p}_af"] for p in POPS},
            })

    merged = pd.DataFrame(merged_rows)
    if merged.empty:
        result = {
            "gene": "ABCA4", "seed": args.seed,
            "reference_template": "population_allele_freq",
            "outcome": "null_result",
            "reason": "no ClinVar variants matched 1000G positions even with ±1bp tolerance",
            "n_clinvar_classified": int(len(clinvar)),
            "n_merged_with_1kg_af": 0,
            "n_pathogenic": 0, "n_benign": 0,
        }
        _write_stats(args.outdir, result)
        _render(args.outdir, np.array([]), np.array([]), result, rng)
        print(json.dumps(result, indent=2))
        return 0

    path_af = merged.loc[merged["call"] == "pathogenic", "af"].to_numpy(dtype=float)
    ben_af = merged.loc[merged["call"] == "benign", "af"].to_numpy(dtype=float)

    result = {
        "gene": "ABCA4",
        "seed": args.seed,
        "reference_template": "population_allele_freq",
        "join_key": "genomic locus chrom:pos (±1bp tolerance)",
        "n_clinvar_classified": int(len(clinvar)),
        "n_merged_with_1kg_af": int(len(merged)),
        "n_pathogenic": int(len(path_af)),
        "n_benign": int(len(ben_af)),
    }

    if len(path_af) < 3 or len(ben_af) < 3:
        result["outcome"] = "null_result"
        result["reason"] = (
            "fewer than 3 variants in one group after the ClinVar x 1000G join; "
            "no test performed (reported honestly, not fabricated)"
        )
        result["median_af_pathogenic"] = float(np.median(path_af)) if len(path_af) else None
        result["median_af_benign"] = float(np.median(ben_af)) if len(ben_af) else None
        _write_stats(args.outdir, result)
        _render(args.outdir, path_af, ben_af, result, rng)
        print(json.dumps(result, indent=2))
        return 0

    u_stat, p_value = stats.mannwhitneyu(path_af, ben_af, alternative="two-sided")
    med_path = float(np.median(path_af))
    med_ben = float(np.median(ben_af))

    path_pop_afs = merged.loc[merged["call"] == "pathogenic", [f"{p}_af" for p in POPS]].to_numpy()
    ben_pop_afs = merged.loc[merged["call"] == "benign", [f"{p}_af" for p in POPS]].to_numpy()
    path_pop_afs = np.nan_to_num(path_pop_afs, nan=0.0)
    ben_pop_afs = np.nan_to_num(ben_pop_afs, nan=0.0)

    path_ac_per_pop = np.round(path_pop_afs * np.array([SUPERPOP_2N[p] for p in POPS])).astype(int)
    ben_ac_per_pop = np.round(ben_pop_afs * np.array([SUPERPOP_2N[p] for p in POPS])).astype(int)
    path_ac_total = path_ac_per_pop.sum(axis=0)
    ben_ac_total = ben_ac_per_pop.sum(axis=0)
    path_an_total = np.array([SUPERPOP_2N[p] for p in POPS]) * len(path_af)
    ben_an_total = np.array([SUPERPOP_2N[p] for p in POPS]) * len(ben_af)
    path_ref_total = path_an_total - path_ac_total
    ben_ref_total = ben_an_total - ben_ac_total

    table = np.array([[int(x) for x in path_ac_total],
                      [int(x) for x in path_ref_total],
                      [int(x) for x in ben_ac_total],
                      [int(x) for x in ben_ref_total]])
    try:
        chi2, chi_p, chi_dof, _ = stats.chi2_contingency(table)
    except ValueError:
        chi2, chi_p, chi_dof = 0.0, 1.0, 1

    path_pop_mean_af = path_pop_afs.mean(axis=0)
    ben_pop_mean_af = ben_pop_afs.mean(axis=0)

    result.update({
        "outcome": "success" if p_value < 0.05 else "null_result",
        "test": "Mann-Whitney U (two-sided) on allele frequency",
        "u_statistic": float(u_stat),
        "p_value": float(p_value),
        "median_af_pathogenic": med_path,
        "median_af_benign": med_ben,
        "direction": ("pathogenic rarer" if med_path < med_ben else "pathogenic not rarer"),
        "significant_at_0.05": bool(p_value < 0.05),
        "headline_statistic": f"p = {p_value:.2e} (Mann-Whitney U)",
        "chi_square_aggregated": {
            "chi2": float(chi2),
            "dof": int(chi_dof),
            "p_value": float(chi_p),
            "description": "2x2x5 chi-square: (pathogenic/benign) x (alt/ref) x 5 super-pops",
        },
        "per_pop_mean_af_pathogenic": {p: float(path_pop_mean_af[i]) for i, p in enumerate(POPS)},
        "per_pop_mean_af_benign": {p: float(ben_pop_mean_af[i]) for i, p in enumerate(POPS)},
    })
    _write_stats(args.outdir, result)
    _render(args.outdir, path_af, ben_af, result, rng)
    print(json.dumps(result, indent=2))
    return 0


def _write_stats(outdir, result):
    with open(f"{outdir}/stats.json", "w") as fh:
        json.dump(result, fh, indent=2)


def _render(outdir, path_af, ben_af, result, rng):
    fig, axes = plt.subplots(1, 2, figsize=(11, 5))

    ax = axes[0]
    groups, labels, colors = [], [], []
    if len(path_af):
        groups.append(np.log10(path_af))
        labels.append(f"Pathogenic\n(n={len(path_af)})")
        colors.append("#c0392b")
    if len(ben_af):
        groups.append(np.log10(ben_af))
        labels.append(f"Benign\n(n={len(ben_af)})")
        colors.append("#2980b9")

    if groups:
        bp = ax.boxplot(groups, patch_artist=True, widths=0.5, showfliers=False)
        for patch, color in zip(bp["boxes"], colors):
            patch.set_facecolor(color)
            patch.set_alpha(0.35)
        for i, (g, color) in enumerate(zip(groups, colors), start=1):
            jitter = rng.uniform(-0.12, 0.12, size=len(g))
            ax.scatter(np.full(len(g), i) + jitter, g, s=14, color=color,
                       alpha=0.6, edgecolors="none", zorder=3)
        ax.set_xticks(range(1, len(labels) + 1))
        ax.set_xticklabels(labels)
    ax.set_ylabel("log10(1000G allele frequency)")
    subtitle = result.get("headline_statistic") or result.get("reason", "")
    ax.set_title(f"ABCA4: ClinVar significance vs 1000G frequency\n{subtitle}",
                 fontsize=10)
    ax.grid(axis="y", alpha=0.3)

    ax = axes[1]
    pops_display = [p.upper() for p in POPS]
    path_pop_mean = result.get("per_pop_mean_af_pathogenic", {})
    ben_pop_mean = result.get("per_pop_mean_af_benign", {})
    x = np.arange(len(POPS))
    w = 0.35
    ax.bar(x - w/2, [path_pop_mean.get(p, 0) for p in POPS], w,
           label="Pathogenic", color="#c0392b", alpha=0.7)
    ax.bar(x + w/2, [ben_pop_mean.get(p, 0) for p in POPS], w,
           label="Benign", color="#2980b9", alpha=0.7)
    ax.set_xticks(x)
    ax.set_xticklabels(pops_display)
    ax.set_ylabel("Mean allele frequency")
    ax.set_title("Per-super-population AF\n(mean across variants)", fontsize=10)
    ax.legend(fontsize=8)
    ax.grid(axis="y", alpha=0.3)

    fig.tight_layout()
    fig.savefig(f"{outdir}/figure.png", dpi=130)
    plt.close(fig)


if __name__ == "__main__":
    sys.exit(main())
