# from glob import glob
from pathlib import Path
from re import sub
from datetime import datetime, timedelta
import numpy as np
from scipy.interpolate import CubicSpline
from cosmic2.tools import dt2float, read_gnss_v1
from traceback import extract_tb
from shutil import copyfile

working_dir = Path(__file__).absolute().parent if '__file__' in dir() else Path.cwd()
def log_message(msg, working_leo):
    with open(working_dir / f"log/progress_rt_{working_leo}_{datetime.now().strftime('%Y%m%d')}.txt", 'a') as fid:
        fid.write(f"{datetime.now().strftime('%H:%M:%S')} {msg}\n")
# lsat_names = dict(C2E1=1, C2E2=2, C2E3=3, C2E4=4, C2E5=5, C2E6=6)
# for keys in list(lsat_names.keys()):
#     lsat_names[lsat_names[keys]] = keys

def cal_zenith(dt, leosat):
    leofile = working_dir / f's01_leoOrb/leoOrb_{leosat}.{dt:%Y.%j}.npz'
    with np.load(leofile, allow_pickle=True) as fid:
        Ldata = fid['data']
        Ldt = fid['dts']
    Gdata, Gdt, Gsats = read_gnss_v1(Ldt[[0,-1]])
    idx = np.all(np.isfinite(Gdata[:-1, :, :3]), axis=(0, 2))
    Gdata = Gdata[:, idx]
    Gsats = Gsats[idx]
    Tbdry = [max(Gdt[0],Ldt[0]), min(Gdt[-1],Ldt[-1])]
    Tbdry_i = [x.replace(second=10*(x.second//10), microsecond=0) for x in Tbdry]
    if Tbdry_i[0] != Tbdry[0]:
        Tbdry_i[0] = Tbdry_i[0] + timedelta(seconds=10)
    Tnew = np.arange(Tbdry_i[0], Tbdry_i[1] + timedelta(seconds=10), timedelta(seconds=10)).astype(datetime)
    Gnew = CubicSpline(dt2float(Gdt[:-1], Tnew[0]), Gdata[:-1, :, :3])(dt2float(Tnew, Tnew[0]))
    Lnew = CubicSpline(dt2float(Ldt, Tnew[0]), Ldata[:, :3])(dt2float(Tnew, Tnew[0]))[:, None, :]
    LG = Gnew - Lnew
    zenith = np.arccos(np.sum(LG * Lnew, axis=2) /
                       np.sum(LG * LG, axis=2) ** 0.5 / np.sum(Lnew * Lnew, axis=2) ** 0.5) / np.pi * 180
    return zenith, Tnew, Gsats

def gnss_priority(gtype):
    if gtype == 'G':
        return -1
    elif gtype == 'R':
        return -2
    elif gtype == 'E':
        return -3
    else:
        return -4

def pair_ro_ref(leo, working_leo):
    rofiles = np.hstack([list((working_dir / f't01_RO_{working_leo}').glob(f'*{leo}*.npz')),
                         list((working_dir / f't03_ROhold_{working_leo}').glob(f'*{leo}*.npz'))])
    reffiles = np.array(list((working_dir / f't01_REF_{working_leo}').glob(f'*{leo}*.npz')))
    ro_dts = np.array([datetime.strptime(sub(r'.*[_\.](\d{4}\.\d{3})\..*',r'\1',rofile.name), '%Y.%j') for rofile in rofiles])
    ref_dts = np.array([datetime.strptime(sub(r'.*[_\.](\d{4}\.\d{3})\..*', r'\1', reffile.name), '%Y.%j') for reffile in reffiles])
    uniq_dts = np.unique(ro_dts)
    for dt in uniq_dts:
        rofiled = rofiles[ro_dts==dt]
        reffiled = reffiles[ref_dts==dt]
        if len(rofiled)==0 or len(reffiled)==0:
            continue
        gnss_ref = np.expand_dims([sub(r'.*\.(...)\.npz', r'\1', x.name) for x in reffiled], axis=0)
        dts_ro = np.expand_dims([datetime.strptime(sub(r'.*\.(\d{14})\.\d{14}\..*', r'\1', x.name),'%Y%m%d%H%M%S') for x in rofiled], axis=1)
        dte_ro = np.expand_dims([datetime.strptime(sub(r'.*\.\d{14}\.(\d{14})\..*', r'\1', x.name),'%Y%m%d%H%M%S') for x in rofiled], axis=1)
        dts_ref = np.expand_dims([datetime.strptime(sub(r'.*\.(\d{14})\.\d{14}\..*', r'\1', x.name),'%Y%m%d%H%M%S') for x in reffiled], axis=0)
        dte_ref = np.expand_dims([datetime.strptime(sub(r'.*\.\d{14}\.(\d{14})\..*', r'\1', x.name),'%Y%m%d%H%M%S') for x in reffiled], axis=0)
        try:
            zenith, Tnew, Gsats = cal_zenith(dt, leo)
        except Exception as e:
            log_message(f'-- Error detected for {leo}.{dt:%Y-%m-%d} ({len(rofiled)} files on hold)', working_leo)
            # TC: updated 2026/07/27 (modified 5 lines into 7 lines)
            for rofile in rofiled:
                if 'hold' not in rofile.parent.name:
                    copyfile(rofile,rofile.parent.with_name(f't03_ROhold_{working_leo}')/rofile.name)
            log_message(f'---  {type(e).__name__} >> {str(e)}', working_leo)
            for tb in extract_tb(e.__traceback__):
                log_message(f"---  line {tb.lineno} of {tb.filename}", working_leo)
            continue
        # with np.load(f's02_zenith/zenith_{leosat}.{dt:%Y.%j}_.npz', allow_pickle=True) as fid:
        #     zenith = fid['zenith']
        #     Tnew = fid['Tnew']
        #     Gsats = fid['Gsats']
        pidx = (dts_ro>dts_ref+timedelta(seconds=30))*(dte_ro<dte_ref-timedelta(seconds=30))*(np.isin(gnss_ref,Gsats))
        for rofile,dts,dte,ppidx in zip(rofiled,dts_ro[:,0],dte_ro[:,0],pidx):
            if ~np.any(ppidx):
                continue
            reffile = reffiled[ppidx]
            tidx = np.where((Tnew>dts)*(Tnew<dte))[0]
            if len(tidx)==0:
                continue
            lidx = [list(Gsats).index(sub(r'.*\.(...)\.npz', r'\1', x.name)) for x in reffile]
            refzenith = np.max(zenith[tidx][:, lidx], axis=0)
            refdt = Tnew[tidx[np.argmax(zenith[tidx][:, lidx], axis=0)]]
            refsnr1 = np.full(reffile.shape, np.nan)
            refant = np.array([int(sub(r'.*\.A(\d\d)\..*',r'\1',x.name)) for x in reffile])
            for pp, rr in enumerate(reffile):
                with np.load(rr, allow_pickle=True) as fid:
                    data = fid['data']
                    header = fid['header']
                sdata = [data[:, x] for x, y in enumerate(header) if 'S1' in y][0]
                tdata = np.array([datetime(1980,1,6)+timedelta(seconds=x) for x in data[:,0]+data[:,1]])
                refsnr1[pp] = CubicSpline(dt2float(tdata,refdt[pp]),sdata)(0)
            idx = np.array(sorted(np.array([range(len(refsnr1)), refsnr1, -refant, [gnss_priority(x.name[-7]) for x in reffile]]).transpose(),
                                  key=lambda x: (x[3], x[2], x[1]), reverse=True))[:, 0].astype('int')
            reffile = reffile[idx]
            refzenith = refzenith[idx]
            refdt = refdt[idx]
            refsnr1 = refsnr1[idx]
            np.savez(sub(r'.*/',f't03_pair_{working_leo}/',rofile.as_posix()),reffile=reffile,refzenith=refzenith,refdt=refdt,refsnr1=refsnr1)
