import os, getopt, sys
from pixell import reproject, enmap
import healpy as hp
import gc 

freqs = [30, 90, 148, 219, 277, 350]

def get_all_output_file_names(sim_idx, output_dir, freqs=freqs):
    ret = []
    ## cmb
    for polidx in ["T","Q","U"]:
        ret.append(get_output_file_name('lensed_cmb', sim_idx, output_dir, freq=None, polidx=polidx))
    for compt_idx in ["kappa","ksz"]:
        ret.append(get_output_file_name(compt_idx, sim_idx, output_dir, freq=None, polidx=None))
    for freq in freqs:
        for compt_idx in ["tsz","rad_pts", "ir_pts"]:
            ret.append(get_output_file_name(compt_idx, sim_idx, output_dir, freq=freq, polidx=None))

        ret.append(get_output_file_name("combined", sim_idx, output_dir, freq=freq, polidx='T'))

    return ret

def get_output_file_name(compt_idx, sim_idx, output_dir, freq=None, polidx=None, freqs=freqs):
    if compt_idx in ["tsz", "rad_pts", "ir_pts"]:
        assert(freq in freqs)
        output_file = os.path.join(output_dir, f"{sim_idx:05d}/{compt_idx}_{freq:03d}ghz_{sim_idx:05d}.fits")
    elif compt_idx in ["kappa", "ksz"]:
        output_file = os.path.join(output_dir, f"{sim_idx:05d}/{compt_idx}_{sim_idx:05d}.fits")
    elif compt_idx in ["lensed_cmb"]:
        assert(polidx in ["T","Q","U"])
        output_file = os.path.join(output_dir, f"{sim_idx:05d}/{compt_idx}_{polidx}_{sim_idx:05d}.fits")
    elif compt_idx in ["combined"]:
        assert(freq in freqs)
        assert(polidx in ["T","Q","U"])
        output_file = os.path.join(output_dir, f"{sim_idx:05d}/{compt_idx}_{polidx}_{freq:03d}ghz_{sim_idx:05d}.fits")
    else:
        raise NotImplemented()

    return output_file

def main(argv):
    simnum = ''
    inputdir = ''
    outpudir = ''
    try:
        opts, args = getopt.getopt(argv,"hs:i:o:",["simnum=", "idir=","odir="])
    except getopt.GetoptError:
        print('mmdl_car2hp.py -s <simnum> -i <inputdir> -o <outputdir>')
        sys.exit(2)
    for opt, arg in opts:
        if opt == '-h':
            print('mmdl_car2hp.py -s <simnum> -i <inputdir> -o <outputdir>')
            sys.exit() 
        elif opt in ("-s", "--simnum"):
            simnum = int(arg)
        elif opt in ("-i", "--idir"):
            inputdir = arg
        elif opt in ("-o", "--odir"):
            outputdir = arg

    os.makedirs(os.path.join(outputdir, f"{simnum:05d}"), exist_ok=True)

    print(f'sim number is {simnum:05d}')
    print(f'Input folder is {inputdir}')
    print(f'Output folder is {outputdir}')

    output_files = get_all_output_file_names(simnum,  outputdir)
    for i, input_file in enumerate(get_all_output_file_names(simnum,  inputdir)):
        gc.collect()
        print(f"processing {input_file}")
        lmax = 10000 if "pts" not in input_file and "combined" not in input_file else 25000
        hp.write_map(output_files[i], 
                     reproject.healpix_from_enmap(enmap.read_map(input_file),
                                            lmax=lmax, nside=8192)
                    )
 


if __name__ == "__main__":
   main(sys.argv[1:])
