"""Offline IceCube probability-area audit.
Usage: python icecube-localization-analysis.py icecube-2026-localization-maps.zip --output reproduced-results.csv
Requires numpy, pandas, astropy, astropy-healpix. No network calls.
The output reproduces numerical areas and circle integrals, not the figure layout.
"""
import argparse
import gzip
import hashlib
import io
import json
import zipfile
import numpy as np
import pandas as pd
from astropy.io import fits
from astropy import units as u
from astropy.coordinates import SkyCoord
from astropy_healpix import uniq_to_level_ipix, healpix_to_lonlat

MOON_DEG2 = np.pi * 0.25**2

def circle_probability(level, ipix, density, ra, dec, radius_deg, target_level=15):
    """Integrate native piecewise-constant density over a spherical cap.
    Refine only boundary cells; child densities inherit parent density.
    2/nside radians conservatively bounds the center-to-edge separation.
    Terminal cells are selected by center. No new reconstruction information
    is inferred by subdivision.
    """
    lev, pix, rho = level.copy(), ipix.copy(), density.copy()
    center = SkyCoord(ra*u.deg, dec*u.deg)
    radius = np.deg2rad(radius_deg)
    result = 0.
    for _ in range(target_level+1):
        lon, lat = healpix_to_lonlat(pix, 2**lev, order="nested")
        distance = center.separation(SkyCoord(lon,lat)).rad
        margin = 2.0/(2**lev)
        area = np.pi/(3*4.0**lev)
        inside = distance+margin < radius
        outside = distance-margin > radius
        result += np.sum(rho[inside]*area[inside])
        border = ~(inside|outside)
        terminal = border & (lev >= target_level)
        result += np.sum(rho[terminal]*(distance[terminal] <= radius)*area[terminal])
        refine = border & ~terminal
        if not np.any(refine):
            break
        lev = np.repeat(lev[refine]+1,4)
        pix = (pix[refine,None]*4+np.arange(4)).ravel()
        rho = np.repeat(rho[refine],4)
    return float(result)

def analyze(zip_path):
    records = []
    with zipfile.ZipFile(zip_path) as archive:
        manifest = json.loads(archive.read("input-manifest.json"))
        for entry in manifest["files"]:
            name = entry["file"]
            compressed = archive.read(name)
            assert len(compressed) == entry["bytes"], name
            assert hashlib.sha256(compressed).hexdigest() == entry["sha256"], name
            if not name.endswith(".fits.gz"):
                continue
            raw = gzip.decompress(compressed)
            assert hashlib.sha256(raw).hexdigest() == entry["uncompressed_sha256"], name
            with fits.open(io.BytesIO(raw)) as hdul:
                header = hdul[1].header.copy()
                uniq = np.asarray(hdul[1].data["UNIQ"],dtype=np.int64)
                rho = np.asarray(hdul[1].data["PROBDENSITY"],dtype=float)
            assert header["ORDERING"] == "NUNIQ"
            assert header["TUNIT2"] == "sr-1"
            level, ipix = uniq_to_level_ipix(uniq)
            omega = np.pi/(3*4.0**level)
            finite = np.isfinite(rho)
            assert np.all(rho[finite] >= 0)
            rho = np.where(finite,rho,0)
            total = np.sum(rho*omega)
            assert total > 0
            rho /= total
            # Verify no duplicate or overlapping native pixels.
            max_level = int(level.max())
            starts = ipix*4**(max_level-level)
            ends = (ipix+1)*4**(max_level-level)
            spatial_order = np.argsort(starts)
            assert len(np.unique(uniq)) == len(uniq)
            assert np.all(starts[spatial_order][1:] >= ends[spatial_order][:-1])
            rank = np.argsort(-rho,kind="stable")
            cumulative = np.cumsum(rho[rank]*omega[rank])
            areas = omega*(180/np.pi)**2
            record = {
                "event": name.split("_")[0],
                "source_url": entry["source_url"],
                "sha256": entry["sha256"],
                "date": header["DATE-OBS"],
                "probability_sum": total,
                "invalid_pixels": int((~finite).sum()),
                "sky_coverage_fraction": float(omega.sum()/(4*np.pi)),
                "status_audit": "revision/retraction unverified; no event notices supplied",
            }
            for cl in (50,90):
                target = cl/100
                k = np.searchsorted(cumulative,target)
                previous = 0. if k == 0 else cumulative[k-1]
                lower = areas[rank[:k]].sum()
                upper = lower+areas[rank[k]]
                interp = lower+(target-previous)/(cumulative[k]-previous)*areas[rank[k]]
                stored = header[f"CONTOUR_AREA_{cl}"]
                ra_mean = (abs(header[f"RA_ERR_PLUS_{cl}"])+abs(header[f"RA_ERR_MINUS_{cl}"]))/2
                dec_mean = (abs(header[f"DEC_ERR_PLUS_{cl}"])+abs(header[f"DEC_ERR_MINUS_{cl}"]))/2
                radius = np.sqrt(ra_mean*np.cos(np.deg2rad(header["DEC"]))*dec_mean)
                cap_area = 2*np.pi*(1-np.cos(np.deg2rad(radius)))*(180/np.pi)**2
                raw_area = np.interp(target,np.r_[0,cumulative*total],
                                     np.r_[0,np.cumsum(areas[rank])])
                record.update({
                    f"area{cl}_lower_deg2": lower,
                    f"area{cl}_upper_deg2": upper,
                    f"area{cl}_interp_deg2": interp,
                    f"area{cl}_rawmass_deg2": raw_area,
                    f"moons{cl}": interp/MOON_DEG2,
                    f"header{cl}_deg2": stored,
                    f"header{cl}_moons": stored/MOON_DEG2,
                    f"diff{cl}_percent": 100*(interp/stored-1),
                    f"circle{cl}_radius_deg": radius,
                    f"circle{cl}_deg2": cap_area,
                    f"circle{cl}_moons": cap_area/MOON_DEG2,
                    f"circle{cl}_over_map": cap_area/interp,
                    f"circle{cl}_mass_refined": circle_probability(
                        level,ipix,rho,header["RA"],header["DEC"],radius),
                })
            records.append(record)
    return pd.DataFrame(records)

if __name__ == "__main__":
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("archive")
    parser.add_argument("--output",default="reproduced-results.csv")
    args = parser.parse_args()
    result = analyze(args.archive)
    result.to_csv(args.output,index=False,float_format="%.10g")
    print(result[["event","moons50","moons90"]].to_string(index=False))
