import numpy as np
import math
import cmath
from scipy.linalg import eig
import scipy.special as sc
import matplotlib.pyplot as plt
import time
import sys

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 matplotlib.colors import BoundaryNorm
from matplotlib.ticker import MaxNLocator

from matplotlib.patches import Circle, Ellipse, Wedge, Polygon
from matplotlib.collections import PatchCollection


np.set_printoptions(precision=4)

def Jb(l,x):
    return sc.jv(l,x)

def jb(l,n):
    return sc.jn_zeros(l,n)[n-1]

def ev(n,l,j):
    return -(jb(l,n)**2+G*(j-l)**2)

def cp(n1,n,l):
    return jb(l,n)/(Jb(l+1,jb(l,n)))*jb(l+1,n1)*Jb(l+1,jb(l,n))*Jb(l,jb(l+1,n1))/(Jb(l+2,jb(l+1,n1))*(jb(l,n)**2-jb(l+1,n1)**2))

def cm(n1,n,l):
    return jb(l,n)/(Jb(l+1,jb(l,n)))*jb(l-1,n1)*Jb(l-1,jb(l,n))*Jb(l,jb(l-1,n1))/(Jb(l,jb(l-1,n1))*(jb(l,n)**2-jb(l-1,n1)**2))

def nn(a):
    return a%Nn+1

def ll(a):
    return int(a/Nn)-Nl

def aa(n,l):
    return (l+Nl)*Nn+(n-1)

def mat(a,b,j,Pe):
    if a==b:
        return ev(nn(a),ll(a),j)
    if ll(b)==(ll(a)+1):
        return Pe*cp(nn(b),nn(a),ll(a))
    if ll(b)==(ll(a)-1):
        return Pe*cm(nn(b),nn(a),ll(a))
    else:
        return 0.00000000
            
def make_figure (v_Pe,v_eigenvalues,v_eigenvalues_imag_ordered,nrows):
	# setto alcune variabili comuni
	axisticslabelfontsize=8
	axisticslabelfontsizeinset=7
	axislabelfontsize=11 
	axislabelfontsizeinset=9
	

	xsize = 3.5
	ysize = 4
	
	with PdfPages('fig1.pdf') as pdf:
		fig = plt.figure(figsize=(xsize,ysize))
		plt.rc('text', usetex=True)
		plt.rc('text.latex', preamble = ','.join('''\usepackage{txfonts} \usepackage{lmodern}'''.split()))
		
		######################### panel 1 ##############################
		panel = fig.add_axes([0.18,0.54,0.76,0.44])
		
		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_ylabel(r'Re$(\lambda_{n,\ell,j}^{\mbox{\scriptsize{Pe}}})$',fontsize=axislabelfontsize)
		panel.set_xticklabels([])
		
		panel.set_xlim(-1,21)
		panel.xaxis.set_major_locator(MultipleLocator(5))
		panel.xaxis.set_minor_locator(MultipleLocator(1))
		
		panel.set_ylim(5,177)
		panel.yaxis.set_major_locator(MultipleLocator(20))
		panel.yaxis.set_minor_locator(MultipleLocator(10))
		
		
		for i in range(nrows):
			y = - np.real(v_eigenvalues[i])
			panel.plot(v_Pe,y,'b-',linewidth=1.0,alpha=0.5)
			
		#~ xs=[]
		#~ ys=[]
		#~ list_e=[]
		#~ for i in range(nrows-1):
			#~ for j in range(i+1,nrows):
				#~ for k in range(len(v_Pe)):
					#~ if (abs(np.real(v_eigenvalues[i][k])-np.real(v_eigenvalues[j][k]))<0.01 and abs(np.imag(v_eigenvalues[i][k])-np.imag(v_eigenvalues[j][k]))<1.0):
						#~ le=1
						#~ for kk in range(len(list_e)):
							#~ if (list_e[kk]==i): le=0
						#~ if (le==1):
							#~ xs.append(v_Pe[k])
							#~ ys.append(np.real(-v_eigenvalues[i][k]))
							#~ list_e.append(i)
							#~ print xs[-1],ys[-1]
						#~ xs.append(v_Pe[k])
						#~ ys.append(np.real(-v_eigenvalues[i][k]))
						#~ print xs[-1],ys[-1]
		#~ print xs
		#~ print ys
		
		xs=[2.7655310621242486, 4.288577154308617, 1.0821643286573146, 1.723446893787575, 3.406813627254509,12.585170340681362]
		ys=[112.63277782766068, 90.29431222303393, 64.01475503668654, 33.18115277183727, 14.286401529867401,52.89084988322897]
		
		
		plt.scatter(xs, ys, s=30, facecolors='None', edgecolors='r')
		
		
		
		
		
		
		
		
		######################### panel 1 ##############################
		panel = fig.add_axes([0.18,0.1,0.76,0.44])
		
		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'Pe',fontsize=axislabelfontsize)
		panel.set_ylabel(r'Im$(\lambda_{n,\ell,j}^{\mbox{\scriptsize{Pe}}})$',fontsize=axislabelfontsize)
		
		panel.set_xlim(-1,21)
		panel.xaxis.set_major_locator(MultipleLocator(5))
		panel.xaxis.set_minor_locator(MultipleLocator(1))
		
		panel.set_ylim(-50,50)
		panel.yaxis.set_major_locator(MultipleLocator(20))
		panel.yaxis.set_minor_locator(MultipleLocator(10))
		
		for i in range(nrows):
			y = - np.asarray(v_eigenvalues_imag_ordered[i])
			panel.plot(v_Pe,y,'b-',linewidth=1.0,alpha=0.5)
			
		#~ panel.text(0.3,1.08,r'$t=0.2$',fontsize=10,transform=panel.transAxes)
		
		
		

		pdf.savefig(fig)
	return


def main():        
	# parameters
	global G
	G=4
	
	# sizes
	global Nn
	global Nl
	global Ntot
	Nn=3
	Nl=2
	Ntot=Nn*(2*Nl+1)
	
	# channel
	j_channel = 1
	
	for n in range(1,Nn+1):
		for l in range(-Nl,Nl+1):
			print n,l,j_channel,-ev(n,l,j_channel)
		
	
	
	NPe=500
	Pemax=20.0

	print('Size of matrix: ',Ntot)

	M = [[0 for i in range(Ntot)] for j in range(Ntot)]

	Pe=np.linspace(0,Pemax,num=NPe)
	Spec=[[0 for i in range(NPe)] for j in range(Ntot)]
	ImSpec=[[0 for i in range(NPe)] for j in range(Ntot)]

	for j in range(NPe):
		for a in range(Ntot):
			for b in range(Ntot):
				M[b][a]=mat(a,b,j_channel,Pe[j])
		values, vectors = eig(M)
		values=np.sort_complex(values)
		for i in range(Ntot):
			Spec[i][j]=values[i]
		Imvalues=np.sort(np.imag(values))
		for i in range(Ntot):
			ImSpec[i][j]=Imvalues[i]

	
	make_figure (Pe,Spec,ImSpec,Ntot)
	
	
main()
	
	
	
	
