import numpy as np, pandas as pd
from dataclasses import dataclass
import polars as pl

inpdir = '../input/ariel-data-challenge-2024'

def read_adc_info():
    # returns adc_info, train_planets, test_planets
    train_adc_info = pd.read_csv(f'{inpdir}/train_adc_info.csv', index_col='planet_id')
    test_adc_info = pd.read_csv(f'{inpdir}/test_adc_info.csv', index_col='planet_id')

    adc_info = pd.concat([train_adc_info, test_adc_info], axis=0)
    train_planets = train_adc_info.index.tolist()
    test_planets = test_adc_info.index.tolist()
    
    return adc_info, train_planets, test_planets

adc_info, train_planets, test_planets = read_adc_info()

@dataclass
class CalibConfig:
    dtype_str:str='np.float32'
    linear_corr:bool=True
    dark:bool=True 
    flat:bool=True
    time_binning:int=30
    hot_or_dead:bool=True # ignore hot_or_dead

# compatible function for sigma_clip(x, sigma=5, maxiter=5)
def sigma_clip_compat(x, sigma):
    smin, s0, smax = np.quantile(x.flatten(), [1-0.954, 0.5, 0.954])  # -2*sigma to 2*sigma
    ds = (smax - smin) / 4 * 7/5.  # 7/5 adhoc
    qmin, qmax = s0 - ds * sigma, s0 + ds * sigma 
    return (x < qmin) | (x > qmax)

def apply_linear_correction(signal, linear_corr):
    # correction coef is reverse of poly1d
    s = None
    for c in reversed(linear_corr):
        if s is None:
            s = np.zeros_like(signal)
        else:
            s *= signal
        s += c
    return s

dt_airs = np.full(11250, 0.1)
dt_airs[1::2] += 4.5

dt_fgs1 = np.full(135000, 0.1)
dt_fgs1[1::2] += 0.1

sensor_const = {
    'airs': ('AIRS-CH0', 1, 356, 1),  # signal prefix, binning multiplier, feat dim, axis to mean
    'fgs1': ('FGS1', 12, 32, (1,2))
}

def read_signal(planet_id, sensor, conf=CalibConfig()):
    dtype = eval(conf.dtype_str)
    dname = f'{inpdir}/train/{planet_id}' if planet_id in train_planets else f'{inpdir}/test/{planet_id}'

    assert sensor in sensor_const
    prefix, tmul, fdim, maxis = sensor_const[sensor]

    signal = pl.read_parquet(f'{dname}/{prefix}_signal.parquet').to_numpy().reshape(-1, 32, fdim)

    flat = pd.read_parquet(f'{dname}/{prefix}_calibration/flat.parquet').values
    dark = pd.read_parquet(f'{dname}/{prefix}_calibration/dark.parquet').values
    dead = pd.read_parquet(f'{dname}/{prefix}_calibration/dead.parquet').values  # bool
    linear_corr = pd.read_parquet(f'{dname}/{prefix}_calibration/linear_corr.parquet').values.reshape(6, 32, fdim)

    # adc convert
    signal = signal / adc_info.loc[planet_id, f'{prefix}_adc_gain'] + adc_info.loc[planet_id, f'{prefix}_adc_offset']

    if conf.linear_corr:
        signal = apply_linear_correction(signal, linear_corr)

    # dark current subtraction
    if conf.dark:
        dt = dt_airs if sensor=='airs' else dt_fgs1
        signal -= dt.reshape(-1, 1, 1) * dark

    # correct flat field
    if conf.flat:
        signal /= flat

    # correlated double sampling (cds)
    signal = signal[1::2] - signal[::2]

    # time binning
    if conf.time_binning is not None and conf.time_binning > 0:
        binning_size = conf.time_binning * tmul
        nt = len(signal)
        signal = signal[:nt//binning_size*binning_size].reshape(-1, binning_size, 32, fdim).mean(1)
    
    # find hot or dead mask
    # hot_or_dead = sigma_clip(dark, sigma=5, maxiters=5).mask | dead
    hot_or_dead = sigma_clip_compat(dark, sigma=5) | dead

    if sensor == 'airs':
        signal = signal[:, 10:22, 39:321]
        hot_or_dead = hot_or_dead[10:22, 39:321]
    else:
        signal = signal[:, 10:22, 10:22]
        hot_or_dead = hot_or_dead[10:22, 10:22]

    if conf.hot_or_dead:
        ii = np.zeros_like(signal, dtype=bool) | hot_or_dead
        signal[ii] = np.nan
        signal = np.nanmean(signal, axis=maxis)
    else:
        signal = signal.mean(axis=maxis)

    return signal.astype(dtype)

if __name__=='__main__':
    import random
    planet_id = random.choice(train_planets)
    conf = CalibConfig(time_binning=30)
    airs_signal = read_signal(planet_id, sensor='airs', conf=conf)
    fgs1_signal = read_signal(planet_id, sensor='fgs1', conf=conf)
    print(airs_signal.shape, fgs1_signal.shape)
    
