import math
import numpy as np
import sys
import time
from datetime import datetime
import csv
import os
from tqdm import tqdm
from netCDF4 import Dataset
from scipy import optimize, stats
import scipy
import matplotlib.pyplot as plt
import matplotlib as mpl
import random

colordict = {
    'orange_alt': (1.0, 0.4980392156862745, 0.054901960784313725),
    'orange_PIK': (227/255,114/255,34/255),
    'red_alt': (0.8392156862745098, 0.15294117647058825, 0.1568627450980392),
    'green_alt': (0.17254901960784313, 0.6274509803921569, 0.17254901960784313),
    'blue_alt': (0.12156862745098039, 0.4666666666666667, 0.7058823529411765),
    'blue_PIK': '#009FDA',
    'purple': '#9467bd',
    'brown': '#8c564b',
    'magenta': '#e377c2',
    'lime': '#bcbd22',
    'cyan': '#17becf',
    'ochre':(0.854901961, 0.647058824, 0.125490196),
    'red_grey': (0.901960784,0.784313725,0.784313725),
    'blue_grey': (0.784313725,0.784313725,0.901960784),
    'grey': '#7f7f7f',
    'black': '#181818',
    'white': '#ffffff',
    'rosa':'#ff007f',
    'blue': '#1f77b4',
    'orange': '#ff7f0e',
    'green': '#2ca02c',
    'yellow':(1,1,0),
    'red': '#d62728',
    'gray': '#7f7f7f',
    'darkblue': '#00008b',
    'darkred': (162./255.,39./255.,48./255),
    'lightblue': '#aec7e8',
    'lightorange': '#ffbb78',
    'lightgreen': '#98df8a',
    'lightred': '#ff9896',
    'lightpurple': '#c5b0d5',
    'lightbrown': '#c49c94',
    'lightmagenta': '#f7b6d2',
    'lightgray': '#c7c7c7',
    'lightgrey': '#c7c7c7',
    'lightlime': '#dbdb8d',
    'lightcyan': '#9edae5',
    'lighterpurple': '#ccccff'
}


bloc_dic = {"CHN":"China","EUR":"Europe","ASEAN":"ASEAN", "E-ASI": "East Asia", "NAFTA" : "NAFTA", "ROW": "RoW", "SAARC":"SAARC","S-ASI":"South Asia","OCE":"Australia \& Oceania","OTHER": "Other countries","L-AM":"Latin America","SS-AFR":"Sub-Saharan Africa","EX-SOWJ":"Post-Soviet states","ARAB":"Arab League", "CHN_ECNS":"\"Northern\" China","CHN_SCNS":"\"Southern\" China" }

alphabet_array = np.array(['a','b','c','d','e','f','g','h','i','j','k','l','m','n','o','p','q','r','s','t','u','v','w','x','y','z'])


group_dic = {'production_quantity':'firms','production':'firms','production_value':'firms','demand':'firms','demand_quantity':'firms','demand_value':'firms','direct_loss':'firms','direct_loss_quantity':'firms','direct_loss_value':'firms','total_loss':'firms','total_loss_quantity':'firms','total_loss_value':'firms','gdp':'regions','gdp_quantity':'regions','gdp_value':'regions','consumption':'regions','consumption_quantity':'regions','consumption_value':'regions','total_flow':'regions','total_flow_quantity':'regions','total_flow_value':'regions','outflow':'regions','outflow_quantity':'regions','outflow_value':'regions'}


def plt_rcParams_sci_adv():
    plt.rcParams.update({
    'font.size': 9,
    'axes.labelsize': 9,
    'legend.fontsize': 7,
    'xtick.labelsize': 7,
    'ytick.labelsize': 7,
    'lines.markeredgewidth': 0,
    'legend.framealpha': 1.,
    'legend.frameon': False,
    'lines.markersize': 1,
    'lines.linewidth': 1.,
    'xtick.direction': 'in',
    'ytick.direction': 'in',
    'xtick.minor.size': 2,
    'ytick.minor.size': 2,
    'xtick.major.size': 3,
    'ytick.major.size': 3,
    'xtick.minor.width': .5,
    'xtick.major.width': .5,
    'ytick.minor.width': .5,
    'ytick.major.width': .5,
    'axes.linewidth': .5,
    'axes.labelpad': 0.5,
    'text.usetex': True,
    'font.serif': 'Computer Modern Roman',
    'font.monospace': 'Computer Modern Typewriter',
    'text.latex.preamble': r'\usepackage{amsmath,amssymb}\usepackage{sansmath}\sansmath',
    'hatch.color': 'lightgrey',
    'hatch.linewidth': 1.,

    })
    
def get_rows_in_file(f,column_number = 0,delimiter = ',',replacer = None,starting_row = 0):
    with open(f, 'r') as csvfile:
        reader = csv.reader(csvfile,delimiter = delimiter)
        table = []
        for num,r in enumerate(reader):
            if num < starting_row:
                continue
            if len(r) != 0 and r[column_number] != '':
                table.append(r[column_number])
            else:
                if replacer is not None:
                    table.append(replacer)
    table = np.array(table)
    return table
    
def netcdf_nan_correcter(arr,replace_value = 0.):
    idl = np.argwhere(arr/arr == 1.0)
    decide_array = np.full(shape = arr.shape,fill_value= True)
    for mem in idl:
        decide_array[tuple(mem)] = False
    arr_return = arr
    arr_return[decide_array] = replace_value
    return arr_return


def subgroup_index(group,subgroup):
    return np.where(np.isin(group,subgroup))[0]

def get_rest(group,unwanted):
    return np.setdiff1d(group,unwanted)
    
def label_quantity(string,differentiate = False):
    if 'gdp' in string:
        return_string = string[0:3].upper() + string[3:]
    else:
        return_string = string[0].upper() + string[1:]
    if differentiate:
        if ('_value' in string):
            return_string = return_string.replace('_value',' (v)')    
        else:
            return_string = return_string + ' (c)'
    else:
        return_string = return_string.replace('_value',' (v)')
    
    
    return return_string.replace('_',' ')

def number_to_row_of_two(number,alignment='vertical'):
    tup = [math.floor(number/2.),number % 2]
    return tup
    


def get_colors(regions):
    color_list = []

    for region in regions:
        if region == "USA" or region == "NAFTA":
            color_list.append(colordict['orange_alt'])
        elif region == "EU" or region == "EU28" or region == "EUR":
            color_list.append(colordict['blue_alt'])
        elif region == "CHN":
            color_list.append(colordict['red_alt'])
        elif region == "ALL":
            color_list.append(colordict['black'])
        elif region == "S-ASI"  or region == 'SAARC':
            color_list.append(colordict['lightpurple'])
        elif region == "E-ASI":
            color_list.append(colordict['cyan'])
        elif region == "ARAB":
            color_list.append(colordict['yellow'])
        elif region == "SS-AFR" or region == "AFR":
            color_list.append(colordict['magenta'])
        elif region == "ASEAN":
            color_list.append(colordict['brown'])
        elif region == "RUS" or region == 'EX-SOWJ':
            color_list.append(colordict['ochre'])
        elif region == "MERCOSUR" or region == "S-AM" or region == "L-AM":
            color_list.append(colordict['green_alt'])
        elif region == "OCE" or region == "ROW":
            color_list.append(colordict['purple'])
    return np.array(color_list)
