# -*- 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



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')
                
        self.l, = plt.plot(y,E(a*y/OE*1e9) , linewidth = 2, color = 'blue')
        self.l2, = plt.plot(y,env(a*y/OE*1e9),linewidth = 2, color = 'red')
        self.l3, = plt.plot(y,1-env(a*y/OE*1e9),linewidth = 2, color = 'red')
        self.ax.set_ylim(-0.02, 1.02)
        

    def maj(self):
        self.l.set_xdata(y)
        self.l.set_ydata(E(a*y/OE*1e9))
        self.l2.set_xdata(y)
        self.l2.set_ydata(env(a*y/OE*1e9))
        self.l3.set_xdata(y)
        self.l3.set_ydata(1-env(a*y/OE*1e9))
        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.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, 0, int(resol/3)])

         
    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)
                  
        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 E(delta): #formule de Fresnel
    return 0.5*(1+cos(2*pi*delta/long0)*cos(pi*ecart0/long0**2*delta))

def env(delta): #enveloppe
    return 0.5*(1+cos(pi*ecart0/long0**2*delta))

#calcule chaque composante de l'éclairement d'un point (y,z) de l'écran
def ecl(y,z):
    delta = a*y/OE*1e9
    Ec = E(delta)
    R, G, B = Ec*WL[0],Ec*WL[1],Ec*WL[2]
    return np.array([ressenti(R),ressenti(G),ressenti(B)])




#Initialisation de toutes les variables
    
long0 = 600
ecart0 = 1.
k0 = 2.0*pi/(long0*1e-9)
k1 = 2.0*pi/((long0 - ecart0 / 2)*1e-9)
k2 = 2.0*pi/((long0 + ecart0 / 2)*1e-9)
resol = 2000 #résolution de l'image calculée
WL = WavelengthToRGB(long0)
### 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
y = np.linspace(-champ, champ, 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

def update_fig(val): #changement figure mais pas spectres
    global champ, y, WL, long0, ecart0, k0, k1, k2
    WL = WavelengthToRGB(long0)
    champ = slider_champ.val * 1e-3
    long0 = slider_moy.val
    ecart0 = slider_ecart.val
    k0 = 2.0*pi/(long0*1e-9)
    k1 = 2.0*pi/((long0 - ecart0 / 2)*1e-9)
    k2 = 2.0*pi/((long0 + ecart0 / 2)*1e-9)
    y = np.linspace(-champ, champ, resol)
    g1.maj()
    g2.maj()

fig = plt.figure(1, figsize = (12, 10)) #création de la figure principale à afficher
fig.suptitle("Trous d'Young éclairés par un doublet", fontsize = 20)

gs = gridspec.GridSpec(3, 1, height_ratios=[1.5 , 2, 1]) #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
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.07, 0.15, 0.8, 0.015], facecolor=axcolor)
axdelta = axes([0.07, 0.1, 0.8, 0.015], facecolor=axcolor)
axchamp = axes([0.07, 0.05, 0.8, 0.015], facecolor=axcolor)


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

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

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