# -*- coding: utf-8 -*-

from __future__ import division
import matplotlib.pyplot as plt
import numpy as np
from math import pi
from RGB import *
from pylab import *
from matplotlib.widgets import Slider
import matplotlib.gridspec as gridspec

#blanche = True #lumière blanche ou spectre corps noir

class eclairement(object):
    def __init__(self, position):
        self.ax = plt.subplot(position)
        self.ax.set_title("Éclairement", fontsize = 12)
        self.ax.tick_params(axis='x', which='both', bottom='off', top='off', \
            labelbottom='off')
        
        plt.ylim((0, 1))
        self.l, = plt.plot(y, (ecl0(y, 0)[0] + ecl0(y, 0)[1] + ecl0(y, 0)[2]) /(
                            ecl0(0, 0)[0] + ecl0(0, 0)[1] + ecl0(0, 0)[2]), linewidth = 2)
        
        self.ax.set_ylim(-0.02, 1.02)
        self.ax.set_xlim(-champ, champ)

    def maj(self):
        self.l.set_xdata(y)
        self.l.set_ydata((ecl0(y, 0)[0] + ecl0(y, 0)[1] + ecl0(y, 0)[2]) / 
            (ecl0(0, 0)[0] + ecl0(0, 0)[1] + ecl0(0, 0)[2]))
        self.ax.set_xlim(-champ, champ)
        draw()

class figure_interf(object):
    def __init__(self, y, z, position, ax=None, im=None, arrow=None, xpos=None, ypos=None):
        #création du graphe donnant les franges d'interférence
        Y,Z = np.meshgrid(y, z) #donne 2 tableaux contenant les coordonnées de chaque
                         #point en fonction de son index (deux dimensions)
        image = np.transpose(ecl(Y,Z), (1, 2, 0))
        self.xpos = delta*1e-9*OE/a/champ*resol+resol
        self.ypos = 75
        self.ax = plt.subplot(position)
        self.ax.get_xaxis().set_visible(False)
        self.ax.get_yaxis().set_visible(False)
        self.im = plt.imshow(image,extent=[0, 2*resol*coeff, 0, 200])
        legende = r"$\delta=%i \,\mathrm{nm}$" %delta
        self.arrow = self.ax.annotate(legende, xy=(self.xpos, self.ypos),
                xytext=(-80,30), textcoords='offset points',fontsize=15,
                bbox=dict(boxstyle="round4", fc="w"),
                arrowprops=dict(arrowstyle="->"))
        fig.canvas.mpl_connect('button_press_event',self.click) #connexion avec gestion click souris
    
    
    
    def maj_arrow(self):
        if self.ypos < 75:
            ytext = 30
        else:
            ytext = -30
        if self.xpos > 100:
            xtext = -80
        else:
            xtext = +50
        if abs(delta) < 1000:
            legende = r"$\delta=%i \,\mathrm{nm}$" %delta
        else:
            legende = r"$\delta=%i,%i \,\mathrm{\mu m}$" %(int(delta/1e3),abs(int((delta-int(delta/1e3)*1e3)/10)))
        self.arrow = self.ax.annotate(legende, xy=(self.xpos, self.ypos),  
                xytext=(xtext,ytext), textcoords='offset points',fontsize=15,
                bbox=dict(boxstyle="round4", fc="w"),
                arrowprops=dict(arrowstyle="->"))
        draw()
        
    def maj(self):
        Y,Z = np.meshgrid(y, z) #donne 2 tableaux contenant les coordonnées de chaque
                         #point en fonction de son index (deux dimensions)
        image = np.transpose(ecl(Y,Z), (1, 2, 0))
        self.im.set_data(image)
        self.xpos = delta*1e-9*OE/a/champ*resol+resol
        self.ax.texts.remove(self.arrow)
        self.maj_arrow()
        
        
        draw()
   
     
    def click(self, event):
            global delta
            if event.inaxes==self.ax:
                #Get nearest data
                self.xpos=event.xdata
                self.ypos=event.ydata
                                
                #Check which mouse button:
                if event.button==1:
                    self.ax.texts.remove(self.arrow)
                    delta = (self.xpos-resol)*a/OE*champ/resol*1e9
                    
                    self.maj_arrow()
                    g4.maj()
                    g5.maj()
                    g6.maj()
                    
class couleur(object):
    def __init__(self, pos):
        self.ax = plt.subplot(pos)
        self.ax.set_title("couleur du point visé", fontsize = 8)
        posx = OE*delta/a*1e-9
        coul = ecl(posx, 0)
        image = np.array([[coul], [coul]])
         
        self.ax.get_xaxis().set_visible(False)
        self.ax.get_yaxis().set_visible(False)
        self.im = plt.imshow(image,extent=[0, 10, 0, 10])
    
    def maj(self):
        posx = OE*delta/a*1e-9
        coul = ecl(posx, 0)
        image = np.array([[coul], [coul]])
         
        self.im.set_data(image)
        draw()   
               

class spectre(object):
    def __init__(self, pos, source):
        update_source()
        self.source = source
        self.ax = plt.subplot(pos)
        self.ax.get_xaxis().set_visible(False)
        self.ax.get_yaxis().set_visible(False)
        self.ax.set_xlabel('$\lambda (\mathrm{nm})$')
        if self.source:
            titre = r"Spectre de la source"
        else:
            titre = r"Spectre pour $\delta = %i\,\mathrm{nm}$" %delta
        image = np.ndarray(shape=(1,381,3))
        for i in visible:
            image[0, i-400, :] = trace_spectre(i, delta, source)
        self.ax.set_title(titre, fontsize = 12)
        self.im = plt.imshow(image, extent=[400, 781, 0, 70])
    
    def maj(self):
        image = np.ndarray(shape=(1,381,3))
        for i in visible:
            image[0, i-400, :] = trace_spectre(i, delta, self.source)
        self.im.set_data(image)
        if not self.source:
            if abs(delta) < 1000:
                legende = r"Spectre pour $\delta=%i \,\mathrm{nm}$" %delta
            else:
                legende = r"Spectre pour $\delta=%i,%i \,\mathrm{\mu m}$" %(int(delta/1e3),
                        abs(int((delta-int(delta/1e3)*1e3)/10)))
            self.ax.set_title(legende, fontsize = 12)
      
            
        draw()


class courbe_spectre(object):
    def __init__(self, pos):
        update_source()
        self.ax = plt.subplot(pos)
        if abs(delta) < 1000:
            legende = r"Spectre pour $\delta=%i \,\mathrm{nm}$" %delta
        else:
            legende = r"Spectre pour $\delta=%i,%i \,\mathrm{\mu m}$" %(int(delta/1e3),
                    abs(int((delta-int(delta/1e3)*1e3)/10)))
        self.ax.set_title(legende, fontsize = 12)
        self.ax.set_xlim(400, 780)
        self.ax.set_ylim(-.02, 1.02)
        self.l, = plt.plot(visible, E(visible, delta) * spectre_source[visible - 400], linewidth = 2)
        
    def maj(self):
        self.l.set_ydata(E(visible, delta) * spectre_source[visible - 400])
        if abs(delta) < 1000:
            legende = r"Spectre pour $\delta=%i \,\mathrm{nm}$" %delta
        else:
            legende = r"Spectre pour $\delta=%i,%i \,\mathrm{\mu m}$" %(int(delta/1e3),
                    abs(int((delta-int(delta/1e3)*1e3)/10)))
        self.ax.set_title(legende, fontsize = 12)
        draw()






#fonction tenant compte de la réponse non linéaire de l'oeil et variant entre 0 et 1
def ressenti(x):
    return np.sqrt(x)


def update_source():
    global spectre_source
    spectre_source = np.zeros(381)
    for i in visible:
        if i<lambda_max and i>lambda_min:
            spectre_source[i-400] = 1

def E(longueur, delta): #formule de Fresnel
    k1 = 2.0*pi/(longueur*1e-9)
    return((1.0+np.cos(k1*delta*1e-9))/2.0)


#calcule chaque composante de l'éclairement d'un point (y,z) de l'écran
def ecl0(y, z): 
    delta = a*y/OE*1e9
    
    if lambda_min == lambda_max:
        Ec = E(lambda_min, delta)
        WL = WavelengthToRGB(lambda_min)
        R = Ec*WL[0]
        G = Ec*WL[1]
        B = Ec*WL[2]
    else:
        R, G, B = 0, 0, 0
        for i in xx:
            Ec = E(i, delta)
            WL = WavelengthToRGB(i)
            R+=Ec*WL[0]
            G+=Ec*WL[1]
            B+=Ec*WL[2]
    return(np.array([(R/RGBtot[0]), (G/RGBtot[1]), (B/RGBtot[2])]))

def trace_spectre(y, delta, source):
    if source:
        Ec = spectre_source[y - 400]
    else:
        Ec = E(y, delta) * spectre_source[y - 400]
    WL = WavelengthToRGB(y)
    R=(Ec*WL[0])
    G=(Ec*WL[1])
    B=(Ec*WL[2])
    return(np.array([R, G, B]))        

def ecl(y, z):
    return ressenti(ecl0(y, z))


def update_fig(val): #changement figure mais pas spectres
    global RGBtot, champ, y
    champ = slider_champ.val * 1e-3
    coeff = champ/champ0 #permet de garder une précision suffisante si on augmente le champ
    y = np.linspace(-champ, champ, int(coeff*resol))
    
    
    g1.maj()
    g2.maj()


def update_spectres(val): #changement de toutes les figures
    global xx, RGBtot, lambda_min, lambda_max, champ, y
    lambda_min = slider_moy.val - int(slider_ecart.val/2)
    lambda_max = slider_moy.val + int(slider_ecart.val/2)
    y = np.linspace(-champ, champ, resol)
    #yy = np.linspace(-champ, champ, coeff*resol)
    xx = np.arange(lambda_min, lambda_max, pas_lambda)
    RGBtot = np.ones(3)
    update_source()
    for i in xx:
        RGBtot+= WavelengthToRGB(i)
    champ = slider_champ.val * 1e-3
    g3.maj()
    g4.maj()
    g5.maj()
    g1.maj()
    g2.maj()
    g6.maj()

#Initialisation de toutes les variables
    
lambda_moy0 = 600
ecart0 = 400
lambda_min = lambda_moy0 - int(ecart0/2) #lambda mini spectre en nm
lambda_max = lambda_moy0 + int(ecart0/2) #lambda maxi spectre en nm
long0 = (lambda_min + lambda_max)/2 # longueur d'onde moyenne du spectre en nm
pas_lambda = 5 #incrément en lambda pour le calcul des interférences en nm
resol = 1000 #résolution de l'image calculée
RGBtot = np.ones(3) #initialisation composantes totales RGB sur tout le spectre

xx = np.arange(lambda_min, lambda_max, pas_lambda) #étendue du spectre visible
visible = np.arange(400, 781)    
    
for i in xx:
    RGBtot+= WavelengthToRGB(i) #calcul composantes totales RGB en lumière blanche

### Fentes d'Young
a = 1e-3              # distance entre les sources en m
OE = 1        # distance entre l'origine O et le centre E de l'écran en m
NbF = 20        # nombre de franges initiales sur lambda_0
champ0 = NbF * long0 * 1e-9 * OE / (2*a) # étendue de la zone de tracé des interférences en m
champ = champ0 #champ initial d'observation
coeff = champ/champ0
delta = 500 #valeur initiale de delta

y = np.linspace(-champ, champ, int(coeff*resol)) #tableau des coordonnées réelles des points

z = np.linspace(0, 1, 1) #ecl indépendant de z donc une seule valeur de z suffit

fig = plt.figure(1, figsize = (12, 10)) #création de la figure principale à afficher
fig.suptitle("Trous d'Young en lumière non monochromatique", fontsize = 20)

gs = gridspec.GridSpec(4, 6, height_ratios=[2, 2, 1.5, 1.5]) #définition de l'arrangement des subplots
#dans la figure principale
g1 = figure_interf(y, z, gs[0, :]) #figure d'interférence pleine largeur en haut
g2 = eclairement(gs[1, :]) #graphe de l'éclairement pleine largeur en 2ième position
g3 = spectre(gs[3,3:], True) #spectre source
g4 = spectre(gs[2,3:], False) #spectre point visé sur la figure d'interférence
g5 = courbe_spectre(gs[2,0:3]) #courbe spectre point visé sur la figure d'interférence
g6 = couleur(gs[3,2]) #couleur point visé

gs.tight_layout(fig) #définit espace entre subplots par rapport à fenêtre principale

#définition des axes des 3 sliders et de leur couleur
axcolor = 'lightgoldenrodyellow'
axmoy = axes([0.03, 0.15, 0.3, 0.015], facecolor=axcolor)
axdelta = axes([0.03, 0.1, 0.3, 0.015], facecolor=axcolor)
axchamp = axes([0.03, 0.05, 0.3, 0.015], facecolor=axcolor)

#Création des 3 sliders
slider_moy = Slider(axmoy, r'$\lambda_0 : $', 400, 780, valinit = lambda_moy0, valfmt = '%i $nm$')
slider_ecart = Slider(axdelta, r'$\delta\lambda : $', 0, 400, valinit = ecart0, valfmt = '%i $nm$')
slider_champ = Slider(axchamp, r'$\mathrm{Champ} : $', champ0 * 1e3, champ0*4e3, valinit = champ0*1e3,
                      valfmt = '%i $mm$')

#Mise à jour des graphes si sliders modifiés
slider_moy.on_changed(update_spectres)
slider_ecart.on_changed(update_spectres)
slider_champ.on_changed(update_fig)

plt.show() #affichage de la fenêtre principale
