#!/usr/bin/env python
"""
Test PSF photometry overfitting directly on CAL frame data.
Fit various configurations and analyze why some stars overfit while others don't.
"""

import numpy as np
from pathlib import Path
from astropy.table import Table
from astropy.modeling.fitting import LevMarLSQFitter
from astropy.stats import sigma_clipped_stats, mad_std
from photutils.background import LocalBackground
from photutils.psf import PSFPhotometry
from stpsf.utils import to_griddedpsfmodel
import sys
sys.path.insert(0, '/orange/adamginsburg/repos/brick-jwst-2221')

from brick2221.analysis.overfitting_experiment_f480m import (
    load_fits_bundle, cutout_slices,
    replace_nan_pixels_for_fitting
)

# Load data
cal_frame = Path('/orange/adamginsburg/jwst/sickle/F480M/pipeline/jw03958007001_03104_00001_nrcblong_cal.fits')
sci_data, sci_wcs, sci_err, sci_dq, sci_wht = load_fits_bundle(cal_frame)

stpsf_grid_file = Path('/orange/adamginsburg/jwst/sickle/psfs/nircam_nrcb5_f480m_fovp512_samp2_npsf16.fits')
psf_model = to_griddedpsfmodel(str(stpsf_grid_file))
fwhm_pix = 2.574

# Load pre-computed residuals
residual_path = Path('/orange/adamginsburg/jwst/sickle/F480M/pipeline/jw03958007001_03104_00001_nrcblong_destreak_o007_crf_satstar_residual.fits')
from astropy.io import fits
with fits.open(residual_path) as hdul:
    residuals = np.asarray(hdul[0].data, dtype=float)

# Detect negative residual stars (oversubtracted) in residuals
inv_residual = -residuals
finite = np.isfinite(inv_residual)
med, _, std = sigma_clipped_stats(inv_residual[finite], sigma=3.0)
robust_std = mad_std(inv_residual[finite], ignore_nan=True)
noise = robust_std if np.isfinite(robust_std) and robust_std > 0 else std

threshold = 3.0 * noise
from photutils.detection import DAOStarFinder
finder = DAOStarFinder(threshold=threshold, fwhm=float(fwhm_pix))
det = finder(inv_residual - med)

print(f"Found {len(det)} negative residual detections in CAL frame")
print(f"Threshold: {threshold:.2f}, Noise: {noise:.2f}")
print()

# Select top 30 by flux
if len(det) > 0:
    det = det[np.argsort(det['flux'])[-30:]]
    det = det[::-1]  # Sort descending
    det = det[:30]

print(f"Top 30 by flux:")
print(f"{'#':<4} {'Peak':<10} {'Flux':<12} {'x':<8} {'y':<8} {'Chi2(6,10)':<12} {'Chi2(2,5)':<12}")
print("-" * 90)

halfsize = 18
results = []

for i, star_row in enumerate(det):
    xc = float(star_row['xcentroid'])
    yc = float(star_row['ycentroid'])

    # Make sure within bounds
    if xc - halfsize < 0 or xc + halfsize >= sci_data.shape[1]:
        continue
    if yc - halfsize < 0 or yc + halfsize >= sci_data.shape[0]:
        continue

    ysl, xsl = cutout_slices(xc, yc, halfsize, sci_data.shape)
    sci_cut = np.asarray(sci_data[ysl, xsl], dtype=float)
    sci_err_cut = np.asarray(sci_err[ysl, xsl], dtype=float)

    sci_fit_cut = replace_nan_pixels_for_fitting(sci_cut, fwhm_pix=fwhm_pix)
    x0 = xc - xsl.start
    y0 = yc - ysl.start

    peak_value = float(sci_fit_cut[int(np.rint(y0)), int(np.rint(x0))])
    flux0 = np.nansum(sci_fit_cut[sci_fit_cut > 0]) / 10

    init_tbl = Table()
    init_tbl['x_0'] = [x0]
    init_tbl['y_0'] = [y0]
    init_tbl['flux_0'] = [flux0]

    # Test two configurations
    chi2_610 = np.nan
    chi2_25 = np.nan

    try:
        # Config 1: LocalBackground(6, 10)
        phot = PSFPhotometry(
            finder=None,
            localbkg_estimator=LocalBackground(6, 10),
            psf_model=psf_model,
            fitter=LevMarLSQFitter(),
            fit_shape=(7, 7),
            aperture_radius=2.0 * fwhm_pix,
            progress_bar=False,
        )
        result = phot(sci_fit_cut, init_params=init_tbl, error=sci_err_cut)
        fvec = np.asarray(phot.fitter.fit_info.get('fvec', []))
        chi2_610 = float(np.sum(fvec**2))
    except:
        pass

    try:
        # Config 2: LocalBackground(2, 5)
        phot = PSFPhotometry(
            finder=None,
            localbkg_estimator=LocalBackground(2, 5),
            psf_model=psf_model,
            fitter=LevMarLSQFitter(),
            fit_shape=(7, 7),
            aperture_radius=2.0 * fwhm_pix,
            progress_bar=False,
        )
        result = phot(sci_fit_cut, init_params=init_tbl, error=sci_err_cut)
        fvec = np.asarray(phot.fitter.fit_info.get('fvec', []))
        chi2_25 = float(np.sum(fvec**2))
    except:
        pass

    print(f"{i+1:<4} {peak_value:<10.1f} {star_row['flux']:<12.0f} {xc:<8.1f} {yc:<8.1f} {chi2_610:<12.0f} {chi2_25:<12.0f}")

    results.append({
        'idx': i,
        'xc': xc,
        'yc': yc,
        'peak': peak_value,
        'flux': star_row['flux'],
        'chi2_610': chi2_610,
        'chi2_25': chi2_25,
    })

print("\n" + "="*90)
print("ANALYSIS")
print("="*90)

if results:
    chi2_610_vals = [r['chi2_610'] for r in results if np.isfinite(r['chi2_610'])]
    chi2_25_vals = [r['chi2_25'] for r in results if np.isfinite(r['chi2_25'])]

    print(f"\nLocalBackground(6, 10): chi2 min={np.min(chi2_610_vals):.0f}, max={np.max(chi2_610_vals):.0f}, median={np.median(chi2_610_vals):.0f}")
    print(f"LocalBackground(2, 5):  chi2 min={np.min(chi2_25_vals):.0f}, max={np.max(chi2_25_vals):.0f}, median={np.median(chi2_25_vals):.0f}")
    print(f"\nConfiguration (2,5) improves fits by: {(np.median(chi2_610_vals) - np.median(chi2_25_vals)) / np.median(chi2_610_vals) * 100:.1f}%")

print("\nDone.")
