import warnings
warnings.filterwarnings("ignore")
import math
import sys
import csv
import numpy as np
from netCDF4 import Dataset
from decimal import *
from tqdm import tqdm
import matplotlib
import matplotlib.pyplot as plt
import os
import accpostproc as accp
from PIL import Image
import lib
import regionfinder
import argparse
lib.plt_rcParams_sci_adv()

parser = argparse.ArgumentParser(description='')
parser.add_argument('--regions', type=str,nargs = '+',default =["CHN","ASEAN","E-ASI","EUR","NAFTA"])
parser.add_argument('--runmean', type=int, default = 14)
parser.add_argument('--starttime', type=int, default = 0)
parser.add_argument('--numbering', type=str, default = None)
parser.add_argument('--numbering_position', type=float,nargs='+', default = [0.03,0.95])
parser.add_argument('--ylower', type=float, default = -3.95)
parser.add_argument('--yupper', type=float, default = 3.95)
parser.add_argument('--yticks', type=float, nargs='+', default = [-2.5,0.0,2.5])
parser.add_argument('--endtime', type=int, default = None)
parser.add_argument('--extra', type=str, default = '')
parser.add_argument('--scenario', type=str, default = "WPTS")
parser.add_argument('--variable', type=str, default = "export")
parser.add_argument('--relation', type=str, default = "perc",choices = ["abs","perc"])

args = parser.parse_args()
scale = pow(10,6.)
seasons = [2000 + i for i in range(21)]
networks = [2000 + i for i in range(16)]

d = Dataset(f"blocs_{args.scenario}-0p25_kl-22_en-2015_2018-05-01-2019-01-01_ar-120_ul.nc")

time = d.variables['time'][args.starttime:args.endtime]
blocs = d.variables['region'][:]
number = math.ceil(len(args.regions) / 2.)


ix = [0,1]
ids = [0,2]

arr = np.array([[[1,2,3],[4,5,6],[7,8,9]],[[10,11,12],[13,14,15],[16,17,18]],[[19,20,21],[22,23,24],[25,26,27]]])

correct_ones = lib.get_rows_in_file("correct_finished.csv")




fig = plt.figure(figsize = (3.5,0.5+1.*number), dpi = 600)

print(f"Make ---- {args.variable} Plot")



quantities = np.full(shape=(len(networks),time.shape[0],blocs.shape[0],blocs.shape[0],len(seasons)),fill_value = np.nan) 
values = np.full(shape=(len(networks),time.shape[0],blocs.shape[0],blocs.shape[0],len(seasons)),fill_value = np.nan)
zeros = np.full(shape=(len(networks),blocs.shape[0],blocs.shape[0],len(seasons)),fill_value = np.nan)
for num_season,season in enumerate(seasons):
    for num_network,network in enumerate(networks):
        name = f"{args.scenario}-0p25_kl-22_en-{network}_{season}-05-01-{season+1}-01-01_ar-120_ul"
        if name in correct_ones:
            d = Dataset(f"blocs_{name}.nc")
            time_shape = d.variables['time'][:].shape[0]
            quantities[num_network,:,:,:,num_season] = (d.groups[f"regions"].variables[f"trade_quantity"][args.starttime:args.endtime,:,:] - d.groups[f"regions"].variables[f"trade_quantity"][0,:,:])/ scale
            values[num_network,:,:,:,num_season] = (d.groups[f"regions"].variables[f"trade_value"][args.starttime:args.endtime,:,:] - d.groups[f"regions"].variables[f"trade_value"][0,:,:] )/ scale
            
            
            zeros[num_network,:,:,num_season] = d.groups[f"regions"].variables[f"trade_quantity"][0,:,:] / scale
            d.close()
            

q_positions = [n-0.2 for n in networks]
v_positions = [n+0.2 for n in networks]

for num, region in enumerate(tqdm(args.regions,desc="Main",leave = False)):
    ax = plt.subplot2grid((number,2), lib.number_to_row_of_two(num),colspan=1,rowspan = 1)

    if num % 2 != 0:
        ax.set_yticklabels([],fontsize = 0)
    ax.axhline(0,color = lib.colordict['gray'])

    
    ix = lib.subgroup_index(blocs,region)
    
    if "trade" not in args.variable and "port"  not in args.variable :
        q = np.nansum(quantities[:,ix,:],axis=(1))
        v = np.nansum(values[:,ix,:],axis=(1))
        z = np.nansum(zeros[:,ix,:],axis=(1))
    elif args.variable == "domestic_trade":
        
        q = np.nansum(np.nansum(quantities[:,ix,:,:],axis=(1))[:,ix,:],axis=(1))
        v = np.nansum(np.nansum(values[:,ix,:,:],axis=(1)),axis=(1))
        z = np.nansum(np.nansum(zeros[ix,:,:],axis=(0))[ix,:],axis=(0))
    elif args.variable == "foreign_trade":
        ix_rest = lib.get_rest(lib.subgroup_index(blocs,blocs),lib.subgroup_index(blocs,region))
        q = np.nansum(np.nansum(quantities[:,:,ix,:,:],axis=(2))[:,:,ix_rest,:],axis=(2)) + np.nansum(np.nansum(quantities[:,:,ix_rest,:,:],axis=(2))[:,:,ix,:],axis=(2))

        v = np.nansum(np.nansum(values[:,:,ix,:,:],axis=(2))[:,:,ix_rest,:],axis=(2)) + np.nansum(np.nansum(values[:,:,ix_rest,:,:],axis=(2))[:,:,ix,:],axis=(2))
        z = np.nansum(np.nansum(zeros[:,ix,:,:],axis=(1))[:,ix_rest,:],axis=(1)) + np.nansum(np.nansum(zeros[:,ix_rest,:,:],axis=(1))[:,ix,:],axis=(1))
    elif args.variable == "export":
        ix_rest = lib.get_rest(lib.subgroup_index(blocs,blocs),lib.subgroup_index(blocs,region))
        q = np.nansum(np.nansum(quantities[:,:,ix,:,:],axis=(2))[:,:,ix_rest,:],axis=(2))
        v = np.nansum(np.nansum(values[:,:,ix,:,:],axis=(2))[:,:,ix_rest,:],axis=(2))
        z = np.nansum(np.nansum(zeros[:,ix,:,:],axis=(1))[:,ix_rest,:],axis=(1))
    elif args.variable == "import":
        ix_rest = lib.get_rest(lib.subgroup_index(blocs,blocs),lib.subgroup_index(blocs,region))

        q = np.nansum(np.nansum(quantities[:,:,ix_rest,:,:],axis=(2))[:,:,ix,:],axis=(2))
        v = np.nansum(np.nansum(values[:,:,ix_rest,:,:],axis=(2))[:,:,ix,:],axis=(2))
        z = np.nansum(np.nansum(zeros[:,ix_rest,:,:],axis=(1))[:,ix,:],axis=(1))
    

   
    
    data_q = []
    data_v = []
    for en,network in enumerate(networks):
        notnan = np.where(z[en,:]/z[en,:]  == 1.)[0]
        z_c = z[en,:]
        z_c = z_c[notnan]
        q_c = q[en,:,:]
        q_c = q_c[:,notnan]
        v_c = v[en,:,:]
        v_c = v_c[:,notnan]
            
        q_en = 100 * np.nansum(q_c[:,:],axis=(0)) / (time.shape[0]*z_c[:])
        v_en = 100 * np.nansum(v_c[:,:],axis=(0)) / (time.shape[0]*z_c[:])
            
        data_q.append(q_en)
        data_v.append(v_en)
    ax.boxplot(data_q,positions = q_positions,whis = (5,95),showfliers=False,widths = 0.25,medianprops =  dict(color=lib.colordict['blue_alt']),boxprops = dict(linewidth = 0.4),whiskerprops = dict(linewidth = 0.4), capprops = dict(linewidth = 0.4))
    ax.boxplot(data_v,positions = v_positions,whis = (5,95),showfliers=False,widths = 0.25,medianprops =  dict(color=lib.colordict['orange_alt']),boxprops = dict(linewidth = 0.4),whiskerprops = dict(linewidth = 0.4), capprops = dict(linewidth = 0.4))
    ax.scatter([],[],color = lib.colordict['black'],label = lib.bloc_dic[region], s = 0)
    ax.legend(loc = 'best')
    ax.set_xlim(1999,2016)
    ax.set_xticks([2000,2002,2004,2006,2008,2010,2012,2014])
    ax.set_yticks(args.yticks)
    if math.ceil(num / 2.) != number -1 :
        ax.set_xticklabels([],fontsize = 0)
    else:
        ax.set_xlabel("Econonmic network [year]")
        ax.set_xticklabels(["",2002,"",2006,"",2010,"",2014])

    ax.set_ylim(args.ylower,args.yupper)
    if num == 0:
        if len(args.regions) % 2 == 0:
            tt = [None,None]
            ll = ['Quantity','Value']
            tt[0], = plt.plot([],[],color = lib.colordict['blue_alt'],linewidth = 1.2)
            tt[1], = plt.plot([],[],color = lib.colordict['orange_alt'],linewidth = 1.2)
            legend2 = plt.legend(tt,ll,loc = 'center',bbox_to_anchor =(1.,1.1),ncol = 2)
            ax.add_artist(legend2)



if len(args.regions) % 2 != 0:
    ax = plt.subplot2grid((number,2), lib.number_to_row_of_two(num+1),colspan=1,rowspan = 1)
    ax.plot([],[],label = "Quantity",color = lib.colordict['blue_alt'],linewidth = 1.2)
    ax.plot([],[],label = "Value",color = lib.colordict['orange_alt'],linewidth = 1.2)
    ax.legend(loc='lower center',handlelength = 0.8)
    ax.set_axis_off()

if args.relation == "perc":
    fig.text(0.02,0.53,f"{lib.label_quantity(args.variable)} change [\%]",rotation = 'vertical',verticalalignment = 'center',horizontalalignment = 'center')
else:
    fig.text(0.02,0.53,f"{lib.label_quantity(args.variable)} change [bn USD]",rotation = 'vertical',verticalalignment = 'center',horizontalalignment = 'center')


plt.subplots_adjust(hspace=0.05, wspace=0.05, left= 0.15, top = 0.94, right = 0.99, bottom = 0.1)

outputname = f"network_ms_{args.scenario}_{args.variable}_{args.relation}_0p25{args.extra}.png"    
fig.savefig(outputname)
fig.savefig(outputname.replace("png","pdf"))    

exit()
