#!/usr/bin/env python
"""
Debug script to compare make_model_image() output vs manually reconstructed model.
Tests if there's a discrepancy in how models are being computed for residuals.
"""

import numpy as np
from pathlib import Path
from astropy.table import Table
from astropy.modeling.fitting import LevMarLSQFitter
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, load_fits_data_and_wcs,
    compute_crowdsource_weight_map, cutout_slices,
    replace_nan_pixels_for_fitting
)

# Load data
science_image = Path('/orange/adamginsburg/jwst/sickle/F480M/pipeline/jw03958-o007_t001_nircam_clear-f480m-nrcb_i2d.fits')
residual_image = Path('/orange/adamginsburg/jwst/sickle/F480M/pipeline/jw03958-o007_t001_nircam_clear-f480m-nrcb_iter2_daophot_basic_residual_i2d.fits')
stpsf_grid_file = Path('/orange/adamginsburg/jwst/sickle/psfs/nircam_nrcb5_f480m_fovp512_samp2_npsf16.fits')

print("Loading data...")
sci_data, sci_wcs, sci_err, sci_dq, sci_wht = load_fits_bundle(science_image)
res_data, res_wcs = load_fits_data_and_wcs(residual_image)
crowd_wht_map = compute_crowdsource_weight_map(sci_data, sci_err, dq=sci_dq, wht=sci_wht)
psf_model = to_griddedpsfmodel(str(stpsf_grid_file))
fwhm_pix = 2.574

# Load overfitting experiment results
exp_outdir = Path('/orange/adamginsburg/jwst/sickle/overfitting_experiments/test_smaller_fit')
stars_tbl = Table.read(exp_outdir / 'cutout_selected_stars.ecsv')
sweep_tbl = Table.read(exp_outdir / 'cutout_parameter_sweep_results.ecsv')

print(f"\nLoaded {len(stars_tbl)} selected stars")
print(f"Loaded {len(sweep_tbl)} sweep results")

# Pick a star with bad residuals (large negative center_resid)
bad_sweep = sweep_tbl[sweep_tbl['config_name'] == 'basic_local5_15_fit7']
worst_idx = np.argmin(bad_sweep['center_resid'])
worst_row = bad_sweep[worst_idx]

star_id = int(worst_row['star_id'])
print(f"\nAnalyzing star_id={star_id} with config={worst_row['config_name']}")
print(f"  center_resid={worst_row['center_resid']:.4f} (very negative = overfitting)")
print(f"  core_median_resid={worst_row['core_median_resid']:.4f}")
print(f"  core_min_resid={worst_row['core_min_resid']:.4f}")

star_row = stars_tbl[stars_tbl['star_id'] == star_id][0]
xc = float(star_row['xpix'])
yc = float(star_row['ypix'])

print(f"  Location: x={xc:.2f}, y={yc:.2f}")

# Extract cutout
halfsize = 18
ysl, xsl = cutout_slices(xc, yc, halfsize, sci_data.shape)
sci_cut = np.asarray(sci_data[ysl, xsl], dtype=float)
sci_fit_cut = replace_nan_pixels_for_fitting(sci_cut, fwhm_pix=fwhm_pix)
sci_err_cut = np.asarray(sci_err[ysl, xsl], dtype=float)
weight_cut = np.asarray(crowd_wht_map[ysl, xsl], dtype=float)

x0 = xc - xsl.start
y0 = yc - ysl.start

print(f"  Cutout: ({sci_cut.shape[0]}x{sci_cut.shape[1]}) at ({x0:.2f}, {y0:.2f}) in cutout coords")

# Re-fit with the problematic config
from astropy.table import Table as AstroTable
from astropy.stats import mad_std

local_noise = mad_std(sci_fit_cut[np.isfinite(sci_fit_cut)], ignore_nan=True)
if not np.isfinite(local_noise) or local_noise <= 0:
    local_noise = 1.0

flux0 = np.nansum(sci_fit_cut[sci_fit_cut > 0]) / 10  # rough estimate
init_tbl = AstroTable()
init_tbl['x_0'] = [x0]
init_tbl['y_0'] = [y0]
init_tbl['flux_0'] = [flux0]

print(f"\nInitial params: x0={x0:.2f}, y0={y0:.2f}, flux0={flux0:.2f}")

# Config that shows overfitting: basic_local5_15_fit7
localbkg = LocalBackground(5, 15)
fit_shape = (7, 7)

phot = PSFPhotometry(
    finder=None,
    localbkg_estimator=localbkg,
    psf_model=psf_model,
    fitter=LevMarLSQFitter(),
    fit_shape=fit_shape,
    aperture_radius=2.0 * fwhm_pix,
    progress_bar=False,
)

print(f"\nRunning photometry with fit_shape={fit_shape}, localbkg=(5,15)...")
result = phot(sci_fit_cut, init_params=init_tbl, error=np.where(np.isfinite(sci_err_cut), sci_err_cut, 1e11))

if len(result) == 0:
    print("ERROR: No result from photometry!")
    sys.exit(1)

xfit = float(result['x_fit'][0]) if 'x_fit' in result.colnames else float(result['x_0'][0])
yfit = float(result['y_fit'][0]) if 'y_fit' in result.colnames else float(result['y_0'][0])
flux_fit = float(result['flux_fit'][0]) if 'flux_fit' in result.colnames else np.nan

print(f"Fitted params: x_fit={xfit:.4f}, y_fit={yfit:.4f}, flux_fit={flux_fit:.4f}")
print(f"Delta: dx={xfit-x0:.4f}, dy={yfit-y0:.4f}, dflux={flux_fit-flux0:.4f}")

# Method 1: Use make_model_image() from the photometry object
print("\n=== Method 1: phot.make_model_image() ===")
model_from_phot = phot.make_model_image(sci_fit_cut.shape, psf_shape=(21, 21), include_localbkg=False)
resid_from_phot = sci_fit_cut - model_from_phot

center_idx_y, center_idx_x = int(np.rint(yfit)), int(np.rint(xfit))
center_val_phot = float(model_from_phot[center_idx_y, center_idx_x])
resid_center_phot = float(resid_from_phot[center_idx_y, center_idx_x])

print(f"Model at center ({center_idx_x}, {center_idx_y}): {center_val_phot:.4f}")
print(f"Residual at center: {resid_center_phot:.4f}")

# Method 2: Manually reconstruct using result parameters
print("\n=== Method 2: Manual reconstruction from fitted params ===")
psf_at_pos = psf_model.evaluate(
    x=np.arange(sci_fit_cut.shape[1], dtype=float),
    y=np.arange(sci_fit_cut.shape[0], dtype=float)[:, np.newaxis],
    flux=1.0,
    x_0=xfit,
    y_0=yfit,
)
model_manual = flux_fit * psf_at_pos
resid_manual = sci_fit_cut - model_manual

center_val_manual = float(model_manual[center_idx_y, center_idx_x])
resid_center_manual = float(resid_manual[center_idx_y, center_idx_x])

print(f"Model at center ({center_idx_x}, {center_idx_y}): {center_val_manual:.4f}")
print(f"Residual at center: {resid_center_manual:.4f}")

# Method 3: Check if there's a background component being added
print("\n=== Method 3: Check background subtraction ===")
from astropy.stats import sigma_clipped_stats

bkg_med, _, _ = sigma_clipped_stats(sci_fit_cut, sigma=3.0)
print(f"Background median in cutout: {bkg_med:.4f}")

# Compare the models
diff = model_from_phot - model_manual
print(f"\n=== Comparison ===")
print(f"Max abs difference in models: {np.nanmax(np.abs(diff)):.6f}")
print(f"Model 1 (phot): center={center_val_phot:.6f}, resid_center={resid_center_phot:.6f}")
print(f"Model 2 (manual): center={center_val_manual:.6f}, resid_center={resid_center_manual:.6f}")
print(f"Difference in residual_center: {resid_center_phot - resid_center_manual:.6f}")

# Check fitter convergence info
print(f"\n=== Fitter convergence info ===")
if hasattr(phot, 'fitter') and hasattr(phot.fitter, 'fit_info') and phot.fitter.fit_info:
    print(f"fit_info: {phot.fitter.fit_info}")
else:
    print("No fit_info available from fitter")

# Check result flags
print(f"\n=== Result flags ===")
for colname in result.colnames:
    if 'flag' in colname.lower() or 'cov' in colname.lower():
        print(f"{colname}: {result[colname][0]}")

print("\nDone.")
