"""Are our far-range references converged? Coverage of the stitched mosaic by distance ring as a function of the
number of rays and of the interaction depth, for Sionna RT and for our tracer.

Neither tool has a distance limit, but both can miss long routes: Sionna finds paths by shooting rays (samples per
source), our tracer stops after max_int interactions and uses a fixed ray / diffraction-fan count. If coverage at
200-800 m keeps growing with rays or depth, the far-field numbers of 06-cloudrf.md are too low.

  python -m rcm_ml.far_convergence select --scene runs/cloudrf/mosaic_2km_anchor --out runs/convergence
  .venv-sionna/bin/python -m rcm_ml.sionna_paths --mosaic 8 --tx-rc R,C --rx-file runs/convergence/rx.npy ...  (scripts/far_convergence.sh)
  python -m rcm_ml.far_convergence tracer --scene runs/cloudrf/mosaic_2km_anchor --out runs/convergence
  python -m rcm_ml.far_convergence report --out runs/convergence --fig docs/research/figures/08_far_convergence.png
  python -m rcm_ml.far_convergence winprop --out runs/convergence   (WinProp DPM / IRT2 / IRT4 coverage by distance)
"""
from __future__ import annotations

import argparse
import json
import time
from pathlib import Path

import numpy as np

from .bigmap import mosaic
from .data import DB_RANGE, P_TRNC, ROOT, THRESH, load_gain, load_map, load_tx
from .raytrace import radio_map
from .train import git_commit

RINGS = [(0, 100), (100, 200), (200, 400), (400, 800)]
PER_RING = [50, 100, 100, 100]
FLOOR_DB = P_TRNC + THRESH * DB_RANGE
# (name, max_int, ray multiplier of 7200 * k, n_fan)
TRACER_CFGS = [("int2", 2, 1, 1440), ("int2_rays4x", 2, 4, 5760), ("int3", 3, 1, 1440), ("int4", 4, 1, 1440)]


def ring_of(d):
    for j, (lo, hi) in enumerate(RINGS):
        if lo <= d < hi:
            return j
    return -1


def select(a):
    scene = json.loads(Path(a.scene, "scene.json").read_text())
    bld, _, _, _ = mosaic(scene["k"])
    tx = scene["tx_rc"]
    n = bld.shape[0]
    rr, cc = np.mgrid[:n, :n]
    d = np.hypot(rr - tx[0], cc - tx[1])
    rng = np.random.default_rng(0)
    pix = []
    for (lo, hi), k in zip(RINGS, PER_RING):
        cand = np.argwhere(~bld & (d >= lo) & (d < hi))
        pix.append(cand[rng.choice(len(cand), k, replace=False)])
    pix = np.concatenate(pix)
    out = Path(a.out)
    out.mkdir(parents=True, exist_ok=True)
    np.save(out / "rx.npy", pix)
    (out / "rx.json").write_text(json.dumps({"commit": git_commit(), "scene": a.scene, "tx_rc": tx, "rings": RINGS, "per_ring": PER_RING, "seed": 0}))
    print(len(pix), "receivers")


def tracer(a):
    scene = json.loads(Path(a.scene, "scene.json").read_text())
    _, polys, _, _ = mosaic(scene["k"])
    tx = tuple(scene["tx_rc"])
    n = scene["size_m"]
    cal = json.loads(Path("runs/mps/raytrace_3_calibrated.json").read_text())
    pix = np.load(Path(a.out, "rx.npy"))
    res = {"commit": git_commit(), "configs": {}}
    vals = {}
    for name, mi, mult, fan in TRACER_CFGS:
        t0 = time.time()
        g, _ = radio_map(polys, tx, n_rays=7200 * scene["k"] * mult, n_fan=fan, max_int=mi, ple=cal["path_loss_exponent"],
                         dslope=float(cal["calibration"]["best"]["dslope"]), rloss=float(cal["calibration"]["best"]["rloss"]), n=n)
        dt = time.time() - t0
        db = np.where(g > 0, P_TRNC + DB_RANGE * g, np.nan)[pix[:, 0], pix[:, 1]]
        vals[name] = db.astype(np.float32)
        res["configs"][name] = {"max_int": mi, "n_rays": 7200 * scene["k"] * mult, "n_fan": fan, "seconds": round(dt, 1)}
        print(name, round(dt, 1), "s", flush=True)
    np.savez_compressed(Path(a.out, "tracer.npz"), **vals)
    Path(a.out, "tracer.json").write_text(json.dumps(res, indent=1))


def sionna_db(path, dd=None, ddloss=None):
    z = np.load(path)
    pw = z["pw_cal"].astype(np.float64)
    if dd is not None:
        pw = pw + dd * 10.0 ** (-ddloss / 10.0)
    with np.errstate(divide="ignore"):
        db = 10 * np.log10(pw)
    return np.where(db > P_TRNC, db, np.nan), json.loads(str(z["meta"]))


def stats(db, ring):
    """Per ring: % of receivers above the floor, median dB of those."""
    out = {}
    for j, (lo, hi) in enumerate(RINGS):
        v = db[ring == j]
        cov = np.nan_to_num(v, nan=-999) > FLOOR_DB
        out[f"{lo}-{hi} m"] = {"n": int(len(v)), "covered_pct": round(100 * float(cov.mean()), 1),
                               "median_db_covered": round(float(np.median(v[cov])), 1) if cov.any() else None}
    return out


def report(a):
    out = Path(a.out)
    meta = json.loads((out / "rx.json").read_text())
    pix = np.load(out / "rx.npy")
    tx = meta["tx_rc"]
    ring = np.array([ring_of(float(np.hypot(r - tx[0], c - tx[1]))) for r, c in pix])
    res = {"commit": git_commit(), "receivers": len(pix), "rings": RINGS, "floor_db": round(FLOOR_DB, 1), "tracer": {}, "sionna": {}}
    tz = np.load(out / "tracer.npz")
    tcfg = json.loads((out / "tracer.json").read_text())["configs"]
    for k in tz.files:
        res["tracer"][k] = {**tcfg[k], "rings": stats(tz[k], ring)}
    # Sionna dumps: sionna_<name>.npz; also with our tracer's two-corner paths added, as in the recommended setup
    scene = json.loads(Path(meta["scene"], "scene.json").read_text())
    _, polys, _, _ = mosaic(scene["k"])
    cal = json.loads(Path("runs/mps/raytrace_3_calibrated.json").read_text())
    dd, _ = radio_map(polys, tuple(tx), n_fan=1440, ple=cal["path_loss_exponent"], dslope=float(cal["calibration"]["best"]["dslope"]),
                      rloss=float(cal["calibration"]["best"]["rloss"]), n=scene["size_m"], dd_only=True, power=True)
    ddloss = json.loads(Path("runs/mps/sionna_cfg/scat09add_r100k/result.json").read_text())["best"]["ddloss"]
    for p in sorted(out.glob("sionna_*.npz")):
        name = p.stem.removeprefix("sionna_")
        db, m = sionna_db(p)
        db_dd, _ = sionna_db(p, dd[pix[:, 0], pix[:, 1]], ddloss)
        assert np.array_equal(np.stack([np.load(p)["rx_r"], np.load(p)["rx_c"]], 1), pix)
        res["sionna"][name] = {"samples": m["samples_per_tx"], "max_depth": m["max_depth"], "seconds": m["s_per_tx_median"],
                               "most_path_slots_per_call": m["most_path_slots_per_call"], "max_paths": m["max_paths"],
                               "rings": stats(db, ring), "rings_with_dd": stats(db_dd, ring)}
    (out / "report.json").write_text(json.dumps(res, indent=1))
    print(json.dumps(res, indent=1))
    if a.fig:
        import matplotlib
        matplotlib.use("Agg")
        import matplotlib.pyplot as plt
        labels = [f"{lo}-{hi} m" for lo, hi in RINGS]
        fig, axs = plt.subplots(1, 2, figsize=(13, 4.6), sharey=True)
        for ax, (title, block, key) in zip(axs, [("Sionna (scat09a calibration, + two-corner paths)", res["sionna"], "rings_with_dd"),
                                                 ("Our tracer (calibrated)", res["tracer"], "rings")]):
            for name, r in block.items():
                ax.plot(labels, [r[key][l]["covered_pct"] for l in labels], marker="o", label=name)
            ax.set_title(title, fontsize=10)
            ax.set_xlabel("distance from Tx")
            ax.grid(alpha=0.3)
            ax.legend(fontsize=8)
        axs[0].set_ylabel("receivers above -127 dB, %")
        plt.savefig(a.fig, dpi=80, bbox_inches="tight")


WINPROP_RINGS = [(0, 50), (50, 100), (100, 150), (150, 200), (200, 260), (260, 400)]


def winprop(a):
    """Outdoor % above the floor by distance ring for DPM, IRT2 and IRT4 on every case that has an IRT4 map."""
    cases = sorted(tuple(map(int, p.stem.split("_"))) for p in (ROOT / "gain/IRT4").glob("*.png"))
    rr, cc = np.mgrid[:256, :256]
    cov = {s: np.zeros(len(WINPROP_RINGS)) for s in ("DPM", "IRT2", "IRT4")}
    cnt = np.zeros(len(WINPROP_RINGS))
    for m, t in cases:
        out = load_map(m) == 0
        tr, tc = load_tx(m, t)
        d = np.hypot(rr - tr, cc - tc)
        masks = [out & (d >= lo) & (d < hi) for lo, hi in WINPROP_RINGS]
        cnt += [k.sum() for k in masks]
        for s in cov:
            g = load_gain(m, t, s) > THRESH
            cov[s] += [(g & k).sum() for k in masks]
    labels = [f"{lo}-{hi} m" for lo, hi in WINPROP_RINGS]
    res = {"commit": git_commit(), "cases": len(cases), "floor_db": round(FLOOR_DB, 1), "outdoor_pixels": dict(zip(labels, cnt.astype(int).tolist())),
           "covered_pct": {s: dict(zip(labels, np.round(100 * v / cnt, 1).tolist())) for s, v in cov.items()}}
    Path(a.out, "winprop_rings.json").write_text(json.dumps(res, indent=1))
    print(json.dumps(res, indent=1))


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("cmd", choices=["select", "tracer", "report", "winprop"])
    ap.add_argument("--scene", default="runs/cloudrf/mosaic_2km_anchor")
    ap.add_argument("--out", default="runs/convergence")
    ap.add_argument("--fig", default=None)
    a = ap.parse_args()
    {"select": select, "tracer": tracer, "report": report, "winprop": winprop}[a.cmd](a)


if __name__ == "__main__":
    main()
