### an example code to search signals in the light curve database ###

## import the necessary modules
import numpy as np
from pathlib import Path
src_dir = str(Path(__file__).parent) + '/src'
lc_dir = str(Path(__file__).parent) + '/data/lightcurve'
cat_dir = str(Path(__file__).parent) + '/data/catalog'
from src.HDF5DataSet import HDF5DataSet
from read_lc import read_lc, plot_cmd, plot_finder_chart
## numpy.loadtxt's UserWarning is so annoying!!!
import warnings
warnings.filterwarnings("ignore", category=UserWarning)


def load_datasets(stamp_comb, bands=['r', 'z'], verbose=False):
    """load light curve datasets for a given stamp combination

    Args:
        stamp_comb (str): stamp combination string
    Returns:
        list: list of HDF5DataSet objects
    """
    if verbose:
        print(f"Loading datasets for stamp combination {stamp_comb} ...")
    stamp1, stamp2 = stamp_comb[0], stamp_comb[1]
    datasets = {}
    for stamp in [stamp1, stamp2]:
        if stamp in ['', '-']: continue
        datasets[stamp] = {}
        for band in bands:
            # result = {'id': star_full_id, 'band': band, 'stamp': stamp}
            datasets[stamp][band] = HDF5DataSet(lc_dir + f'/lc_{stamp}_{band}.h5')
            if verbose:
                print(f"  - Loaded dataset {lc_dir + f'/lc_{stamp}_{band}.h5'}")
    if verbose: 
        print("All datasets loaded.")
    return datasets

def load_star_info_list(stamp_comb, verbose=False):
    """load star info list for a given stamp combination

    Args:
        stamp_comb (str): stamp combination string
    Returns:
        list: list of star info dicts
    """
    stamp1, stamp2 = stamp_comb[0], stamp_comb[1]
    raw_list = np.loadtxt(cat_dir + f'/group/stargroup_{stamp1}_{stamp2}.cat',
                          dtype=str)
    # if star_info_list has only one star (one dimension), convert it to two dimensions
    if raw_list.ndim == 1:
        raw_list = raw_list[np.newaxis, :]
    star_info_list = []
    for s in raw_list:
        star_info = {'global_id': int(s[0]),
                     'ra': float(s[1]), 'dec': float(s[2]),
                     'mag': float(s[3]), 'merr': float(s[4]),
                     'stamp1': s[5], 'internal_id1': int(s[6]),
                     'stamp2': s[11], 'internal_id2': int(s[12])}
        star_info_list.append(star_info)
    if verbose:
        print(f"Loaded {len(star_info_list)} stars from {cat_dir + f'/stargroup_{stamp1}_{stamp2}.cat'}")
    return star_info_list

# # a toy function to search for events
def event_search(light_curves):
    """a toy function to search for high rms light curves

    Args:
        light_curves (dict): dict of light curves
    Returns:
        bool: whether an event is found
    """
    # a toy example: check if there is any point with flux > 1000
    rms = 0
    for lcdata in light_curves:
        if lcdata['band'] == 'z':
            rms += np.sqrt(np.sum(lcdata['flux']/lcdata['ferr'])**2)
    return (rms > 20)

if __name__ == "__main__":
    import time

    # load all stamp combinations
    stamp_comb_list = np.loadtxt(cat_dir + '/stamp_comb.list', dtype=str)

    # load light curve datasets for one stamp combination
    stamp_comb = stamp_comb_list[0] # choose the first one
    datasets = load_datasets(stamp_comb, verbose=True)

    # load star info list for this stamp combination
    star_info_list = load_star_info_list(stamp_comb, verbose=True)
    
    # read light curves for the first star in the list
    star_info = star_info_list[0]
    for i, star_info in enumerate(star_info_list[:100]):
        time_begin = time.time()
        light_curves = read_lc(star_info['global_id'], 
                               star_info=star_info, 
                               datasets=datasets, 
                               verbose=False)
        has_signal = event_search(light_curves)
        if has_signal:
            light_curves = read_lc(star_info['global_id'], 
                                   star_info=star_info, 
                                   datasets=datasets,
                                   read_cmd=True, 
                                   verbose=False)
            plot_cmd(light_curves, saveto=f'test/cmd_{star_info['global_id']:08d}.png')
            print(f"star {star_info['global_id']}, has_signal: {has_signal:d}, cost:{time.time()-time_begin:.3f}, saved to test/cmd_{star_info['global_id']:08d}.png")
        else:
            print(f"star {star_info['global_id']}, has_signal: {has_signal:d}, cost:{time.time()-time_begin:.3f}")
        