#!/usr/bin/env python3
"""Derive every number the page prints from the source files. Python 3.9+, numpy.

Inputs (never published):
  GRDC_DIR/6742200_Q_Day.Cmd.txt   Orsova daily mean discharge 1840-2022   (GRDC, owner INHGA)
  GRDC_DIR/6742201_Q_Day.Cmd.txt   Bazias daily mean discharge 1976-2022
  GRDC_DIR/6742500_Q_Day.Cmd.txt   Zimnicea 1931-2022
  GRDC_DIR/6742900_Q_Day.Cmd.txt   Ceatal Izmail 1921-2021
Inputs (public, in the project):
  research/inhga-bazias-observed-2026.csv   INHGA bulletin values, 1 Jun - 2 Oct 2026
  research/eurostat-ro-hydro.json           Eurostat nrg_ind_peh, RO, hydro, gross electricity, GWh
Output:
  research/record.json      derived statistics only (build.mjs publishes a hashed copy under assets/). GRDC terms forbid redistributing the data
                            themselves, so no daily series leaves this script: only per-year and
                            per-month statistics, threshold counts and a short list of the lowest days.
Usage:  GRDC_DIR=/path/to/grdc-raw python3 scripts/compute.py [--check]
"""
import csv, datetime as dt, json, math, os, statistics as st, sys
import numpy as np

ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
GRDC = os.environ.get("GRDC_DIR") or os.path.join(ROOT, "..", "_automation", "danube-water-energy-upgrade", "grdc-raw")
FILES = {
    "orsova": "z1/6742200_Q_Day.Cmd.txt",
    "bazias": "z2/6742201_Q_Day.Cmd.txt",
    "zimnicea": "z3/6742500_Q_Day.Cmd.txt",
    "ceatal": "z3/6742900_Q_Day.Cmd.txt",
}
GRID_LO, GRID_HI, GRID_STEP = 1000, 3000, 10
WINDOWS = {"all": range(1, 13), "junsep": range(6, 10), "julaug": range(7, 9)}
WINDOWS.update({f"m{m:02d}": [m] for m in range(1, 13)})
THRESH = (1500, 1800, 2000, 2500)


def load(path):
    d = {}
    for line in open(path, encoding="latin-1"):
        if line[0] == "#" or line.startswith("YYYY"):
            continue
        a = line.strip().split(";")
        v = float(a[2])
        d[dt.date.fromisoformat(a[0])] = None if v <= -998 else v
    return d


def corr(x, y):
    return float(np.corrcoef(x, y)[0, 1])


def mk_tau_p(y):
    """Kendall tau with normal-approximation p (no tie correction beyond Sen's); y evenly spaced."""
    n = len(y)
    s = sum(np.sign(y[j] - y[i]) for i in range(n) for j in range(i + 1, n))
    var = n * (n - 1) * (2 * n + 5) / 18
    z = (s - np.sign(s)) / math.sqrt(var) if s else 0.0
    p = math.erfc(abs(z) / math.sqrt(2))
    return float(s) / (n * (n - 1) / 2), p


def stats_for(d, first_year, last_year):
    """Per-year, per-window and calendar-day statistics from one daily series."""
    by_year = {}
    for day, v in d.items():
        if v is not None:
            by_year.setdefault(day.year, []).append((day, v))
    years = [y for y in range(first_year, last_year + 1) if y in by_year]
    out = {"years": [], "windows": {}}
    for y in years:
        vals = [v for _, v in by_year[y]]
        mn = min(by_year[y], key=lambda t: t[1])
        mon_min = []
        for m in range(1, 13):
            mv = [v for day, v in by_year[y] if day.month == m]
            mon_min.append(int(min(mv)) if mv else None)
        out["years"].append({
            "y": y, "n": len(vals), "mean": round(st.mean(vals)), "min": int(mn[1]), "minDate": mn[0].strftime("%m-%d"),
            "mon": mon_min, **{f"d{t}": sum(1 for v in vals if v < t) for t in THRESH},
        })
    grid = list(range(GRID_LO, GRID_HI + 1, GRID_STEP))
    for wid, months in WINDOWS.items():
        months = set(months)
        vals = [v for day, v in d.items() if v is not None and day.month in months]
        arr = np.array(vals)
        cum = [int((arr <= g).sum()) for g in grid]
        ymin = {}
        for y in years:
            wv = [v for day, v in by_year[y] if day.month in months]
            ymin[y] = int(min(wv)) if wv else None
        lowest = min(((v, day) for day, v in d.items() if v is not None and day.month in months), key=lambda t: t[0])
        out["windows"][wid] = {"n": len(vals), "cum": cum, "ymin": [ymin[y] for y in years], "lowest": [int(lowest[0]), lowest[1].isoformat()]}
    out["grid"] = [GRID_LO, GRID_HI, GRID_STEP]
    out["yearList"] = years
    return out, by_year


def lowest_days(d, n, months=None):
    items = [(v, day) for day, v in d.items() if v is not None and (months is None or day.month in months)]
    items.sort(key=lambda t: (t[0], t[1]))
    return [[int(v), day.isoformat()] for v, day in items[:n]]


def calendar_envelope(by_year, half=3):
    """For each calendar day: lowest, 10th percentile and median flow over all years, within +-half days."""
    env = []
    base = dt.date(2001, 1, 1)  # non-leap anchor; 29 Feb folded into 28 Feb/1 Mar neighbours by +-half
    table = {}
    for y, rows in by_year.items():
        for day, v in rows:
            doy = (dt.date(2001, day.month, min(day.day, 28 if day.month == 2 else day.day)) - base).days
            table.setdefault(doy, []).append(v)
    for doy in range(365):
        pool = []
        for k in range(doy - half, doy + half + 1):
            pool.extend(table.get(k % 365, []))
        a = np.array(pool)
        env.append([int(a.min()), int(np.percentile(a, 10)), int(np.median(a))])
    return env


def main():
    D = {k: load(os.path.join(GRDC, v)) for k, v in FILES.items()}
    o, b = D["orsova"], D["bazias"]
    rec = {"generator": "scripts/compute.py", "stations": {}}

    # ---- station blocks
    for key, (fy, ly) in {"orsova": (1840, 2022), "bazias": (1976, 2022)}.items():
        s, _ = stats_for(D[key], fy, ly)
        s["lowest"] = lowest_days(D[key], 12)
        s["lowestJunSep"] = lowest_days(D[key], 8, set(range(6, 10)))
        s["lowestOct"] = lowest_days(D[key], 6, {10})
        s["monthRecord"] = [lowest_days(D[key], 1, {m})[0] for m in range(1, 13)]
        s["missingDays"] = sum(1 for v in D[key].values() if v is None)
        rec["stations"][key] = s
    _, by_year_o = stats_for(o, 1840, 2022)
    rec["stations"]["orsova"]["envelope"] = calendar_envelope(by_year_o)
    _, by_year_b = stats_for(b, 1976, 2022)
    rec["stations"]["bazias"]["envelope"] = calendar_envelope(by_year_b)

    # ---- cold-season share of the lowest days
    low1300 = [day for day, v in o.items() if v is not None and v <= 1300]
    rec["orsova1300"] = {"days": len(low1300), "months": sorted({d.month for d in low1300}), "years": sorted({d.year for d in low1300})}

    # ---- two gauges
    both = [day for day in b if b[day] is not None and o.get(day) is not None]
    bx = np.array([b[d] for d in both]); ox = np.array([o[d] for d in both])
    lag1 = [(b[d], o[d + dt.timedelta(1)]) for d in both if o.get(d + dt.timedelta(1)) is not None]
    rec["cross"] = {
        "days": len(both), "identical": int((bx == ox).sum()), "ratioMean": round(float(np.mean(ox / bx)), 4),
        "corr": round(corr(bx, ox), 4), "corrLag1": round(corr([x for x, _ in lag1], [y for _, y in lag1]), 4),
        "medianAbsDiff": float(np.median(np.abs(bx - ox))), "p95AbsDiff": float(np.percentile(np.abs(bx - ox), 95)),
        "medianAbsDiffPct": round(float(np.median(np.abs(bx - ox) / bx * 100)), 1),
        "jan1985": [[d.isoformat(), int(b[d]), int(o[d])] for d in sorted(both) if d.year == 1985 and d.month == 1 and d.day in (13, 19, 20, 21)],
        "annualMinCorr": None,
    }
    yb = {y["y"]: y["min"] for y in rec["stations"]["bazias"]["years"]}
    yo = {y["y"]: y["min"] for y in rec["stations"]["orsova"]["years"]}
    common = sorted(set(yb) & set(yo))
    rec["cross"]["annualMinCorr"] = round(corr([yb[y] for y in common], [yo[y] for y in common]), 3)
    # downstream stations: annual mean in dry years, to show the low water is one river-wide event
    dry = {}
    for key in ("orsova", "zimnicea", "ceatal"):
        s, _ = stats_for(D[key], 1921, 2022)
        dry[key] = {y["y"]: [y["mean"], y["min"]] for y in s["years"]}
    rec["downstream"] = {str(y): {k: dry[k].get(y) for k in dry} for y in (1921, 1947, 1954, 2003, 2011, 2022)}

    # ---- the dam (1971): blocks and tests
    ann = rec["stations"]["orsova"]["years"]
    def block(a, z):
        ys = [x for x in ann if a <= x["y"] <= z]
        return {"from": a, "to": z, "n": len(ys), "meanOfMeans": round(st.mean(x["mean"] for x in ys)), "meanOfMins": round(st.mean(x["min"] for x in ys)),
                "lowestMin": min(x["min"] for x in ys)}
    rec["dam"] = {"before": block(1840, 1970), "after": block(1971, 2022), "from1991": block(1991, 2022)}
    pre = np.array([x["min"] for x in ann if x["y"] <= 1970], float); post = np.array([x["min"] for x in ann if x["y"] >= 1971], float)
    rng = np.random.default_rng(20261002)
    allv = np.concatenate([pre, post]); obs = abs(post.mean() - pre.mean()); cnt = 0
    for _ in range(20000):
        rng.shuffle(allv); cnt += abs(allv[len(pre):].mean() - allv[:len(pre)].mean()) >= obs
    rec["dam"]["permP_annualMin"] = round(cnt / 20000, 3)
    pre_m = np.array([x["mean"] for x in ann if x["y"] <= 1970], float); post_m = np.array([x["mean"] for x in ann if x["y"] >= 1971], float)
    allv = np.concatenate([pre_m, post_m]); obs = abs(post_m.mean() - pre_m.mean()); cnt = 0
    for _ in range(20000):
        rng.shuffle(allv); cnt += abs(allv[len(pre_m):].mean() - allv[:len(pre_m)].mean()) >= obs
    rec["dam"]["permP_annualMean"] = round(cnt / 20000, 3)
    def dstd(a, z):
        xs = []
        for day in sorted(o):
            if a <= day.year <= z and o[day] is not None:
                p = o.get(day - dt.timedelta(1))
                if p: xs.append(math.log(o[day] / p))
        return round(float(np.std(xs)), 4)
    rec["dam"]["dayToDayLogSd"] = {"1940-1970": dstd(1940, 1970), "1972-1989": dstd(1972, 1989), "1990-2022": dstd(1990, 2022)}

    # ---- trends (Kendall, whole record) -- reported with p, not interpreted
    ys = [x["y"] for x in ann]
    junsep_min = rec["stations"]["orsova"]["windows"]["junsep"]["ymin"]
    tau_m, p_m = mk_tau_p([x["mean"] for x in ann]); tau_n, p_n = mk_tau_p(junsep_min); tau_a, p_a = mk_tau_p([x["min"] for x in ann])
    rec["trend"] = {"annualMean": [round(tau_m, 3), round(p_m, 3)], "annualMin": [round(tau_a, 3), round(p_a, 3)], "junSepMin": [round(tau_n, 3), round(p_n, 3)], "n": len(ann)}
    # 1840-1930 vs 1931-2022 mean of the Jun-Sep minimum
    rec["junSepMinByHalf"] = {"1840-1930": round(st.mean(v for y, v in zip(ys, junsep_min) if y <= 1930)), "1931-2022": round(st.mean(v for y, v in zip(ys, junsep_min) if y >= 1931))}

    # ---- the operator's pre-regulation list (Mediafax 17 Aug 2026) against the validated record
    ymin = {x["y"]: (x["min"], x["minDate"]) for x in ann}
    rec["operatorList"] = {str(y): ymin[y][0] for y in (1947, 1902, 1858, 1866)}

    # ---- Eurostat hydro vs Orsova annual mean
    eu = json.load(open(os.path.join(ROOT, "research", "eurostat-ro-hydro.json")))
    tidx = eu["dimension"]["time"]["category"]["index"]
    gwh = {int(y): eu["value"].get(str(i)) for y, i in tidx.items()}  # operator TOTAL is the first block
    mean_by_y = {x["y"]: x["mean"] for x in ann}
    pairs = [(y, gwh[y], mean_by_y[y]) for y in sorted(gwh) if y in mean_by_y and gwh[y] is not None]
    x = np.array([p[2] for p in pairs], float); yv = np.array([p[1] for p in pairs], float)
    slope, icpt = np.polyfit(x, yv, 1)
    order = sorted(pairs, key=lambda p: p[2])
    rec["hydro"] = {
        "pairs": [[p[0], round(p[1]), p[2]] for p in pairs], "r": round(corr(x, yv), 3), "slopeGWhPer1000": round(float(slope) * 1000), "intercept": round(float(icpt)),
        "n": len(pairs), "driest5": [p[0] for p in order[:5]], "wettest5": [p[0] for p in order[-5:]],
        "driest5MeanGWh": round(st.mean(p[1] for p in order[:5])), "wettest5MeanGWh": round(st.mean(p[1] for p in order[-5:])),
        "source": "Eurostat nrg_ind_peh, geo=RO, siec=RA100, nrg_bal=GEP, plants=TOTAL, operator=TOTAL, GWh",
    }

    # ---- 2026 INHGA bulletins against the validated record
    rows = [(r["bulletin_date"], int(r["bazias_m3s"])) for r in csv.DictReader(open(os.path.join(ROOT, "research", "inhga-bazias-observed-2026.csv"))) if r["bazias_m3s"]]
    env_o = rec["stations"]["orsova"]["envelope"]
    month_rec_o = {m: rec["stations"]["orsova"]["monthRecord"][m - 1][0] for m in range(1, 13)}
    month_rec_b = {m: rec["stations"]["bazias"]["monthRecord"][m - 1][0] for m in range(1, 13)}
    series = []
    for dstr, v in rows:
        day = dt.date.fromisoformat(dstr)
        doy = (dt.date(2001, day.month, min(day.day, 28 if day.month == 2 else day.day)) - dt.date(2001, 1, 1)).days
        series.append([dstr, v, env_o[doy][0], env_o[doy][2]])
    rec["y2026"] = {
        "series": series, "n": len(series), "multiannualMean": {"6": 5900, "7": 4700, "8": 3900, "9": 3800, "10": 3900},
        "belowMonthRecordOrsova": {str(m): sum(1 for dstr, v in rows if int(dstr[5:7]) == m and v < month_rec_o[m]) for m in range(6, 11)},
        "belowEnvelope": sum(1 for s in series if s[1] < s[2]),
        "belowRecordBy100": {str(m): sum(1 for dstr, v in rows if int(dstr[5:7]) == m and v <= month_rec_o[m] - 100) for m in range(6, 11)},
        "belowBaziasRecord": {str(m): sum(1 for dstr, v in rows if int(dstr[5:7]) == m and v < month_rec_b[m]) for m in range(6, 11)},
        "daysBelow": {str(t): sum(1 for _, v in rows if v < t) for t in THRESH},
        "first1250": [d for d, v in rows if v == min(r[1] for r in rows)],
        "daysByMonth": {str(m): sum(1 for dstr, _ in rows if int(dstr[5:7]) == m) for m in range(6, 11)},
        "monthRecordOrsova": {str(m): month_rec_o[m] for m in range(6, 11)},
        "monthRecordBazias": {str(m): month_rec_b[m] for m in range(6, 11)},
        "lowest": min(rows, key=lambda r: r[1]), "lastDate": rows[-1][0],
    }

    # ---- sources
    rec["meta"] = {
        "retrieved": "2026-10-02", "owner": "National Institute of Hydrology and Water Management (INHGA), Romania",
        "citation": "The Global Runoff Data Centre, 56068 Koblenz, Germany",
        "orsovaDays": len(o), "bazias": {"first": min(b).isoformat(), "last": max(b).isoformat()},
    }
    out = os.path.join(ROOT, "research", "record.json")
    with open(out, "w") as f:
        json.dump(rec, f, separators=(",", ":"), ensure_ascii=False)
    print("wrote", out, os.path.getsize(out), "bytes")
    if "--check" in sys.argv:
        summarize(rec)


def summarize(rec):
    o = rec["stations"]["orsova"]
    print("lowest", o["lowest"][:4]); print("junsep", o["lowestJunSep"][:4]); print("oct", o["lowestOct"][:3])
    print("monthRecord", o["monthRecord"])
    print("1300 days", rec["orsova1300"]); print("cross", rec["cross"]); print("downstream", rec["downstream"])
    print("dam", rec["dam"]); print("trend", rec["trend"], rec["junSepMinByHalf"]); print("operator", rec["operatorList"])
    print("hydro", {k: v for k, v in rec["hydro"].items() if k != "pairs"})
    y = rec["y2026"]; print("2026", {k: v for k, v in y.items() if k != "series"})
    b = rec["stations"]["bazias"]; print("bazias lowest", b["lowest"][:4], b["lowestJunSep"][:3], b["monthRecord"][5:10])


if __name__ == "__main__":
    main()
