"""Decode benchmark: recast-radar against the other Python radar readers.

Every timed sample is: file path -> every field of every sweep decoded to
float arrays in memory, file read included. Each reader is compared with the
recast-radar call that returns the same kind of object:

- xradar, MetPy, wradlib: `recast_radar.open(path).load()`, an xarray
  DataTree with every field loaded as floats (as xradar returns).
- Py-ART: `recast_radar.to_pyart(path)`, the same Py-ART Radar object that
  Py-ART's own reader returns.
 Each reader runs in its own
process, pinned to one CPU with one thread. The readers run in rounds, in a
different order every round, and the result is the median over rounds of each
round's median.

    pip install recast-radar arm_pyart metpy xradar wradlib psutil
    python decode_bench.py fetch             # download the files into ./data
    python decode_bench.py run               # prints a Markdown table
    python decode_bench.py run --json out.json --rounds 5 --iters 5

Every file is public: NOAA NEXRAD on AWS, OPERA ODIM_H5 (CC BY 4.0) and the
open-radar-data S-Pol CfRadial file.
"""

import argparse
import gc
import hashlib
import json
import os
import platform
import statistics
import subprocess
import sys
import time
import urllib.request
from pathlib import Path

HERE = Path(__file__).resolve().parent
DATA = HERE / "data"

# label, file name, format, url, sha256, readers
CASES = [
    ("Level II, KTLX 2024 (bzip2 records)", "KTLX20240315_000217_V06", "l2",
     "https://unidata-nexrad-level2.s3.amazonaws.com/2024/03/15/KTLX/KTLX20240315_000217_V06",
     "366e63315c7f541cadbac3e800dd5a0bd4880865f103e79db654a93c3ac09856",
     ["recast", "recast-pyart", "metpy", "pyart", "xradar"]),
    ("Level II, KTLX 2013 (gzip file)", "KTLX20130520_201643_V06.gz", "l2",
     "https://unidata-nexrad-level2.s3.amazonaws.com/2013/05/20/KTLX/KTLX20130520_201643_V06.gz",
     "772e01b154a5c966982a6d0aa2fc78bc64f08a9b77165b74dc02d7aa5aa69275",
     ["recast", "recast-pyart", "metpy", "pyart", "xradar"]),
    ("Level III, TLX N0B (super-res reflectivity)", "TLX_N0B_2026_06_22_08_06_23", "l3",
     "https://unidata-nexrad-level3.s3.amazonaws.com/TLX_N0B_2026_06_22_08_06_23",
     "630f1da18a8c75b151ce58dfa2bcb15fd2dd28515e574e14891510e26462ca30",
     ["recast", "recast-pyart", "metpy", "pyart"]),
    ("ODIM_H5 volume, DMI Rømø (dual-pol)", "dkrom.pvol.20260820T1130.h5", "odim",
     "https://s3.waw3-1.cloudferro.com/openradar-archive/2026/08/20/DK/dkrom/PVOL/"
     "dkrom%4020260820T1130%400.47_0.65_0.96_1.46_2.36_4.8_8.4_9.99_12.97_15.0"
     "%40DBZH_LDR_PHIDP_RHOHV_TH_VRAD_WRAD_ZDR.h5",
     "e3fafbdadc0e270c379845b29beb47af1278986f446823c3d9dbc420f8c78ce8",
     ["recast", "recast-pyart", "wradlib", "pyart", "xradar"]),
    ("CfRadial 1, NCAR S-Pol", "cfrad.20080604_002217_000_SPOL_v36_SUR.nc", "cfrad1",
     "https://raw.githubusercontent.com/openradar/open-radar-data/"
     "9cc67e67efa44134170d3a2f2b54f1aebe934483/data/cfrad.20080604_002217_000_SPOL_v36_SUR.nc",
     "67821b6c2bb0f27b5de49dee636f36e6e5bbad95f1ee168cb2d1af48e98992fe",
     ["recast", "recast-pyart", "wradlib", "pyart", "xradar"]),
]
NAMES = {"recast": "recast-radar open()", "recast-pyart": "recast-radar to_pyart()",
         "xradar": "xradar", "pyart": "Py-ART", "metpy": "MetPy", "wradlib": "wradlib"}
PACKAGES = {"recast": "recast-radar", "recast-pyart": "recast-radar", "metpy": "MetPy",
            "pyart": "arm_pyart", "xradar": "xradar", "wradlib": "wradlib"}
# The recast-radar reader each library's ratio divides by.
BASE = {"xradar": "recast", "metpy": "recast", "wradlib": "recast", "pyart": "recast-pyart"}
ONE_THREAD = {k: "1" for k in ["RAYON_NUM_THREADS", "OMP_NUM_THREADS", "OPENBLAS_NUM_THREADS",
                               "MKL_NUM_THREADS", "NUMEXPR_NUM_THREADS"]}


# ---------------------------------------------------------------- readers

def tree_elements(tree):
    return sum(int(v.size) for node in tree.subtree for v in node.ds.data_vars.values())


def make_reader(lib, fmt):
    """A function path -> number of decoded elements (a sanity count)."""
    if lib == "recast":
        import recast_radar
        # Values are lazy until read; load() decodes every field to floats.
        return lambda path: tree_elements(recast_radar.open(path).load())
    if lib == "recast-pyart":
        import recast_radar
        return lambda path: sum(int(f["data"].size)
                                for f in recast_radar.to_pyart(path).fields.values())
    if lib == "xradar":
        import xradar
        opener = {"l2": xradar.io.open_nexradlevel2_datatree,
                  "odim": xradar.io.open_odim_datatree,
                  "cfrad1": xradar.io.open_cfradial1_datatree}[fmt]
        return lambda path: tree_elements(opener(path).load())
    if lib == "pyart":
        import pyart
        reader = {"l2": pyart.io.read_nexrad_archive, "l3": pyart.io.read_nexrad_level3,
                  "odim": pyart.aux_io.read_odim_h5, "cfrad1": pyart.io.read_cfradial}[fmt]
        return lambda path: sum(int(f["data"].size) for f in reader(path).fields.values())
    if lib == "metpy":
        from metpy.io import Level2File, Level3File
        if fmt == "l2":
            def level2(path):  # Level2File scales every moment when it reads
                f = Level2File(path)
                return sum(int(d.size) for sweep in f.sweeps for ray in sweep
                           for _, (_, d) in ray[4].items())
            return level2
        def level3(path):
            f = Level3File(path)
            sym = getattr(f, "sym_block", None)
            return sum(len(row) for p in (sym[0] if sym else []) for row in p.get("data", []))
        return level3
    if lib == "wradlib":
        import wradlib
        if fmt == "odim":
            return lambda path: sum(int(v.size) for v in wradlib.io.read_opera_hdf5(path).values()
                                    if getattr(v, "ndim", 0) == 2)
        def generic(path):
            variables = wradlib.io.read_generic_netcdf(path).get("variables", {})
            return sum(int(v["data"].size) for v in variables.values()
                       if hasattr(v.get("data"), "size"))
        return generic
    raise ValueError(f"no {lib} reader for {fmt}")


def worker(lib, fmt, path, warmup, iters, cpu):
    """One process: decode `path` warmup + iters times, print one JSON line."""
    import warnings
    warnings.filterwarnings("ignore")
    import psutil
    if cpu is not None:
        psutil.Process().cpu_affinity([cpu])
    read = make_reader(lib, fmt)
    samples, elements = [], 0
    for i in range(warmup + iters):
        gc.collect()
        t = time.perf_counter()
        elements = read(path)
        ms = (time.perf_counter() - t) * 1000.0
        if i >= warmup:
            samples.append(ms)
    print(json.dumps({"median_ms": statistics.median(samples), "samples_ms": samples,
                      "elements": elements}))


# ---------------------------------------------------------------- driver

def fetch(_args):
    DATA.mkdir(exist_ok=True)
    for label, name, _fmt, url, sha, _libs in CASES:
        target = DATA / name
        if not target.exists():
            print("downloading", name)
            urllib.request.urlretrieve(url, target)
        if sha and hashlib.sha256(target.read_bytes()).hexdigest() != sha:
            raise SystemExit(f"{name}: SHA-256 mismatch")
        print("ok", label)


def version(lib):
    from importlib.metadata import version as v
    try:
        return v(PACKAGES[lib])
    except Exception:
        return "?"


def run(args):
    env = dict(os.environ, **ONE_THREAD, PYTHONWARNINGS="ignore")
    results = {}
    for r in range(args.rounds):
        jobs = [(c, lib) for c in CASES for lib in c[5]]
        shift = r % len(jobs)
        for (label, name, fmt, *_), lib in jobs[shift:] + jobs[:shift]:
            cmd = [sys.executable, __file__, "_worker", lib, fmt, str(DATA / name),
                   str(args.warmup), str(args.iters), str(args.cpu)]
            out = subprocess.run(cmd, env=env, capture_output=True, text=True)
            key = (label, lib)
            if out.returncode != 0:
                err = (out.stderr.strip().splitlines() or ["failed"])[-1]
                results.setdefault(key, {"error": err[:200]})
                continue
            row = json.loads(out.stdout.strip().splitlines()[-1])
            results.setdefault(key, {"rounds": [], "elements": row["elements"]})
            results[key].setdefault("rounds", []).append(row["median_ms"])
            print(f"round {r + 1}/{args.rounds}  {lib:8s} {row['median_ms']:9.1f} ms  {label}",
                  file=sys.stderr)
    libs = ["recast", "recast-pyart", "xradar", "pyart", "metpy", "wradlib"]
    meta = {
        "python": platform.python_version(), "os": platform.platform(),
        "cpu": platform.processor(), "pinned_cpu": args.cpu,
        "rounds": args.rounds, "iters": args.iters, "warmup": args.warmup,
        "versions": {PACKAGES[l]: version(l) for l in libs},
    }
    lines = ["| File | " + " | ".join(NAMES[l] for l in libs) + " |",
             "|---|" + "---:|" * len(libs)]
    table = []
    for label, *_rest in CASES:
        cells, row = [], {"file": label}
        for lib in libs:
            res = results.get((label, lib))
            if res is None:
                cells.append("")
            elif "rounds" in res:
                ms = statistics.median(res["rounds"])
                row[lib] = {"ms": ms, "rounds_ms": res["rounds"], "elements": res["elements"]}
                base = results.get((label, BASE.get(lib)), {}).get("rounds")
                ratio = f" ({ms / statistics.median(base):.1f}×)" if base else ""
                cells.append(f"{ms:,.0f}{ratio}" if ms >= 10 else f"{ms:.1f}{ratio}")
            else:
                row[lib] = {"error": res["error"]}
                cells.append("fails")
        table.append(row)
        lines.append(f"| {label} | " + " | ".join(cells) + " |")
    print("\n".join(lines))
    print("\n" + json.dumps(meta, ensure_ascii=False))
    if args.json:
        Path(args.json).write_text(json.dumps({"meta": meta, "results": table}, indent=1,
                                              ensure_ascii=False), encoding="utf-8")


def main():
    if len(sys.argv) > 1 and sys.argv[1] == "_worker":
        lib, fmt, path, warmup, iters, cpu = sys.argv[2:8]
        return worker(lib, fmt, path, int(warmup), int(iters), None if cpu == "None" else int(cpu))
    p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    sub = p.add_subparsers(dest="cmd", required=True)
    sub.add_parser("fetch")
    pr = sub.add_parser("run")
    pr.add_argument("--rounds", type=int, default=5)
    pr.add_argument("--iters", type=int, default=5)
    pr.add_argument("--warmup", type=int, default=1)
    pr.add_argument("--cpu", type=int, default=2, help="CPU to pin every reader to")
    pr.add_argument("--json")
    args = p.parse_args()
    {"fetch": fetch, "run": run}[args.cmd](args)


if __name__ == "__main__":
    main()
