#!/usr/bin/env python3
"""Orchestrates the MSA 2-piece-affine-gap experiment: simulate data, grid-search gap
params on dev, evaluate baseline vs proposed on test, write results.json + figs."""
import json
import os
import subprocess
import sys
import time
from itertools import product

import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt

ROOT = os.path.dirname(os.path.abspath(__file__))
BIN_SIM = os.path.join(ROOT, "src", "simulate")
BIN_MSA = os.path.join(ROOT, "src", "msa")
DATA = os.path.join(ROOT, "data")
FIGS = os.path.join(ROOT, "figs")
os.makedirs(DATA, exist_ok=True)
os.makedirs(FIGS, exist_ok=True)

N_TAXA = 6
ROOT_LEN = 220
DEV_SEEDS = [100, 101, 102]
TEST_SEEDS = [0, 1, 2]
FAMS_PER_SEED = 4  # deviation from plan's 10, see NOTES.md / deviations
PI_LONGS = [0.00, 0.05, 0.20]

BASELINE_GRID = list(product([4, 6, 8, 10, 12], [0.5, 1, 2, 3]))          # 20 combos
PROPOSED_GRID_O2E2 = list(product([12, 16, 20, 24, 28], [0.05, 0.1, 0.2, 0.4]))  # 20 combos

LONG_INDEL_MIN_LEN = 10
FLANK_WINDOW = 5


def read_fasta(path):
    names, seqs = [], []
    cur = []
    with open(path) as f:
        for line in f:
            line = line.rstrip("\n")
            if not line:
                continue
            if line.startswith(">"):
                if cur:
                    seqs.append("".join(cur))
                names.append(line[1:])
                cur = []
            else:
                cur.append(line)
        if cur:
            seqs.append("".join(cur))
    return names, seqs


def simulate(seed, pi_long, n_fam, prefix):
    out_prefix = os.path.join(DATA, prefix)
    cmd = [BIN_SIM, str(seed), str(pi_long), str(n_fam), str(N_TAXA), str(ROOT_LEN), out_prefix]
    r = subprocess.run(cmd, capture_output=True, text=True, check=True)
    print(f"[simulate] {' '.join(cmd)} :: {r.stderr.strip()}")
    return out_prefix


def run_msa(unaligned_fa, mode, params):
    cmd = [BIN_MSA, unaligned_fa, str(mode)] + [str(p) for p in params]
    r = subprocess.run(cmd, capture_output=True, text=True, check=True)
    runtime_ms, peak_bytes = None, None
    for tok in r.stderr.strip().split():
        if tok.startswith("RUNTIME_MS="):
            runtime_ms = float(tok.split("=")[1])
        if tok.startswith("PEAK_DP_BYTES="):
            peak_bytes = int(tok.split("=")[1])
    names, seqs = _parse_fasta_text(r.stdout)
    return names, seqs, runtime_ms, peak_bytes


def _parse_fasta_text(text):
    names, seqs = [], []
    cur = []
    for line in text.split("\n"):
        if not line:
            continue
        if line.startswith(">"):
            if cur:
                seqs.append("".join(cur))
            names.append(line[1:])
            cur = []
        else:
            cur.append(line)
    if cur:
        seqs.append("".join(cur))
    return names, seqs


def build_col_maps(names, seqs):
    """For each seq: map ungapped residue index -> column index."""
    order = {n: i for i, n in enumerate(names)}
    maps = {}
    for n, s in zip(names, seqs):
        m = []
        for col, ch in enumerate(s):
            if ch != "-":
                m.append(col)
        maps[n] = m  # maps[n][pos] = column index
    return maps, order


def true_pairs_by_column(names, seqs, restrict_cols=None):
    """Return set of (seqA,posA,seqB,posB) with seqA<seqB for residues co-occurring
    in the same column. If restrict_cols given, only pairs whose column is in it."""
    L = len(seqs[0])
    n = len(seqs)
    # per-seq running ungapped position counter
    pos_counter = [0] * n
    pairs = set()
    for col in range(L):
        if restrict_cols is not None and col not in restrict_cols:
            for i in range(n):
                if seqs[i][col] != "-":
                    pos_counter[i] += 1
            continue
        present = []
        for i in range(n):
            if seqs[i][col] != "-":
                present.append((i, pos_counter[i]))
                pos_counter[i] += 1
        for a in range(len(present)):
            for b in range(a + 1, len(present)):
                ia, pa = present[a]
                ib, pb = present[b]
                pairs.add((names[ia], pa, names[ib], pb))
    return pairs


def computed_pair_lookup(names, seqs):
    """Same as true_pairs_by_column but no restriction; used to check membership fast."""
    return true_pairs_by_column(names, seqs, restrict_cols=None)


def long_indel_flank_columns(true_names, true_seqs):
    """Columns within FLANK_WINDOW of a true gap-run (in any sequence) of length>=LONG_INDEL_MIN_LEN."""
    L = len(true_seqs[0])
    cols = set()
    for s in true_seqs:
        i = 0
        while i < L:
            if s[i] == "-":
                j = i
                while j < L and s[j] == "-":
                    j += 1
                run_len = j - i
                if run_len >= LONG_INDEL_MIN_LEN:
                    for c in range(max(0, i - FLANK_WINDOW), i):
                        cols.add(c)
                    for c in range(j, min(L, j + FLANK_WINDOW)):
                        cols.add(c)
                i = j
            else:
                i += 1
    return cols


def sp_score(true_names, true_seqs, comp_names, comp_seqs, restrict_cols=None):
    tp = true_pairs_by_column(true_names, true_seqs, restrict_cols=restrict_cols)
    if not tp:
        return None
    cp = computed_pair_lookup(comp_names, comp_seqs)
    hit = len(tp & cp)
    return hit / len(tp)


def tc_score(true_names, true_seqs, comp_names, comp_seqs):
    order_t = {n: i for i, n in enumerate(true_names)}
    order_c = {n: i for i, n in enumerate(comp_names)}
    n = len(true_names)
    # reorder comp seqs to true_names order
    comp_by_name = dict(zip(comp_names, comp_seqs))
    comp_seqs_ord = [comp_by_name[n_] for n_ in true_names]

    def col_tuples(seqs):
        L = len(seqs[0])
        pos_counter = [0] * n
        tuples = []
        for col in range(L):
            key = []
            for i in range(n):
                if seqs[i][col] != "-":
                    key.append(pos_counter[i])
                    pos_counter[i] += 1
                else:
                    key.append(None)
            tuples.append(tuple(key))
        return tuples

    true_cols = col_tuples(true_seqs)
    comp_cols = col_tuples(comp_seqs_ord)
    comp_set = set(comp_cols)
    n_all_gap = 0
    hit = 0
    total = 0
    for t in true_cols:
        if all(x is None for x in t):
            continue
        total += 1
        if t in comp_set:
            hit += 1
    return hit / total if total else None


def evaluate_family(true_path, comp_names, comp_seqs):
    true_names, true_seqs = read_fasta(true_path)
    sp = sp_score(true_names, true_seqs, comp_names, comp_seqs)
    tc = tc_score(true_names, true_seqs, comp_names, comp_seqs)
    flank_cols = long_indel_flank_columns(true_names, true_seqs)
    sp_flank = sp_score(true_names, true_seqs, comp_names, comp_seqs, restrict_cols=flank_cols) if flank_cols else None
    return sp, tc, sp_flank, (len(flank_cols) > 0)


def eval_grid_point(prefix, n_seeds_fams, mode, params):
    """Run msa with given params over all (seed,fam) dev families, return mean SP."""
    sps = []
    for seed, fam in n_seeds_fams:
        unaligned = f"{prefix}{seed}_fam{fam}_unaligned.fa"
        true_path = f"{prefix}{seed}_fam{fam}_true.fa"
        names, seqs, _, _ = run_msa(unaligned, mode, params)
        sp, _, _, _ = evaluate_family(true_path, names, seqs)
        if sp is not None:
            sps.append(sp)
    return float(np.mean(sps)) if sps else 0.0


def main():
    t_start = time.time()
    log = []
    results = {
        "task_slug": "msa-2piece-affine-gap",
        "question": "In progressive MSA, how does the SP-score gain of 2-piece affine gap cost over 1-piece affine vary with the fraction of long-indel events (pi_long), and does the gain survive the greedy progressive propagation?",
        "falsifiable_prediction": (
            "pi_long=0 -> 2-piece SP gain <=0.5pp and sign inconsistent across 3 seeds; "
            "pi_long=0.20 -> 2-piece SP gain >=2pp and same sign across all 3 seeds. "
            "Refuted if pi_long=0.20 gain <1pp or signs disagree, or if pi_long=0 shows >1pp gain."
        ),
        "prediction_outcome": None,
        "negative_result": None,
        "dataset": f"synthetic: {N_TAXA} taxa, root {ROOT_LEN}bp, K2P substitutions + mixed-geometric indels; "
                   f"dev seeds {DEV_SEEDS} x {FAMS_PER_SEED} fam, test seeds {TEST_SEEDS} x {FAMS_PER_SEED} fam, per pi_long in {PI_LONGS}",
        "env": {"python": sys.version.split()[0], "torch": "n/a (C++/no-torch experiment)", "extra_packages": []},
        "seeds": TEST_SEEDS,
        "arms": [],
        "headline": {},
        "deviations": [
            "n_taxa reduced 8->6, root_len 800->220bp, families/seed 10->4, and proposed grid E2 kept at 4 values "
            "(not narrowed) to fit the 45-minute compute budget on shared 4-core CPU; see NOTES.md for full trace.",
            "match/mismatch score rescaled from an initial +2/-1 to +1/-3 and branch lengths shortened "
            "(U(0.03,0.18)->U(0.01,0.05)) after discovering the original scale made naive pairwise alignment "
            "prefer spurious tiny gaps at high sequence divergence (verified against ground truth, see NOTES.md).",
            "sp_near_long_indel is undefined for a (seed, arm) when none of that seed's 4 test families contain "
            "a true gap of length>=10 (this happens legitimately at low pi_long, since long indels are drawn "
            "with probability pi_long and pi_long=0.00 makes them rare-but-not-impossible under the short "
            "geometric tail). In that case the metric falls back to that seed's whole-alignment sp_score, "
            "logged explicitly in logs/run.log as '... falls back to sp_score=...'. Affects pi_long=0.00 (all "
            "3 seeds) and pi_long=0.05 (1 of 3 seeds); pi_long=0.20 always had real flank data. This metric is "
            "reported for mechanism inspection only and is not the basis of the headline claim.",
        ],
        "runtime_sec": None,
    }

    per_pi = {}

    for pi_long in PI_LONGS:
        pi_tag = f"{pi_long:.2f}"
        print(f"\n=== pi_long={pi_tag} : generating dev/test data ===")
        dev_prefix = f"pi{pi_tag}_dev_seed"
        test_prefix = f"pi{pi_tag}_test_seed"
        for seed in DEV_SEEDS:
            simulate(seed, pi_long, FAMS_PER_SEED, f"pi{pi_tag}_dev_seed{seed}")
        for seed in TEST_SEEDS:
            simulate(seed, pi_long, FAMS_PER_SEED, f"pi{pi_tag}_test_seed{seed}")

        dev_items = [(seed, fam) for seed in DEV_SEEDS for fam in range(FAMS_PER_SEED)]

        def dev_path_prefix(seed):
            return os.path.join(DATA, f"pi{pi_tag}_dev_seed{seed}_")

        def dev_family_paths(seed, fam):
            p = dev_path_prefix(seed)
            return f"{p}fam{fam}_unaligned.fa", f"{p}fam{fam}_true.fa"

        # ---- baseline grid search on dev ----
        print(f"[grid] baseline 1-piece: {len(BASELINE_GRID)} combos x {len(dev_items)} dev families")
        best_base, best_base_sp = None, -1
        for (O, E) in BASELINE_GRID:
            sps = []
            for seed, fam in dev_items:
                un, tr = dev_family_paths(seed, fam)
                names, seqs, _, _ = run_msa(un, 1, [O, E])
                sp, _, _, _ = evaluate_family(tr, names, seqs)
                if sp is not None:
                    sps.append(sp)
            mean_sp = float(np.mean(sps)) if sps else 0.0
            log.append(f"pi={pi_tag} baseline grid O={O} E={E} dev_mean_sp={mean_sp:.4f}")
            if mean_sp > best_base_sp:
                best_base_sp = mean_sp
                best_base = (O, E)
        print(f"[grid] best baseline (O,E)={best_base} dev_sp={best_base_sp:.4f}")

        # ---- proposed grid search on dev (O1,E1 fixed to best baseline) ----
        O1, E1 = best_base
        print(f"[grid] proposed 2-piece: {len(PROPOSED_GRID_O2E2)} combos x {len(dev_items)} dev families "
              f"(O1={O1},E1={E1} fixed)")
        best_prop, best_prop_sp = None, -1
        for (O2, E2) in PROPOSED_GRID_O2E2:
            sps = []
            for seed, fam in dev_items:
                un, tr = dev_family_paths(seed, fam)
                names, seqs, _, _ = run_msa(un, 2, [O1, E1, O2, E2])
                sp, _, _, _ = evaluate_family(tr, names, seqs)
                if sp is not None:
                    sps.append(sp)
            mean_sp = float(np.mean(sps)) if sps else 0.0
            log.append(f"pi={pi_tag} proposed grid O2={O2} E2={E2} dev_mean_sp={mean_sp:.4f}")
            if mean_sp > best_prop_sp:
                best_prop_sp = mean_sp
                best_prop = (O2, E2)
        print(f"[grid] best proposed (O2,E2)={best_prop} dev_sp={best_prop_sp:.4f}")
        O2, E2 = best_prop

        # ---- test evaluation, per seed ----
        metrics_base = {"sp_score": [], "tc_score": [], "sp_near_long_indel": [], "runtime_ms": [], "peak_dp_bytes": []}
        metrics_prop = {"sp_score": [], "tc_score": [], "sp_near_long_indel": [], "runtime_ms": [], "peak_dp_bytes": []}

        for seed in TEST_SEEDS:
            p = os.path.join(DATA, f"pi{pi_tag}_test_seed{seed}_")
            base_sp, base_tc, base_spflank, base_rt, base_pk = [], [], [], [], []
            prop_sp, prop_tc, prop_spflank, prop_rt, prop_pk = [], [], [], [], []
            for fam in range(FAMS_PER_SEED):
                un = f"{p}fam{fam}_unaligned.fa"
                tr = f"{p}fam{fam}_true.fa"
                names, seqs, rt, pk = run_msa(un, 1, [O1, E1])
                sp, tc, spf, has_flank = evaluate_family(tr, names, seqs)
                base_sp.append(sp); base_tc.append(tc)
                if has_flank: base_spflank.append(spf)
                base_rt.append(rt); base_pk.append(pk)

                names, seqs, rt, pk = run_msa(un, 2, [O1, E1, O2, E2])
                sp, tc, spf, has_flank = evaluate_family(tr, names, seqs)
                prop_sp.append(sp); prop_tc.append(tc)
                if has_flank: prop_spflank.append(spf)
                prop_rt.append(rt); prop_pk.append(pk)

            metrics_base["sp_score"].append(float(np.mean(base_sp)))
            metrics_base["tc_score"].append(float(np.mean(base_tc)))
            if base_spflank:
                metrics_base["sp_near_long_indel"].append(float(np.mean(base_spflank)))
            else:
                fallback = float(np.mean(base_sp))
                log.append(f"pi={pi_tag} seed={seed} baseline: no true gap>=10 in any of "
                            f"{FAMS_PER_SEED} test families -> sp_near_long_indel falls back to sp_score={fallback:.4f}")
                metrics_base["sp_near_long_indel"].append(fallback)
            metrics_base["runtime_ms"].append(float(np.mean(base_rt)))
            metrics_base["peak_dp_bytes"].append(float(np.mean(base_pk)))

            metrics_prop["sp_score"].append(float(np.mean(prop_sp)))
            metrics_prop["tc_score"].append(float(np.mean(prop_tc)))
            if prop_spflank:
                metrics_prop["sp_near_long_indel"].append(float(np.mean(prop_spflank)))
            else:
                fallback = float(np.mean(prop_sp))
                log.append(f"pi={pi_tag} seed={seed} proposed: no true gap>=10 in any of "
                            f"{FAMS_PER_SEED} test families -> sp_near_long_indel falls back to sp_score={fallback:.4f}")
                metrics_prop["sp_near_long_indel"].append(fallback)
            metrics_prop["runtime_ms"].append(float(np.mean(prop_rt)))
            metrics_prop["peak_dp_bytes"].append(float(np.mean(prop_pk)))

            log.append(f"pi={pi_tag} seed={seed} baseline sp={metrics_base['sp_score'][-1]:.4f} "
                        f"proposed sp={metrics_prop['sp_score'][-1]:.4f}")

        def arm_metrics(d):
            out = {}
            for k, vals in d.items():
                if any(v is None for v in vals):
                    out[k] = {"per_seed": vals, "mean": None, "std": None}
                else:
                    out[k] = {"per_seed": vals, "mean": float(np.mean(vals)), "std": float(np.std(vals))}
            return out

        arm_base = {
            "name": f"baseline-affine-1p_pi{pi_tag}",
            "is_baseline": True,
            "pi_long": pi_long,
            "params": {"O": O1, "E": E1},
            "what": f"1-piece affine gap g(l)=O+l*E, (O,E)=({O1},{E1}) selected by dev grid search over "
                    f"{len(BASELINE_GRID)} combos at pi_long={pi_tag}.",
            "metrics": arm_metrics(metrics_base),
        }
        arm_prop = {
            "name": f"proposed-affine-2p_pi{pi_tag}",
            "is_baseline": False,
            "pi_long": pi_long,
            "params": {"O1": O1, "E1": E1, "O2": O2, "E2": E2},
            "what": f"2-piece affine gap g(l)=min(O1+l*E1, O2+l*E2), (O1,E1) fixed from baseline, "
                    f"(O2,E2)=({O2},{E2}) selected by dev grid search over {len(PROPOSED_GRID_O2E2)} combos "
                    f"at pi_long={pi_tag}.",
            "metrics": arm_metrics(metrics_prop),
        }
        results["arms"].append(arm_base)
        results["arms"].append(arm_prop)
        per_pi[pi_long] = {
            "base_sp_mean": arm_base["metrics"]["sp_score"]["mean"],
            "prop_sp_mean": arm_prop["metrics"]["sp_score"]["mean"],
            "base_sp_per_seed": arm_base["metrics"]["sp_score"]["per_seed"],
            "prop_sp_per_seed": arm_prop["metrics"]["sp_score"]["per_seed"],
        }

    # ---- headline & prediction check ----
    d0 = per_pi[0.00]
    d20 = per_pi[0.20]
    gain0_pp = (d0["prop_sp_mean"] - d0["base_sp_mean"]) * 100
    gain20_pp = (d20["prop_sp_mean"] - d20["base_sp_mean"]) * 100
    signs20 = [np.sign(p - b) for p, b in zip(d20["prop_sp_per_seed"], d20["base_sp_per_seed"])]
    signs0 = [np.sign(p - b) for p, b in zip(d0["prop_sp_per_seed"], d0["base_sp_per_seed"])]
    consistent20 = len(set(signs20)) == 1 and signs20[0] != 0
    consistent0 = len(set(signs0)) == 1 and signs0[0] != 0

    refuted = (gain20_pp < 1.0) or (not consistent20) or (gain0_pp > 1.0)
    if refuted:
        outcome = "refuted"
    elif gain0_pp <= 0.5 and (not consistent0) and gain20_pp >= 2.0 and consistent20:
        outcome = "confirmed"
    else:
        outcome = "inconclusive"

    results["prediction_outcome"] = outcome
    results["negative_result"] = (outcome != "confirmed")
    results["headline"] = {
        "metric": "sp_score",
        "baseline_mean": d20["base_sp_mean"],
        "proposed_mean": d20["prop_sp_mean"],
        "delta": d20["prop_sp_mean"] - d20["base_sp_mean"],
        "claim": (
            f"At pi_long=0.20, 2-piece affine gains {gain20_pp:.2f}pp SP over 1-piece affine "
            f"(seed signs consistent={consistent20}); at pi_long=0.00 the gain is {gain0_pp:.2f}pp "
            f"(seed signs consistent={consistent0}). Prediction {outcome}."
        ),
    }
    results["runtime_sec"] = round(time.time() - t_start, 1)

    with open(os.path.join(ROOT, "results.json"), "w") as f:
        json.dump(results, f, indent=2)

    print("\n".join(log))
    print("\n=== SUMMARY ===")
    for pl in PI_LONGS:
        d = per_pi[pl]
        print(f"pi_long={pl:.2f}  baseline_sp={d['base_sp_mean']:.4f}  proposed_sp={d['prop_sp_mean']:.4f}  "
              f"gain_pp={(d['prop_sp_mean']-d['base_sp_mean'])*100:.2f}")
    print(f"prediction_outcome={outcome}  runtime_sec={results['runtime_sec']}")

    # ---- figure ----
    fig, axes = plt.subplots(1, 2, figsize=(11, 4.5))
    xs = PI_LONGS
    base_means = [per_pi[p]["base_sp_mean"] for p in xs]
    prop_means = [per_pi[p]["prop_sp_mean"] for p in xs]
    base_stds = [float(np.std(per_pi[p]["base_sp_per_seed"])) for p in xs]
    prop_stds = [float(np.std(per_pi[p]["prop_sp_per_seed"])) for p in xs]
    ax = axes[0]
    ax.errorbar(xs, base_means, yerr=base_stds, marker="o", label="baseline 1-piece", capsize=3)
    ax.errorbar(xs, prop_means, yerr=prop_stds, marker="s", label="proposed 2-piece", capsize=3)
    ax.set_xlabel("pi_long (fraction of long-indel events)")
    ax.set_ylabel("SP score (test set, mean +/- std over 3 seeds)")
    ax.set_title("SP score vs long-indel fraction")
    ax.legend()
    ax.grid(alpha=0.3)

    ax2 = axes[1]
    gains_pp = [(per_pi[p]["prop_sp_mean"] - per_pi[p]["base_sp_mean"]) * 100 for p in xs]
    ax2.bar([str(p) for p in xs], gains_pp, color=["#888" if g < 1 else "#2a7" for g in gains_pp])
    ax2.axhline(0, color="black", lw=0.8)
    ax2.set_xlabel("pi_long")
    ax2.set_ylabel("2-piece SP gain over 1-piece (pp)")
    ax2.set_title("Gain vs long-indel fraction")
    ax2.grid(alpha=0.3, axis="y")

    fig.tight_layout()
    fig.savefig(os.path.join(FIGS, "sp_gain_vs_pi_long.png"), dpi=150)
    print("wrote figs/sp_gain_vs_pi_long.png")


if __name__ == "__main__":
    main()
