#!/usr/bin/env python3
"""Quantify the superseded inclusive-PI versus corrected half-open Fe-L products."""

from __future__ import annotations

import csv
import json
from pathlib import Path

import astropy.units as u
import numpy as np
from astropy.coordinates import SkyCoord
from astropy.io import fits
from astropy.wcs import WCS
from scipy.stats import spearmanr

ROOT = Path(__file__).resolve().parent
OUT = ROOT / "joint_spectrum_fitting_2T_basedon_region_v22/r47_mos2_fel_ratio_map_pilot_20260713"
OLD = OUT / "superseded_inclusive_pi_boundary_20260713"
NEW_FITS = OUT / "r47_mos2_fel_ratio_point_sources_included_background_corrected.fits"
OLD_FITS = OLD / NEW_FITS.name
NEW_SECTORS = OUT / "r47_mos2_fel_ratio_point_sources_included_sector_comparison.csv"
OLD_SECTORS = OLD / NEW_SECTORS.name
NEW_SUMMARY = OUT / "r47_mos2_fel_ratio_background_corrected_summary.json"
OLD_SUMMARY = OLD / NEW_SUMMARY.name
M104 = SkyCoord(189.9976 * u.deg, -11.6231 * u.deg, frame="fk5")


def load_ratio(path: Path) -> tuple[np.ndarray, np.ndarray, WCS]:
    ratio = np.asarray(fits.getdata(path, "RATIO"), dtype=float)
    valid = np.asarray(fits.getdata(path, "VALID"), dtype=bool)
    return ratio, valid, WCS(fits.getheader(path, "RATIO"))


def load_csv(path: Path) -> list[dict[str, str]]:
    with path.open(newline="") as stream:
        return list(csv.DictReader(stream))


def main() -> None:
    old_ratio, old_valid, wcs = load_ratio(OLD_FITS)
    new_ratio, new_valid, new_wcs = load_ratio(NEW_FITS)
    if not wcs.wcs.compare(new_wcs.wcs, cmp=0, tolerance=0.0):
        raise RuntimeError("Old/new ratio WCS differs")

    yy, xx = np.indices(old_ratio.shape)
    sky = wcs.pixel_to_world(xx, yy)
    radius_arcmin = sky.separation(M104).to_value(u.arcmin)
    common = old_valid & new_valid & np.isfinite(old_ratio) & np.isfinite(new_ratio)
    annulus = common & (radius_arcmin >= 4.0) & (radius_arcmin < 7.0)

    old_rows = {row["sector_id"]: row for row in load_csv(OLD_SECTORS)}
    new_rows = {row["sector_id"]: row for row in load_csv(NEW_SECTORS)}
    columns = ("full_ratio", "point_source_masked_full_ratio")
    impact_rows: list[dict[str, object]] = []
    for sector_id in sorted(new_rows):
        row: dict[str, object] = {"sector_id": sector_id}
        for column in columns:
            old_value = float(old_rows[sector_id][column])
            new_value = float(new_rows[sector_id][column])
            row[f"old_{column}"] = old_value
            row[f"new_{column}"] = new_value
            row[f"delta_{column}"] = new_value - old_value
            row[f"fractional_delta_{column}"] = (new_value - old_value) / old_value
        impact_rows.append(row)

    output_csv = OUT / "r47_mos2_fel_halfopen_pi_boundary_sector_impact.csv"
    with output_csv.open("w", newline="") as stream:
        writer = csv.DictWriter(stream, fieldnames=list(impact_rows[0]))
        writer.writeheader()
        writer.writerows(impact_rows)

    old_summary = json.loads(OLD_SUMMARY.read_text())
    new_summary = json.loads(NEW_SUMMARY.read_text())
    delta = new_ratio[annulus] - old_ratio[annulus]
    payload = {
        "superseded_definition": {
            "low": "700 <= PI <= 875",
            "high": "875 <= PI <= 1050",
            "input_counts": old_summary["input_counts"],
        },
        "corrected_definition": {
            "low": "700 <= PI < 875",
            "high": "875 <= PI < 1050",
            "input_counts": new_summary["input_counts"],
        },
        "removed_boundary_events": {
            "low_PI_875": 50,
            "high_PI_1050": 29,
        },
        "r4_r7_pixel_morphology": {
            "common_valid_pixels": int(annulus.sum()),
            "spearman_rho": float(spearmanr(old_ratio[annulus], new_ratio[annulus]).statistic),
            "median_signed_ratio_change": float(np.median(delta)),
            "median_absolute_ratio_change": float(np.median(np.abs(delta))),
            "p95_absolute_ratio_change": float(np.percentile(np.abs(delta), 95)),
        },
        "full_map_median_ratio": {
            "old": old_summary["ratio_statistics"]["full"]["median"],
            "new": new_summary["ratio_statistics"]["full"]["median"],
        },
        "sector_impact_csv": str(output_csv),
        "maximum_absolute_sector_full_ratio_change": float(
            max(abs(float(row["delta_full_ratio"])) for row in impact_rows)
        ),
        "interpretation": (
            "The inclusive-PI release is superseded. The correction is mandatory for provenance, "
            "while these metrics quantify whether the R4-R7 morphology changed materially."
        ),
    }
    output_json = OUT / "r47_mos2_fel_halfopen_pi_boundary_impact_summary.json"
    output_json.write_text(json.dumps(payload, indent=2) + "\n")
    print(json.dumps(payload, indent=2))


if __name__ == "__main__":
    main()
