#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Created on Thu Sep 28 14:47:37 2023

@author: victorcolas
"""

 

import numpy as np
from math import pi
import matplotlib.pyplot as plt


def CalculEnergie(x,y):
    npt=len(x)
    pas=x[1]-x[0]
    Energie=0
    for k in range(npt-1):
        Energie=Energie+y[k+1]*pas
    return Energie

def TF(x,y):
    dx=x[1]-x[0]
    N=len(y)
    TFy=np.fft.fftshift(np.fft.fft(y))
    freq=np.fft.fftshift(np.fft.fftfreq(N,dx))
    return freq,TFy

def iTF(TFy):
    y=np.fft.ifft(TFy)
    return y

def propagationFourier(freq,TFy,k,z):
    alpha=2*np.pi*freq
    argument =  (k** 2 - alpha ** 2)
    tmp = np.sqrt(np.abs(argument))
    gamma = np.where(argument >= 0, tmp, 1j*tmp)
    TFy_en_z=TFy*np.exp(-1j*gamma*z)
    return TFy_en_z





#### Description du faisceau ####
largeur = 30000.e-6           # Largeur de l'axe des abscisses
w0 = 20.e-6                  # Largeur du faisceau Gaussien
lambd = 633.e-9             # Longueur d'onde
Lr = (pi*w0**2) / lambd     # Longueur de Rayleigh
k = 2*pi/(lambd)           # Vecteur d'onde
N = 30001                   # Nombre de points de la fenêtre de calcul
i = complex (0,1)           # Nombre complexe
facteur = 0.01             # Pour définir dz en dessous
dz = facteur*Lr             # Distance de propagation


#### Description geométriques ####

NbRayleigh = 5     # Distance de propagation totale en nombre de longueur de Rayleigh
x=np.linspace(-largeur/2,largeur/2,N)
dx=abs(x[2]-x[1])
zFinal = NbRayleigh*Lr    # Distance de propagation mètres
Nb_cellules = int(NbRayleigh/facteur) # Nombre de cellules pour la BPM
wz=w0*np.sqrt(1+(zFinal/Lr)**2)   # formule largeur faisceau gaussien

# ## test TF ##
# nu1=10/(largeur)
# nu2=30/(largeur)
# y=np.cos(2*np.pi*nu1*x)+np.cos(2*np.pi*nu2*x)
# freq,TFy=TF(x,y)
# plt.figure()
# plt.plot(x,y, label = 'initial')
# plt.figure()
# plt.plot(freq,TFy, label = 'initial')
# plt.xlim([-3*nu2,3*nu2])


#### Ampltude gaussienne en z=0 ####
E_x_0 = np.exp((-2*x**2)/(w0**2))  

## optionnel : passage par une fente de largeur n*w0 en z=0
# porte=np.zeros(N,dtype = int)
# milieu=N//2
# n=0.3
# porte[int(milieu-(N*n*w0/largeur)//2):int(milieu+(N*n*w0/largeur)//2)] = 1 
# E_x_0=E_x_0*porte
  

#### Définition de la TF spatiale du champs E en z=0 ####
fx,E_alpha_0=TF(x,E_x_0)

#### Définition du propagateur puis multiplication = propagation fourier ####
E_alpha_z=propagationFourier(fx,E_alpha_0,k,zFinal)

#### Calcul de la TF-1 champs E en z ###
E_x_z=iTF(E_alpha_z)

#### Vérif de la conservation d'energie ###
EnergieInitiale=CalculEnergie(x,abs(E_x_0)**2)
EnergieFinale=CalculEnergie(x,abs(E_x_z)**2)
print(f'Energie initiale : {EnergieInitiale:.4e}')
print(f'Energie finale : {EnergieFinale:.4e}')

#### Affichage des intensité en z=0 pui z=zFinal ###
plt.figure()
plt.plot(x,abs(E_x_0)**2, label = 'initial')
plt.plot(x,abs(E_x_z)**2, label = 'final',ls='--')
plt.vlines(w0, 0, 1, color='gray' ,label = 'initial beam diameter th.',ls='--')
plt.vlines(-w0, 0, 1, color='gray' ,ls='--')
plt.vlines(wz, 0, 1, color='black' ,label = 'final beam diameter th.',ls='--')
plt.vlines(-wz, 0, 1, color='black' ,ls='--')
plt.grid(True) 
plt.legend()
plt.xlim([-5*wz ,5*wz])
plt.xlabel('x (m)')
plt.ylabel('Intensity distribution (A.U.)')

### Suivi de la propagation en 1+1D
print(f'Distance simulée : {zFinal:.4e} m')
print("Lancement de la propagation ")

Nb_cellules = int(NbRayleigh/facteur) # Nombre de cellules pour la BPM
z = np.linspace(0, zFinal, Nb_cellules)  # vecteur z de propagation mètres

TABLEAU =np.zeros([N,Nb_cellules],dtype= complex)  # tableau vide dont le nb de colonnes correspond aux nombres de cellules de calcul BPM
progression = 0    # debut de la propagation
TABLEAU [:,0] = abs(E_x_0)**2  #Premiere colonne du tableau est le faisceau de depart



for j in range(1,Nb_cellules):
    if (int(j)%(Nb_cellules/5)==0): # compteur pour l'affichage
        print("Calcul à ",int(j/Nb_cellules *100)," %")
    zj=z[j]
    E_alpha_zj=propagationFourier(fx,E_alpha_0,k,zj) 
    E_x_zj=iTF(E_alpha_zj)
    TABLEAU[:,j] = abs(E_x_zj)**2                  # On inclut le faisceau apres la derniere cellule dans notre tableau
    progression = progression + dz          # On a donc avancé de dz sur notre chemin
    j+=1

print("Calcul à  100 %")

fig, ax = plt.subplots()
plt.imshow(np.abs(TABLEAU),aspect='auto',cmap='jet',extent=[z.min(),z.max(),x.min(),x.max()])
plt.ylabel('x (m)')
plt.xlabel('Distance de propagation z (m)')
plt.colorbar()
plt.ylim([-2*wz ,2*wz])
plt.show()
