diff --git a/src/xdust/halos/halo.py b/src/xdust/halos/halo.py index 11bbe74..8bc910d 100644 --- a/src/xdust/halos/halo.py +++ b/src/xdust/halos/halo.py @@ -403,23 +403,18 @@ def fake_image(self, arf, src_flux, exposure, arf = InterpolatedUnivariateSpline(arf_x, arf_y, k=1) # Source counts to use for each energy bin - if self.lam_unit == 'angs': - ltemp = self.lam * u.angstrom - ltemp_kev = ltemp.to(u.keV, equivalencies=u.spectral()).value - arf_temp = arf(ltemp_kev)[::-1] - src_counts = src_flux * arf_temp * exposure - else: - src_counts = src_flux * arf(self.lam) * exposure + lam_kev = self.lam.to(u.keV, equivalencies=u.spectral()).value + src_counts = src_flux * arf(lam_kev) * exposure # Decide which energy indexes to use if lmin is None: imin = 0 else: - imin = min(np.arange(len(self.lam))[self.lam >= lmin]) + imin = min(np.arange(len(self.lam))[self.lam >= lmin * self.lam.unit]) if lmax is None: iend = len(self.lam) else: - iend = max(np.arange(len(self.lam))[self.lam <= lmax]) + iend = max(np.arange(len(self.lam))[self.lam <= lmax * self.lam.unit]) #iend = imax #if imax < 0: diff --git a/tests/test_galhalo.py b/tests/test_galhalo.py index 45858bd..36094c7 100644 --- a/tests/test_galhalo.py +++ b/tests/test_galhalo.py @@ -2,6 +2,8 @@ import numpy as np from scipy.integrate import trapezoid as trapz import astropy.units as u +from astropy.io import fits +from astropy.table import Table from xdust.halos import * from xdust import grainpop @@ -167,3 +169,41 @@ def test_halo_io(test): assert np.all(percent_diff(new_halo.taux,test.taux) < 0.01) assert np.all(percent_diff(new_halo.norm_int.flatten(),test.norm_int.flatten()) < 0.01) assert new_halo.lam.unit == 'keV' + +def create_arf(tmp_path, outfile, evals, effarea) -> str: + """ + Create a dummy ARF file for testing + """ + t = Table() + t['ENERG_LO'] = evals[:-1] + t['ENERG_HI'] = evals[1:] + t['SPECRESP'] = np.full(len(evals) - 1, effarea) + + hdu_table = fits.BinTableHDU(t.as_array(), name='SPECRESP') + hdu_primary = fits.PrimaryHDU() + hdul = fits.HDUList([hdu_primary, hdu_table]) + hdul.writeto(tmp_path / outfile, overwrite=True) + return str(tmp_path / outfile) + +@pytest.fixture +def halo(): + new_halo = galhalo.UniformGalHalo(EVALS, THVALS) + new_halo.calculate(GPOP) + return new_halo + +@pytest.fixture +def arf(tmp_path): + arf = create_arf(tmp_path, 'test_arf.fits', EVALS, 100.0) + return arf + +def test_fake_image(halo, arf): + image = halo.fake_image(arf, src_flux=FABS, exposure=1e4, pix_scale=1.0, num_pix=[8, 16]) + assert image.shape == (16, 8) + assert np.all(image >= 0.0) + assert np.all(image == np.floor(image)) + assert image.sum() > 0 + +def test_fake_image_lmin_lmax(halo, arf): + restricted = halo.fake_image(arf, src_flux=FABS, exposure=1e4, pix_scale=1.0, num_pix=[8, 16], lmin=0.5, lmax=5.0) + unrestricted = halo.fake_image(arf, src_flux=FABS, exposure=1e4, pix_scale=1.0, num_pix=[8, 16]) + assert 0 < restricted.sum() < unrestricted.sum()