#Code for computing ion-neutral collision and plotting
#By Amadi Brians Chinonso

#IMPORT MODULES AND PACKAGES
import scipy
import netCDF4
import datetime
import statistics
import numpy as np
import xarray as xr
from netCDF4 import Dataset
from matplotlib import rcParams
import matplotlib.pyplot as plt
import matplotlib.dates as dates
from mpl_toolkits.mplot3d import Axes3D
from matplotlib.dates import HourLocator, MinuteLocator, DateFormatter

import warnings
warnings.filterwarnings("ignore")


def time_conv(variable):
    wac_time_2 = variable
    #print(wac_time_2)

    #convert the integer to HH:MM:SS format
    w_time_2 = []

    for i in wac_time_2:

        time_int = i
        time = time_int/3600
        hours = int(time)
        minutes = (time*60) % 60
        seconds = (time*3600) % 60

        w_time_2 = np.append(w_time_2, "%d:%02d:%02d" % (hours, minutes, seconds))

    return w_time_2

def time_2d(wtime):
    wac_time_2d = []
    for i in wtime:
        waccy = i[0:5]
        wac_time_2d = np.append(wac_time_2d, waccy)

    return wac_time_2d
    
def sec_to_hm(timestamp, yy, mm, dd):
    # Sample masked numpy array
    masked_array = timestamp  # Example array with one masked value
    
    # Convert seconds to hours and minutes
    hours = masked_array / 3600
    minutes = (masked_array % 3600) / 60
    
    # Convert to datetime format
    dtv = []
    for h, m in zip(hours, minutes):
        if np.ma.is_masked(h) or np.ma.is_masked(m):
            dtv.append(None)  # If the value is masked, append None
        else:
            dtv.append(datetime.datetime(yy, mm, dd, int(h), int(m)))
    return dtv

def latym(sx):
    
    det = sx['datesec'][:]
    #CONVERT TO HOURS
    #CONVERT DATESEC TO HOUR
    time_int = det
    time = time_int/3600
    
    latt = sx['lat'][:][138:192]
    #lonn = wacod0['time'][:] * np.pi/180
    theta = time[:] / 24.0 * (2.0 * np.pi)
    #print(theta)
    theta, latt = np.meshgrid(theta, latt)
    #print(theta.shape)
    #print(latt.shape)
    
    mlatt = sx['mlat'][48:97]
    theta2 = time[:] / 24.0 * (2.0 * np.pi)
    theta2, mlatt = np.meshgrid(theta2, mlatt)
    #print(theta2.shape)
    #print(mlatt.shape)
    #print(wac7p['PHIM2D'][:, 48:97, 55].shape)
    
    X, Y = np.meshgrid(sx['lat'][:][63:128], sx['Z3'][0, 0:32, 63, 252]/1000)

    return theta, theta2, latt, mlatt, X, Y

def split_number(number):
    number_str = str(number)
    if len(number_str) != 8:
        return "Number must be 8 digits long"
    
    first_four = number_str[:4]
    middle_two = number_str[4:6]
    last_two = number_str[6:]
    
    return first_four, middle_two, last_two
    
#Adopted from Incoherent Scatter Radar (ISR) Summer School Materials organized 
#by NSF and MIT HayStack Observatory scientitsts and Researchers

def getCin(i,n,Ti=1000.0,Tn=1000.0):
    '''
    Non-resonant and resonant collision frequency coefficients (Cin x 10^10) from SN00 Table 4.4 and 4.5
    Collision frequencies can then be calculated as vin = Cin x 1e-10 x Nn where Nn is the neutral density in cm^-3
    
    Inputs:
        i - ion mass in amu (1=H+,4=He+,12=C+,14=N+,16=O+,29=CO+,28=N2+,30=NO+,32=O2+,44=CO2+)
        n - neutral mass in amu (1=H,4=He,14=N,16=O,29=CO,28=N2,32=O2,44=CO2)
        Ti - ion temp (for resonant collisions)
        Tn - neutral temp (for resonant collisions)

    Note for CO+, need to set i or n to 29 amu (because of duplicate N2+ value) 
    
    '''
    
    Amu2Ion = {'1':'H+','4':'He+','12':'C+','14':'N+','16':'O+','29':'CO+','28':'N2+','30':'NO+','32':'O2+','44':'CO2+'}
    Amu2Ntrl = {'1':'H','4':'He','14':'N','16':'O','29':'CO','28':'N2','32':'O2','44':'CO2'}

    # resonant
    Tr = (Ti+Tn)/2.0
    HpH = 2.65*Tr**0.5*(1.0-0.083*np.log10(Tr))**2.0
    HpO = 0.661*Ti**0.5*(1.0-0.047*np.log10(Ti))**2.0
    HepHe = 0.873*Tr**0.5*(1.0-0.093*np.log10(Tr))**2.0
    NpN = 0.383*Tr**0.5*(1.0-0.063*np.log10(Tr))**2.0
    OpH = 0.661*Ti**0.5*(1.0-0.047*np.log10(Ti))**2.0
    OpO = 0.367*Tr**0.5*(1.0-0.064*np.log10(Tr))**2.0
    COpCO = 0.342*Tr**0.5*(1.0-0.085*np.log10(Tr))**2.0
    N2pN2 = 0.514*Tr**0.5*(1.0-0.073*np.log10(Tr))**2.0
    O2pO2 = 0.259*Tr**0.5*(1.0-0.063*np.log10(Tr))**2.0
    CO2pCO2 = 0.285*Tr**0.5*(1.0-0.083*np.log10(Tr))**2.0
    
    Cin = {
        'H+':   {'H':HpH, 'He':10.6, 'N':26.1,'O':HpO, 'CO':35.6, 'N2':33.6, 'O2':32.0, 'CO2':41.4},
        'He+':  {'H':4.71,'He':HepHe,'N':11.9,'O':10.1,'CO':16.9, 'N2':16.0, 'O2':15.3, 'CO2':20.0},
        'C+':   {'H':1.69,'He':1.71, 'N':5.73,'O':4.94,'CO':8.74, 'N2':8.26, 'O2':8.01, 'CO2':10.7},
        'N+':   {'H':1.45,'He':1.49, 'N':NpN, 'O':4.42,'CO':7.90, 'N2':7.47, 'O2':7.25, 'CO2':9.73},
        'O+':   {'H':OpH, 'He':1.32, 'N':4.62,'O':OpO, 'CO':7.22, 'N2':6.82, 'O2':6.64, 'CO2':8.95},
        'CO+':  {'H':0.74,'He':0.79, 'N':2.95,'O':2.58,'CO':COpCO,'N2':4.24, 'O2':4.49, 'CO2':6.18},
        'N2+':  {'H':0.74,'He':0.79, 'N':2.95,'O':2.58,'CO':4.84, 'N2':N2pN2,'O2':4.49, 'CO2':6.18},
        'NO+':  {'H':0.69,'He':0.74, 'N':2.79,'O':2.44,'CO':4.59, 'N2':4.34, 'O2':4.27, 'CO2':5.89},
        'O2+':  {'H':0.65,'He':0.70, 'N':2.64,'O':2.31,'CO':4.37, 'N2':4.13, 'O2':O2pO2,'CO2':5.63},
        'CO2+': {'H':0.47,'He':0.51, 'N':2.00,'O':1.76,'CO':3.40, 'N2':3.22, 'O2':3.18, 'CO2':CO2pCO2},
        }

    try:
        i=int(i)
        n=int(n)
        ion=Amu2Ion[str(i)]
        ntrl=Amu2Ntrl[str(n)]        
        return Cin[ion][ntrl]
    except:
        return 0.0
        
#Import files
sx1 = Dataset('/glade/campaign/hao/itmodel/bamadi/archive/f.e22.FXSD.f09_f09_mg17_Feb22_5MinOut.001/atm/hist/f.e22.FXSD.f09_f09_mg17_Feb22_5MinOut.001.cam.h1.2022-02-01xc.nc')
sx2 = Dataset('/glade/campaign/hao/itmodel/bamadi/archive/f.e22.FXSD.f09_f09_mg17_Feb22_5MinOut.001/atm/hist/f.e22.FXSD.f09_f09_mg17_Feb22_5MinOut.001.cam.h1.2022-02-02xc.nc')
sx3 = Dataset('/glade/campaign/hao/itmodel/bamadi/archive/f.e22.FXSD.f09_f09_mg17_Feb22_5MinOut.001/atm/hist/f.e22.FXSD.f09_f09_mg17_Feb22_5MinOut.001.cam.h1.2022-02-03xc.nc')
sx4 = Dataset('/glade/campaign/hao/itmodel/bamadi/archive/f.e22.FXSD.f09_f09_mg17_Feb22_5MinOut.001/atm/hist/f.e22.FXSD.f09_f09_mg17_Feb22_5MinOut.001.cam.h1.2022-02-04xc.nc')
sx5 = Dataset('/glade/campaign/hao/itmodel/bamadi/archive/f.e22.FXSD.f09_f09_mg17_Feb22_5MinOut.001/atm/hist/f.e22.FXSD.f09_f09_mg17_Feb22_5MinOut.001.cam.h1.2022-02-05xc.nc')

rcParams['font.weight'] = 'bold'
coll = ['r', 'g', 'k', 'orange', 'blue']
rcParams['font.size'] = '10'

#Plot Joule Heating and Electron density
fig, (ax1, ax2, ax3, ax4,ax5, ax6, ax7, ax8, ax9, ax10) = plt.subplots(10, 1, figsize=(30, 40))
plt.subplots_adjust(wspace = 0.05)

pad = 0.035
fr = 0.045
fst = 20
lp = 20
fsc = 15

for j in np.arange(len(sxx)):
    f = sxx[j]
    sx1t, sx1t2, sx1l, sx1ml, sx1X, sx1Y = latym(f)

    ax1 = plt.subplot(5,2,(2*j)+1, projection='polar')
    ax2 = plt.subplot(5,2,2*(1+j), projection='3d')

    vmin1 = 0
    vmax1 = 0.125
    levels1 = np.linspace(vmin1, vmax1)
    im1 = ax1.contourf(sx1t[:], sx1l[::-1], f['QJOULE'][:, 12, 138:192, 260].T, levels = levels1, cmap = 'rainbow')
    ax1.set_theta_offset(3 * np.pi/2)
    ax1.set_xticklabels(['{:.1f}'.format(xlabel) \
                        for xlabel in np.arange(0.0,24.0,(24.0 / len(ax1.get_xticklabels())))])
    ax1.set_yticklabels(np.arange(80, 0, -10), color = 'white')
    cb1 = fig.colorbar(im1, ax = ax1, fraction = fr, pad = 0.035)
    cb1.formatter.set_powerlimits((0, 0))
    cb1.ax.set_ylabel('$K/s$', rotation=270, labelpad=lp, fontsize = fsc)
    
    vmin3 = 0
    vmax3 = 2.8E6
    levels3 = np.linspace(vmin3, vmax3)
    for i in np.arange(144, 288, 60):
        im3 = ax2.contourf(sx1X, sx1Y, f['EDens'][i, 0:32, 63:128, 260], 50, levels = levels3, zdir='z', offset= i/12, cmap=plt.get_cmap('jet'))    
        #levels=np.linspace(Z1.min(), Z1.max(), 100)
        if i == 264:
            cb3 = fig.colorbar(im3, ax = ax2, fraction = fr, pad=0.01)
            cb3.formatter.set_powerlimits((0, 0))
            cb3.ax.set_ylabel('${cm^{-3}}$', rotation=270, labelpad=lp, fontsize = fsc)
        #ax.contourf(X, Y, wacod0['EDens'][6, :, 63:128, 252], 50, zdir='z', offset=6, cmap=plt.get_cmap('rainbow'))
    
    # Set labels and title
    ax2.set_xlabel('Latitude')
    ax2.set_ylabel('Altitude')
    ax2.set_zlabel('Hour (UT)')
    #if j == 0:
        #ax2.set_title('Electron density', fontsize = fst)
    ax2.set_zlim3d(12, 24)

fig.text(0.4, 0.9, 'Joule heating', ha='center', fontsize = '30', weight = 'bold')
fig.text(0.7, 0.9, 'Electron density', ha='center', fontsize = '30', weight = 'bold')
plt.savefig("spacex_fig.jpg", bbox_inches="tight",
            )
plt.show()

#FIND MAX VALUE OF QJOULE
for j in np.arange(len(sxx)):
    f = sxx[j]
    print(f['QJOULE'][:, 12, 138:192, 260].max())
    
#Ion-neutral col Coefficients

#for O
nu_coo = getCin(16, 16)

#for O2
nu_coo2 = getCin(16, 32)

O = sxxww[0]['O'][144:, 10, 96, 280]
O2 = sxxww[0]['O2'][144:, 10, 96, 280]
#Ion-neutral col
nu_op = nu_coo*O + nu_coo2*O2
nu_op

#PEDERSEN CONDUCTANCE
coll = ['g', 'k', 'orange', 'red']
rcParams['font.weight'] = 'bold'
rcParams['font.size'] = '10'
fig, (ax1, ax2, ax3) = plt.subplots(3, 1, figsize=(10, 12), sharex = True)
for k in np.arange(len(sxxww)):
    
    ds = sxxww[k]
    det = sxxww[k]['date'][2]
    yy, mm, dd = split_number(det)
    
    O = ds['O'][200:, 10, 96, 280]
    O2 = ds['O2'][200:, 10, 96, 280]
    #Ion-neutral col
    nu_op = nu_coo*O + nu_coo2*O2
    
    t = sec_to_hm(ds['datesec'][:], 2022, 2, 2)
    tt = t[:][200:]

    ax1.plot(tt, ds['ED1'][200:,40, 47]*1.e+4, color = coll[k],
             label = str(dd) + '/' + str(mm) + '/' + str(yy))
    
    ax2.plot(tt, ds['EDYN_ZIGM11_PED'][200:,40, 47], color = coll[k], 
             label = str(dd) + '/' + str(mm) + '/' + str(yy))

    ax3.plot(tt, nu_op, color = coll[k],
         label = str(dd) + '/' + str(mm) + '/' + str(yy)) #I divided by 7.4 to normalize the col freq
    
    # Set x-axis format
    ax3.xaxis.set_major_locator(HourLocator(interval=2))
    ax3.xaxis.set_minor_locator(MinuteLocator(interval=30))
    ax3.xaxis.set_major_formatter(DateFormatter('%H:%M'))  # Display only hour and minute

    #ZOOM PLOT
    axins = zoomed_inset_axes(ax2, 2.5, loc=1)
    axins.plot(tt, ds['EDYN_ZIGM11_PED'][200:,40, 47], color = coll[k])
    axins.set_xlim(tt[45], tt[60])
    axins.set_ylim(0, 20)
    mark_inset(ax2, axins, loc1=2, loc2=4, fc="none", ec="0.75")
    axins.set_xticklabels([])
    axins.set_xticks([])
    axins.set_yticklabels([])
    axins.set_yticks([])
    axins.set_facecolor('none')
    #ax2.set_ylim(0, 20)

    if k == len(sxxww)-1:
        ax1.set_title('a) Eastward Electric Field', weight = 'bold')
        ax1.set_ylabel('E (mV/m*1e-4)', weight = 'bold')
        ax2.set_ylabel(r' ${\sigma_p}$ (S)', weight = 'bold')
        ax3.set_ylabel(r' ${\nu_{in}}$', weight = 'bold')
        ax2.set_title('b) Pedersen Conductance', weight = 'bold')
        ax3.set_title('c) Ion-neutral collision', weight = 'bold')
        ax1.legend()
        #ax2.legend()
        ax1.grid()
        ax2.grid()
        ax3.grid()
        ax3.set_xlim(t[:][200],t[-1])
        ax3.set_xlabel('Time (UT)', weight = 'bold')
        
   
#print(tt[20])
#plt.legend()
plt.show()
