import numpy as np
import matplotlib.pyplot as plt
import astropy.io.fits as pyfits

from grizli.jwst_utils import flag_nircam_hot_pixels

signal = np.zeros((48,48), dtype=np.float32)

# hot
signal[16,16] = 10

# plus
for off in [-1,1]:
    signal[32+off, 32] = 10
    signal[32, 32+off] = 7

err = np.ones_like(signal)
np.random.seed(1)
noise = np.random.normal(size=signal.shape)*err

dq = np.zeros(signal.shape, dtype=int)
dq[32,32] = 2048 # HOT

header = pyfits.Header()
header['MDRIZSKY'] = 0.

hdul = pyfits.HDUList([
    pyfits.ImageHDU(data=signal+noise, name='SCI', header=header),
    pyfits.ImageHDU(data=err, name='ERR'),
    pyfits.ImageHDU(data=dq, name='DQ'),
])

sn, dq_flag, count = flag_nircam_hot_pixels(hdul)

fig, axes = plt.subplots(1,2,figsize=(8,4), sharex=True, sharey=True)

axes[0].imshow(signal + noise, vmin=-2, vmax=9, cmap='gray')
axes[0].set_xlabel('Simulated data')
axes[1].imshow(dq_flag, cmap='magma')
axes[1].set_xlabel('Flagged pixels')

for ax in axes:
    ax.set_xticklabels([])
    ax.set_yticklabels([])

fig.tight_layout(pad=1)

plt.show()