#!/usr/bin/env python3
"""Run the locked Genesis pedigree analysis.

The calculations are conditional on treating the supplied ages as phenotypic
observations.  They do not test the historicity of the named events.
"""

from __future__ import annotations

import json
import math
from pathlib import Path

import numpy as np
import pandas as pd
from scipy.optimize import minimize, least_squares
from scipy.special import gammaln, logsumexp
from scipy.stats import norm, f as f_dist


ROOT = Path(__file__).resolve().parent
RESULTS = ROOT / "results"
RESULTS.mkdir(exist_ok=True)

NAMES = np.array(
    ["Shem", "Arphaxad", "Cainan", "Shelah", "Eber", "Peleg", "Reu", "Serug", "Nahor", "Terah"]
)
FATHER_AGES = np.array([100, 135, 130, 130, 134, 130, 132, 130, 79, 130], dtype=float)
LIFE_PRIMARY = np.array([600, 565, 460, 533, 504, 339, 339, 330, 208, 205], dtype=float)
LIFE_ALT = np.array([600, 565, 460, 460, 504, 339, 339, 330, 208, 205], dtype=float)
AMBIENTS = [40, 60, 80, 100]
N_GRID = [2, 3, 5, 8, 10, 15, 20, 30, 50, 75, 100, 200, 500]


def transition_matrix(n: int, retention: float) -> np.ndarray:
    """P(child k | parent j) for irreversible retention of j active factors."""
    retention = float(np.clip(retention, 1e-10, 1 - 1e-10))
    out = np.zeros((n + 1, n + 1), dtype=float)
    for j in range(n + 1):
        k = np.arange(j + 1)
        logc = gammaln(j + 1) - gammaln(k + 1) - gammaln(j - k + 1)
        lp = logc + k * math.log(retention) + (j - k) * math.log1p(-retention)
        out[j, : j + 1] = np.exp(lp)
    return out


def emission(y: float, means: np.ndarray, sigma: float) -> np.ndarray:
    return np.maximum(norm.pdf(y, loc=means, scale=sigma), 1e-300)


def train_filter(y: np.ndarray, ambient: float, n: int, retention: float, sigma: float):
    """Condition on Shem as K=N and fit observations Arphaxad..Serug."""
    d = y[0] - ambient
    means = ambient + d * np.arange(n + 1) / n
    trans = transition_matrix(n, retention)
    alpha = np.zeros(n + 1)
    alpha[n] = 1.0
    loglik = 0.0
    # Shem fixes the scale and is not reused as a zero-residual observation.
    for obs in y[1:8]:
        alpha = alpha @ trans
        alpha *= emission(obs, means, sigma)
        z = alpha.sum()
        if not np.isfinite(z) or z <= 0:
            return -np.inf, alpha
        loglik += math.log(z)
        alpha /= z
    return loglik, alpha


def fit_training(y: np.ndarray, ambient: float, n: int):
    def objective(theta):
        retention = 0.5 + 0.4999 / (1 + math.exp(-float(theta[0])))
        sigma = math.exp(float(theta[1]))
        ll, _ = train_filter(y, ambient, n, retention, sigma)
        return -ll if np.isfinite(ll) else 1e100

    starts = [
        [1.5, math.log(15)],
        [2.2, math.log(35)],
        [3.0, math.log(60)],
        [0.5, math.log(30)],
    ]
    best = None
    for start in starts:
        z = minimize(objective, start, method="Nelder-Mead", options={"maxiter": 1000, "xatol": 1e-8})
        if best is None or z.fun < best.fun:
            best = z
    retention = 0.5 + 0.4999 / (1 + math.exp(-float(best.x[0])))
    sigma = math.exp(float(best.x[1]))
    ll, alpha = train_filter(y, ambient, n, retention, sigma)
    return {"retention": retention, "sigma": sigma, "train_loglik": ll, "posterior_k7": alpha}


def heldout_logscore(
    y: np.ndarray,
    ambient: float,
    n: int,
    fit: dict,
    h8: float,
    h9: float,
):
    d = y[0] - ambient
    means = ambient + d * np.arange(n + 1) / n
    alpha7 = fit["posterior_k7"]
    sigma = fit["sigma"]
    a8_prior = alpha7 @ transition_matrix(n, h8)
    e8 = emission(y[8], means, sigma)
    p8 = float(a8_prior @ e8)
    a8 = a8_prior * e8
    a8 /= max(a8.sum(), 1e-300)
    a9_prior = a8 @ transition_matrix(n, h9)
    p9_cond = float(a9_prior @ emission(y[9], means, sigma))

    pred8_mean = float(a8_prior @ means)
    pred8_sd = float(np.sqrt(a8_prior @ ((means - pred8_mean) ** 2) + sigma**2))
    pred8_lower_tail = float(a8_prior @ norm.cdf(y[8], loc=means, scale=sigma))
    pred9_prior_mean = float(a9_prior @ means)
    pred9_prior_sd = float(np.sqrt(a9_prior @ ((means - pred9_prior_mean) ** 2) + sigma**2))
    return {
        "logscore_nahor": math.log(max(p8, 1e-300)),
        "logscore_joint": math.log(max(p8 * p9_cond, 1e-300)),
        "pred_nahor_mean": pred8_mean,
        "pred_nahor_sd": pred8_sd,
        "P_Nahor_at_or_below_observed": pred8_lower_tail,
        "pred_terah_cond_mean": pred9_prior_mean,
        "pred_terah_cond_sd": pred9_prior_sd,
    }


def run_real_series(label: str, y: np.ndarray):
    profiles = []
    predictions = []
    for ambient in AMBIENTS:
        for n in N_GRID:
            fit = fit_training(y, ambient, n)
            profiles.append(
                {
                    "series": label,
                    "ambient": ambient,
                    "N": n,
                    "retention_pre": fit["retention"],
                    "outloss_pre": 1 - fit["retention"],
                    "sigma_non_genetic": fit["sigma"],
                    "train_loglik": fit["train_loglik"],
                    "posterior_mean_K7": float(fit["posterior_k7"] @ np.arange(n + 1)),
                }
            )
            models = {
                "continuation": (fit["retention"], fit["retention"]),
                "babel_pulse_0.5": (0.5, fit["retention"]),
                "babel_persistent_0.5": (0.5, 0.5),
            }
            for model, (h8, h9) in models.items():
                score = heldout_logscore(y, ambient, n, fit, h8, h9)
                predictions.append(
                    {
                        "series": label,
                        "ambient": ambient,
                        "N": n,
                        "model": model,
                        "h8": h8,
                        "h9": h9,
                        **score,
                    }
                )
            for h in np.arange(0.35, 0.751, 0.05):
                score = heldout_logscore(y, ambient, n, fit, float(h), fit["retention"])
                predictions.append(
                    {
                        "series": label,
                        "ambient": ambient,
                        "N": n,
                        "model": f"babel_pulse_sensitivity_{h:.2f}",
                        "h8": h,
                        "h9": fit["retention"],
                        **score,
                    }
                )
    return pd.DataFrame(profiles), pd.DataFrame(predictions)


def half_excess_table(y: np.ndarray, label: str):
    rows = []
    for g in range(1, len(y)):
        implied_a = 2 * y[g] - y[g - 1]
        rows.append(
            {
                "series": label,
                "boundary": f"{NAMES[g-1]}->{NAMES[g]}",
                "child_generation": g,
                "parent_lifespan": y[g - 1],
                "child_lifespan": y[g],
                "ambient_implied_by_half_outcross": implied_a,
                "inside_fixed_40_100_range": 40 <= implied_a <= 100,
                "distance_from_80": abs(implied_a - 80),
            }
        )
    return pd.DataFrame(rows)


def maturation_table():
    rows = []
    for m in range(1, 7):
        for q in np.arange(0.1, 1.01, 0.1):
            rows.append(
                {
                    "M_recessive_factors": m,
                    "maternal_transmission_q": q,
                    "P_Terah_immediate_rebound": (q / 2) ** m,
                }
            )
    return pd.DataFrame(rows)


def shock_scan(y: np.ndarray, label: str):
    """Exploratory smooth exponential versus one multiplicative shock."""
    rows = []
    g = np.arange(len(y), dtype=float)
    for ambient in AMBIENTS:
        def resid0(theta):
            b, loglam = theta
            return y - (ambient + math.exp(b) * np.exp(-math.exp(loglam) * g))

        z0 = least_squares(resid0, [math.log(600 - ambient), math.log(0.1)])
        rss0 = float(np.sum(resid0(z0.x) ** 2))
        for boundary in range(1, 10):
            def resid1(theta):
                b, loglam, logitshock = theta
                shock = 1 / (1 + math.exp(-logitshock))
                retained = np.exp(-math.exp(loglam) * g) * np.where(g >= boundary, shock, 1.0)
                return y - (ambient + math.exp(b) * retained)

            z1 = least_squares(resid1, [math.log(600 - ambient), math.log(0.1), 1.0])
            shock = 1 / (1 + math.exp(-float(z1.x[2])))
            rss1 = float(np.sum(resid1(z1.x) ** 2))
            fstat = ((rss0 - rss1) / 1) / (rss1 / 7)
            rows.append(
                {
                    "series": label,
                    "ambient": ambient,
                    "boundary_generation": boundary,
                    "boundary": f"{NAMES[boundary-1]}->{NAMES[boundary]}",
                    "rss_smooth": rss0,
                    "rss_shock": rss1,
                    "shock_retention": shock,
                    "F_illustrative": fstat,
                    "p_fixed_boundary_illustrative": 1 - f_dist.cdf(fstat, 1, 7),
                }
            )
    out = pd.DataFrame(rows)
    out["rank_within_ambient"] = out.groupby(["series", "ambient"])["rss_shock"].rank(method="min")
    return out


def summarize(profile: pd.DataFrame, pred: pd.DataFrame):
    rows = []
    for (series, ambient), group in profile.groupby(["series", "ambient"]):
        best_train = group.loc[group.train_loglik.idxmax()]
        p = pred[(pred.series == series) & (pred.ambient == ambient)]
        # Training-likelihood weights across the pre-fixed N grid. These are an
        # empirical-Bayes sensitivity summary, not a formal Bayes factor.
        lls = group.set_index("N").train_loglik
        w = np.exp(lls - lls.max())
        w /= w.sum()
        agg = {}
        tail = {}
        for model in ["continuation", "babel_pulse_0.5", "babel_persistent_0.5"]:
            pm = p[p.model == model].set_index("N")
            # Mixture predictive densities must be averaged on the density scale.
            agg[(model, "nahor")] = math.log(float(np.sum(w * np.exp(pm.logscore_nahor))))
            agg[(model, "joint")] = math.log(float(np.sum(w * np.exp(pm.logscore_joint))))
            tail[model] = float(np.sum(w * pm.P_Nahor_at_or_below_observed))
        rows.append(
            {
                "series": series,
                "ambient": ambient,
                "best_training_N": int(best_train.N),
                "best_retention_pre": best_train.retention_pre,
                "best_sigma": best_train.sigma_non_genetic,
                "logscore_nahor_continuation": agg[("continuation", "nahor")],
                "logscore_nahor_babel_pulse": agg[("babel_pulse_0.5", "nahor")],
                "log_predictive_ratio_nahor_pulse_vs_continuation": agg[("babel_pulse_0.5", "nahor")] - agg[("continuation", "nahor")],
                "predictive_ratio_nahor_pulse_vs_continuation": math.exp(agg[("babel_pulse_0.5", "nahor")] - agg[("continuation", "nahor")]),
                "P_Nahor_at_or_below_observed_under_continuation": tail["continuation"],
                "P_Nahor_at_or_below_observed_under_babel_pulse": tail["babel_pulse_0.5"],
                "logscore_joint_continuation": agg[("continuation", "joint")],
                "logscore_joint_babel_pulse": agg[("babel_pulse_0.5", "joint")],
                "logscore_joint_babel_persistent": agg[("babel_persistent_0.5", "joint")],
                "predictive_ratio_joint_pulse_vs_continuation": math.exp(agg[("babel_pulse_0.5", "joint")] - agg[("continuation", "joint")]),
                "predictive_ratio_joint_persistent_vs_continuation": math.exp(agg[("babel_persistent_0.5", "joint")] - agg[("continuation", "joint")]),
            }
        )
    return pd.DataFrame(rows)


def main():
    all_profiles = []
    all_predictions = []
    all_halves = []
    all_scans = []
    for label, y in [("Shelah_533_primary", LIFE_PRIMARY), ("Shelah_460_sensitivity", LIFE_ALT)]:
        profile, pred = run_real_series(label, y)
        all_profiles.append(profile)
        all_predictions.append(pred)
        all_halves.append(half_excess_table(y, label))
        all_scans.append(shock_scan(y, label))

    profile = pd.concat(all_profiles, ignore_index=True)
    pred = pd.concat(all_predictions, ignore_index=True)
    halves = pd.concat(all_halves, ignore_index=True)
    scans = pd.concat(all_scans, ignore_index=True)
    maturation = maturation_table()
    summary = summarize(profile, pred)

    profile.to_csv(RESULTS / "training_likelihood_profiles.csv", index=False)
    pred.to_csv(RESULTS / "heldout_predictive_scores.csv", index=False)
    halves.to_csv(RESULTS / "half_excess_boundaries.csv", index=False)
    scans.to_csv(RESULTS / "exploratory_shock_scan.csv", index=False)
    maturation.to_csv(RESULTS / "maturation_rebound_probabilities.csv", index=False)
    summary.to_csv(RESULTS / "headline_summary.csv", index=False)

    metadata = {
        "primary_lifespans": LIFE_PRIMARY.tolist(),
        "sensitivity_lifespans": LIFE_ALT.tolist(),
        "fathering_ages": FATHER_AGES.tolist(),
        "names": NAMES.tolist(),
        "ambient_values": AMBIENTS,
        "N_grid": N_GRID,
        "training_generations": [0, 7],
        "heldout_generations": [8, 9],
        "babel_primary_boundary": "Serug->Nahor",
    }
    (RESULTS / "run_metadata.json").write_text(json.dumps(metadata, indent=2) + "\n")
    print(summary.to_string(index=False))


if __name__ == "__main__":
    main()
