#!/usr/bin/env python3
"""
03_country_tests.py — for every country (and for the whole catalogue),
does the Moon's nakṣatra at the moment of the earthquake follow the
time-weighted uniform null?

For each country × {raw catalogue, Gardner–Knopoff mainshocks} ×
{27 nakṣatras, 28 with Abhijit}:
  n, chi-square, df, p, Cohen's w, and the power of the test to detect a
  medium departure (w = 0.3) at alpha = 0.05.  A country is "testable" when
  that power is at least 0.8 (n >= 258 for 27 bins, 262 for 28).  No
  multiple-comparison correction is applied; the essay states instead how
  many false alarms the number of tests would produce by chance.
Also, for each country, one-sided binomial tests for each of the four
circles of Bṛhat Saṁhitā 32 (Vāyu, Agni, Indra, Varuṇa).

Input : catalog_annotated.csv.gz
Output: country_results.csv, country_data.json (counts for the essay page)
"""
import argparse
import json

import numpy as np
import pandas as pd

import bhukampa_common as bc

W_EFFECT = 0.3
ALPHA = 0.05
POWER_GATE = 0.8


def bins(df, col, k):
    return np.bincount(df[col].to_numpy(), minlength=k)[:k]


def analyse(sub, p27, p28):
    out = {}
    for tag, d in (("raw", sub), ("main", sub[sub.mainshock])):
        c27 = bins(d, "nak27_i", 27)
        c28 = bins(d, "nak28_i", 28)
        r27 = bc.chisq_gof(c27, p27)
        r28 = bc.chisq_gof(c28, p28)
        out[tag] = {
            "n": int(len(d)),
            "counts27": c27.tolist(), "counts28": c28.tolist(),
            "chi2_27": r27["chi2"], "p27": r27["p"], "w27": r27["w"],
            "chi2_28": r28["chi2"], "p28": r28["p"], "w28": r28["w"],
            "power27": bc.chisq_power(len(d), 26, W_EFFECT, ALPHA),
            "power28": bc.chisq_power(len(d), 27, W_EFFECT, ALPHA),
            "circles": {c: bc.circle_test(d["nak28_i"].to_numpy(), p28, c)
                        for c in bc.MANDALA},
        }
    return out


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--catalog", default="catalog_annotated.csv.gz")
    ap.add_argument("--csv", default="country_results.csv")
    ap.add_argument("--json", default="country_data.json")
    ap.add_argument("--min-n", type=int, default=50,
                    help="skip countries with fewer raw events than this")
    a = ap.parse_args()

    df = pd.read_csv(a.catalog)
    df["nak27_i"] = df["nak27"].map({n: i for i, n in enumerate(bc.NAK27)})
    df["nak28_i"] = df["nak28"].map({n: i for i, n in enumerate(bc.NAK28)})
    p27, p28 = bc.time_weighted_expectation()

    results = {"_expected": {"p27": p27.tolist(), "p28": p28.tolist(),
                             "w_effect": W_EFFECT, "alpha": ALPHA,
                             "power_gate": POWER_GATE,
                             "n_gate27": bc.n_for_power(26, W_EFFECT, ALPHA, POWER_GATE),
                             "n_gate28": bc.n_for_power(27, W_EFFECT, ALPHA, POWER_GATE)}}
    results["World"] = analyse(df, p27, p28)
    results["World (on land only)"] = analyse(df[df.onland], p27, p28)

    rows = []
    for country, sub in df.groupby("country"):
        if country == "(open ocean)" or len(sub) < a.min_n:
            continue
        res = analyse(sub, p27, p28)
        results[country] = res
        lat, lon = sub.latitude.median(), sub.longitude.median()
        for tag in ("raw", "main"):
            r = res[tag]
            rows.append({"country": country, "catalog": tag, "n": r["n"],
                         "lat": lat, "lon": lon,
                         "chi2_27": r["chi2_27"], "p27": r["p27"], "w27": r["w27"],
                         "power27": r["power27"],
                         "chi2_28": r["chi2_28"], "p28": r["p28"], "w28": r["w28"],
                         "power28": r["power28"],
                         "varuna_k": r["circles"]["Varuṇa"]["k"],
                         "varuna_ratio": r["circles"]["Varuṇa"]["ratio"],
                         "varuna_p": r["circles"]["Varuṇa"]["p_greater"]})
    tab = pd.DataFrame(rows)
    tab["testable27"] = tab.power27 >= POWER_GATE
    tab["testable28"] = tab.power28 >= POWER_GATE
    tab = tab.sort_values(["catalog", "p27"])
    tab.to_csv(a.csv, index=False, float_format="%.6g")

    # centroids into the json too
    for _, r in tab.iterrows():
        d = results[r.country]
        d.setdefault("lat", float(r.lat)); d.setdefault("lon", float(r.lon))
    with open(a.json, "w", encoding="utf-8") as fh:
        json.dump(results, fh, ensure_ascii=False, separators=(",", ":"))

    for tag in ("raw", "main"):
        t = tab[(tab.catalog == tag) & tab.testable27]
        print(f"\n== {tag}: {len(t)} testable countries; "
              f"{(t.p27 < .05).sum()} with p<0.05 (chance alone: {len(t)/20:.1f})")
        print(t[["country", "n", "p27", "p28", "varuna_ratio", "varuna_p"]]
              .head(15).to_string(index=False))


if __name__ == "__main__":
    main()
