from __future__ import print_function
import numpy as np
import sys
import os
import random
from random import seed
from random import random
from random import gauss
import datetime
import matplotlib.pyplot as plt
from matplotlib.backends.backend_pdf import PdfPages
from matplotlib.ticker import MultipleLocator
from matplotlib.ticker import ScalarFormatter
from matplotlib.ticker import NullFormatter
import matplotlib.image as image
import matplotlib.patches as patches
import matplotlib.lines as lines
import matplotlib.ticker as ticker
from scipy.optimize import curve_fit

def define_action_times(dt):
	action_times=np.zeros((5),dtype=np.int)
	for a in range(1,6):
		taui = 1.0/4**a
		action_times[a-1] = (int) (taui/dt)
	return action_times

def make_figure (factorx,dt,avg_time_passive,avg_time_casual,avg_time_optimal,Results):
	# data
	xdata, fraction_of_unfit_individuals, fraction_of_BP_individuals, fraction_of_advanced_individuals, number_of_advanced_individuals, advanced_individuals, avg_time_advanced, avg_fraction_time_BP_of_advanced_individuals, data_count_individuals_advanced, data_times, data_fractions = Results[0],Results[1],Results[2],Results[3],Results[4],Results[5],Results[6],Results[7],Results[8],Results[9],Results[10]

	# xlimits
	xlima,xlimb=-1.8*factorx, 40*factorx-2*factorx
	xtic,xtic_minor = 4*factorx,2*factorx
	xwidth = 2.5*factorx
	
	# benchmarks lines
	va_passive=[]
	va_casual=[]
	va_optimal=[]
	va_passive_f=[]
	va_casual_f=[]
	va_optimal_f=[]
	va_casual_BP=[]
	va_optimal_BP=[]
	va_casual_ABP=[]
	va_optimal_ABP=[]
	
	action_times = define_action_times(dt)
	va_casual_value = 0.0
	staying_prob = 0.5
	weighted_avg_action_time = 0.0
	for a_time in range(5):
		weighted_avg_action_time += action_times[a_time]*dt / 5
	for i in range(10000):
		value_to_add = (i+1) * staying_prob**i * (1.0-staying_prob)
		va_casual_value += value_to_add
		if (value_to_add < 0.00001): break
	va_casual_value *= weighted_avg_action_time
	print ('va_casual_value =', va_casual_value)
	
	a_time_opt0 = 4
	a_time_opt1 = 4
	
	for x in xdata:
		va_passive.append(avg_time_passive)
		va_casual.append(avg_time_casual)
		va_optimal.append(avg_time_optimal)
		va_passive_f.append(1.0)
		va_casual_f.append(0.5)
		va_optimal_f.append(0.5)
		va_casual_BP.append(va_casual_value)
		va_optimal_BP.append(dt*action_times[a_time_opt0])
		va_casual_ABP.append(va_casual_value)
		va_optimal_ABP.append(dt*action_times[a_time_opt1])
		
	va_passive=np.asarray(va_passive)
	va_casual=np.asarray(va_casual)
	va_optimal=np.asarray(va_optimal)
	va_passive_f=np.asarray(va_passive_f)
	va_casual_f=np.asarray(va_casual_f)
	va_optimal_f=np.asarray(va_optimal_f)
	va_casual_BP=np.asarray(va_casual_BP)
	va_optimal_BP=np.asarray(va_optimal_BP)
	va_casual_ABP=np.asarray(va_casual_ABP)
	va_optimal_ABP=np.asarray(va_optimal_ABP)
	
	# setto alcune variabili comuni
	axisticslabelfontsize=9
	axisticslabelfontsizeinset=7
	axislabelfontsize=11 
	axislabelfontsizeinset=9
	
	xfig,yfig=7.0,2.5
	factor = xfig/yfig
	
	with PdfPages('fig2.pdf') as pdf:
		fig = plt.figure(figsize=(xfig,yfig))
		plt.rc('text', usetex=True)
		plt.rc('text.latex', preamble = ','.join('''\usepackage{txfonts} \usepackage{lmodern}'''.split()))
		
		
		
		######################## panel A ########################################
		panel = fig.add_axes([0.08, 0.16, 0.31, 0.82])
		panel.tick_params(axis='both',which='both',direction='in',bottom=True,top=True,left=True,right=True)
		for tick in panel.xaxis.get_major_ticks(): tick.label.set_fontsize(axisticslabelfontsize)
		for tick in panel.yaxis.get_major_ticks(): tick.label.set_fontsize(axisticslabelfontsize)
		panel.set_xlabel(r'generation',fontsize=axislabelfontsize)
		panel.set_ylabel(r'$N_{\alpha}/N$',fontsize=axislabelfontsize)
		panel.set_xlim(xlima,xlimb)
		panel.set_ylim(-0.02, 1.02)
		panel.xaxis.set_major_locator(MultipleLocator(xtic))
		panel.yaxis.set_major_locator(MultipleLocator(0.2))
		panel.yaxis.set_minor_locator(MultipleLocator(0.1))
		t = panel.text(0.04,0.9,r'(a)',fontsize=axislabelfontsize,transform=panel.transAxes)
		t.set_bbox(dict(facecolor='white', alpha=0.8, edgecolor='none',pad=1.5))
		
		
		panel.plot(xdata,fraction_of_unfit_individuals,'r-',markersize=2,linewidth=1.5,zorder=3,label=r'ABP-like individuals')
		panel.plot(xdata,fraction_of_BP_individuals,'b--',markersize=2,linewidth=1.5,zorder=3,label=r'BP-like individuals')
		panel.plot(xdata,fraction_of_advanced_individuals,'g-',markersize=2,linewidth=1.5,zorder=3,label=r'switching individuals')
		
		panel.legend(loc='upper right', bbox_to_anchor=(0.99, 0.87),ncol=1,fontsize=8,handlelength=1.5,labelspacing=0.2)
		
		######################## panel A1 ########################################
		panel = fig.add_axes([0.2, 0.32, 0.18, 0.32])
		panel.tick_params(axis='both',which='both',direction='in',bottom=True,top=True,left=True,right=True)
		for tick in panel.xaxis.get_major_ticks(): tick.label.set_fontsize(axisticslabelfontsizeinset)
		for tick in panel.yaxis.get_major_ticks(): tick.label.set_fontsize(axisticslabelfontsizeinset)
		panel.set_xlabel(r'generation',fontsize=axislabelfontsizeinset,labelpad=-1)
		panel.set_ylabel(r'$N_{\alpha_{i,j}}/N$',fontsize=axislabelfontsizeinset,labelpad=-1)
		panel.set_xlim(xlima,xlimb)
		panel.set_ylim(-0.05, 1.05)
		panel.xaxis.set_major_locator(MultipleLocator(xtic))
		panel.yaxis.set_major_locator(MultipleLocator(0.2))
		panel.yaxis.set_minor_locator(MultipleLocator(0.1))
		t.set_bbox(dict(facecolor='white', alpha=0.8, edgecolor='none',pad=1.5))
		
		print (data_count_individuals_advanced[9][9][-1],data_count_individuals_advanced[4][9][-1],data_count_individuals_advanced[4][8][-1])
		
		value = 0
		i0,j0=0,0
		for i in range(5,10):
			for j in range(5,10):
				if (data_count_individuals_advanced[i][j][-1]>value): value,i0,j0=data_count_individuals_advanced[i][j][-1],i,j
				
		value = 0
		i1,j1=0,0
		for i in range(5,10):
			for j in range(5,10):
				if (i==i0 and j==j0):
					value2=0.0
				else:
					if (data_count_individuals_advanced[i][j][-1]>value): value,i1,j1=data_count_individuals_advanced[i][j][-1],i,j
					
		value = 0
		i2,j2=0,0
		for i in range(5,10):
			for j in range(5,10):
				if (i==i0 and j==j0 or i==i1 and j==j1):
					value2=0.0
				else:
					if (data_count_individuals_advanced[i][j][-1]>value): value,i2,j2=data_count_individuals_advanced[i][j][-1],i,j
					
		value = 0
		i3,j3=0,0
		for i in range(5,10):
			for j in range(5,10):
				if (i==i0 and j==j0 or i==i1 and j==j1 or i==i2 and j==j2):
					value2=0.0
				else:
					if (data_count_individuals_advanced[i][j][-1]>value): value,i3,j3=data_count_individuals_advanced[i][j][-1],i,j
					
		value = 0
		i4,j4=0,0
		for i in range(5,10):
			for j in range(5,10):
				if (i==i0 and j==j0 or i==i1 and j==j1 or i==i2 and j==j2 or i==i3 and j==j3):
					value2=0.0
				else:
					if (data_count_individuals_advanced[i][j][-1]>value): value,i4,j4=data_count_individuals_advanced[i][j][-1],i,j
		
		for index in range(len(xdata)):
			n = int(round(number_of_advanced_individuals[index]))
			n_individuals = int(round(n/fraction_of_advanced_individuals[index]))
		
		
		#~ panel.plot(xdata,data_count_individuals_advanced[i1][j1]/n_individuals,'b-',markersize=2,linewidth=1.,zorder=3,label=(r'$i=%d \; j=%d$' %(i1+1,j1+1)) )
		#~ panel.plot(xdata,data_count_individuals_advanced[i2][j2]/n_individuals,'r-',markersize=2,linewidth=1.,zorder=3,label=(r'$i=%d \; j=%d$' %(i2+1,j2+1)) )
		#~ panel.plot(xdata,data_count_individuals_advanced[i3][j3]/n_individuals,'k-',markersize=2,linewidth=1.,zorder=3,label=(r'$i=%d \; j=%d$' %(i3+1,j3+1)) )
		#~ panel.plot(xdata,data_count_individuals_advanced[i4][j4]/n_individuals,'m-',markersize=2,linewidth=1.,zorder=3,label=(r'$i=%d \; j=%d$' %(i4+1,j4+1)) )
		for i in range(5,10):
			for j in range(5,10):
				if (i==i0 and j==j0):
					value2=0.0
				else:
					panel.plot(xdata,data_count_individuals_advanced[i][j]/n_individuals,'-',color='gray',markersize=2,linewidth=0.5,zorder=3,label=(r'$i=%d \; j=%d$' %(i1+1,j1+1)) )
		
		panel.plot(xdata,data_count_individuals_advanced[i0][j0]/n_individuals,'g-',markersize=2,linewidth=1.,zorder=3,label=(r'$i=%d \; j=%d$' %(i0+1,j0+1)) )
		print (i0,j0)
		
		panel.annotate(r'$\alpha_{5,5}$', xy = (4.8,0.73), xycoords='data', xytext=(8, 0.5), textcoords='data', fontsize=8, arrowprops=dict(arrowstyle="->"),horizontalalignment='right', verticalalignment='top')
		
		######################## panel B ########################################
		panel = fig.add_axes([0.475, 0.16, 0.51, 0.82])
		panel.tick_params(axis='both',which='both',direction='in',bottom=True,top=True,left=True,right=True)
		for tick in panel.xaxis.get_major_ticks(): tick.label.set_fontsize(axisticslabelfontsize)
		for tick in panel.yaxis.get_major_ticks(): tick.label.set_fontsize(axisticslabelfontsize)
		panel.set_xlabel(r'generation',fontsize=axislabelfontsize)
		panel.set_ylabel(r'time to reach target $[\tau]$',fontsize=axislabelfontsize)
		panel.set_xlim(xlima,xlimb)
		panel.set_ylim(0.0, 6.15)
		panel.xaxis.set_major_locator(MultipleLocator(xtic))
		panel.yaxis.set_major_locator(MultipleLocator(1))
		panel.yaxis.set_minor_locator(MultipleLocator(0.5))
		t = panel.text(0.93,0.9,r'(b)',fontsize=axislabelfontsize,transform=panel.transAxes)
		t.set_bbox(dict(facecolor='white', alpha=0.8, edgecolor='none',pad=1.5))
		
		c = "green"
		for index in range(len(xdata)):
			x = xdata[index]
			n = int(round(number_of_advanced_individuals[index]))
			n_individuals = int(round(n/fraction_of_advanced_individuals[index]))
			print ('check: Number of individuals = ',n_individuals)
			data = np.zeros((n))
			j=0
			for i in range(n_individuals):
				if (advanced_individuals[i][index] == 1.0):
					data[j]=data_times[i][index]
					j+=1
			if (n>1):
				panel.boxplot(data, sym='', whis=[10,90], positions=[x], widths = xwidth, manage_xticks=False, autorange=False, patch_artist=True,boxprops=dict(facecolor='lightgreen', color=c),capprops=dict(color=c),whiskerprops=dict(color=c),medianprops=dict(color=c))
		
		panel.plot(xdata,avg_time_advanced,'go-',markersize=2,linewidth=1,zorder=3,label=r'switching individuals')
		panel.plot(xdata,va_passive,'b--',markersize=2,linewidth=1.5,zorder=3,label=r'passive particle (avg.)')
		panel.plot(xdata,va_casual,'m--',markersize=2,linewidth=2.0,zorder=3, label=r'casual particle (avg.)')
		panel.plot(xdata,va_optimal,'k--',markersize=2,linewidth=1.0,zorder=3,label=r'optimal particle (avg.)')
		
		panel.legend(loc='upper right', bbox_to_anchor=(0.93, 0.99),ncol=1,fontsize=8,handlelength=1.5,labelspacing=0.2)
		
		
		xr,yr=4.9*factorx,5.0
		dxr,dyr=5.8*factorx,0.3
		epsxr,epsyr = 1*factorx,0.42
		rect = patches.Rectangle((xr,yr),dxr,dyr, facecolor='lightgreen',edgecolor='g')
		panel.add_patch(rect)
		line = lines.Line2D([xr, xr-dxr/2], [yr+dyr/2, yr+dyr/2], color='g', linewidth=0.8)
		panel.add_line(line)
		line = lines.Line2D([xr-dxr/2,xr-dxr/2], [yr+dyr/3, yr+2*dyr/3], color='g', linewidth=0.8)
		panel.add_line(line)
		line = lines.Line2D([xr+dxr, xr+dxr+dxr/2], [yr+dyr/2, yr+dyr/2], color='g', linewidth=0.8)
		panel.add_line(line)
		line = lines.Line2D([xr+dxr+dxr/2,xr+dxr+dxr/2], [yr+dyr/3, yr+2*dyr/3], color='g', linewidth=0.8)
		panel.add_line(line)
		line = lines.Line2D([xr+dxr/2,xr+dxr/2], [yr, yr+dyr], color='g', linewidth=0.8)
		panel.add_line(line)
		panel.text(xr-dxr/2-epsxr,yr+epsyr,r'$10$\%',fontsize=8)
		panel.text(xr-epsxr,yr+epsyr,r'$25$\%',fontsize=8)
		panel.text(xr+dxr/2-epsxr,yr+epsyr,r'$50$\%',fontsize=8)
		panel.text(xr+dxr-epsxr,yr+epsyr,r'$75$\%',fontsize=8)
		panel.text(xr+dxr+dxr/2-epsxr,yr+epsyr,r'$90$\%',fontsize=8)
		
		
		pdf.savefig(fig)
	return

def read_results (evaluation_time,n_individuals_ini,n_generations):
	f=open('Results_Pe100.dat','r')
	lines=f.readlines()
	f.close()
	
	n_individuals = n_individuals_ini*2
	
	n_individuals_gen = np.zeros((n_generations))
	xdata = np.zeros((n_generations))
	fraction_of_unfit_individuals = np.zeros((n_generations))
	fraction_of_BP_individuals = np.zeros((n_generations))
	fraction_of_advanced_individuals = np.zeros((n_generations))
	number_of_advanced_individuals = np.zeros((n_generations))
	avg_time_advanced = np.zeros((n_generations))
	avg_fraction_time_BP_of_advanced_individuals = np.zeros((n_generations))
	data_times = np.zeros((n_individuals,n_generations))
	data_fractions = np.zeros((n_individuals,n_generations))
	data_count_individuals_advanced = np.zeros((10,10,n_generations))
	advanced_individuals = np.zeros((n_individuals,n_generations))
	

	
	for line in lines:
		p = line.split()
		if (len(p)<4):
			generation = int(p[-1])
			if (generation>0): print ('number of individuals in generation ',generation-1,' = ',individual+1)
			if (generation>=n_generations): break
			individual = -1
		else:
			individual += 1
			n_individuals_gen[generation]+=1.0
			genome_id,num_target,fraction_of_time_passive,a0,a1 = int(p[0]),float(p[1]),float(p[2]),int(p[3]),int(p[4])
			if (individual<n_individuals):
				if(a1<=5 and a0>5):
					fraction_of_unfit_individuals[generation]+=1.0
				elif(a1<=5 and a0<=5):
					fraction_of_BP_individuals[generation]+=1.0
				else:
					if (fraction_of_time_passive>0.999):
						fraction_of_BP_individuals[generation]+=1.0
					else:
						avg_time_advanced[generation]+=evaluation_time/num_target
						number_of_advanced_individuals[generation]+=1.0
						avg_fraction_time_BP_of_advanced_individuals[generation]+=fraction_of_time_passive
						advanced_individuals[individual][generation] = 1.0
				data_count_individuals_advanced[a0-1][a1-1][generation]+=1.0
				if (num_target>0.0): data_times[individual][generation] = evaluation_time/num_target
				data_fractions[individual][generation] = fraction_of_time_passive

	for generation in range(n_generations):
		xdata[generation] = generation
		fraction_of_unfit_individuals[generation] /= n_individuals_gen[generation]
		fraction_of_BP_individuals[generation] /= n_individuals_gen[generation]
		fraction_of_advanced_individuals[generation] = number_of_advanced_individuals[generation]/n_individuals_gen[generation]
		avg_time_advanced[generation] /= number_of_advanced_individuals[generation]
		avg_fraction_time_BP_of_advanced_individuals[generation] /= number_of_advanced_individuals[generation]
		
	Results=[xdata,fraction_of_unfit_individuals,fraction_of_BP_individuals,fraction_of_advanced_individuals,number_of_advanced_individuals,advanced_individuals,avg_time_advanced,avg_fraction_time_BP_of_advanced_individuals,data_count_individuals_advanced,data_times,data_fractions]
	return Results

def main():
	evaluation_time = 500.0
	n_individuals_ini = 1000
	n_generations = 10
	dt = 0.0001
	factorx = 0.25 # =1 for 40 episodes
	
	Results = read_results(evaluation_time,n_individuals_ini,n_generations)
	
	avg_time_passive = 1.14166
	avg_time_casual = 1.91283
	avg_time_optimal = 0.28629
	# plots
	make_figure(factorx,dt,avg_time_passive,avg_time_casual,avg_time_optimal,Results)


main()
