import numpy as np
import matplotlib.pyplot as plt
from matplotlib.widgets import Slider

# --- Paramètres physiques initiaux ---
wavelength = 632.8e-9  # Longueur d'onde (laser hélium-néon) en mètres
d = 0.5e-3             # Distance entre les deux trous d'Young (0.5 mm)
D1 = 0.5               # Distance Source -> Trous (0.5 m)
D2 = 1.0               # Distance Trous -> Écran (1.0 m)

# --- Grille de calcul ---
x = np.linspace(-0.005, 0.005, 1000)  # Position sur l'écran (de -5 mm à +5 mm)
y = np.linspace(-0.005, 0.005, 200)   # Hauteur sur l'écran pour le rendu 2D

# --- Fonction de calcul d'intensité ---
def compute_intensity(a_val, n_points=100):
    """
    Calcule l'intensité moyenne sur l'écran en sommant les contributions
    incohérentes des points d'une source étendue de largeur a_val.
    """
    if a_val == 0:
        source_points = [0.0]
    else:
        source_points = np.linspace(-a_val / 2, a_val / 2, n_points)

    I_total = np.zeros_like(x)

    for xs in source_points:
        # Déphasage introduit par la position xs du point source
        phi = (2 * np.pi / wavelength) * d * (x / D2 + xs / D1)
        # Figure d'interférence pour un point source (I = 2 * I0 * (1 + cos(phi)))
        I_total += 1 + np.cos(phi)

    return I_total / len(source_points)

# --- Création de la figure ---
fig, (ax_profile, ax_image) = plt.subplots(2, 1, figsize=(9, 7), gridspec_kw={'height_ratios': [2, 1]})
plt.subplots_adjust(bottom=0.25, hspace=0.4)

# 1. Plot du profil 1D d'intensité
I_init = compute_intensity(a_val=0)
line, = ax_profile.plot(x * 1e3, I_init, color='red', lw=1.5)
ax_profile.set_title("Profil d'intensité lumineuse $I(x)$")
ax_profile.set_xlabel("Position sur l'écran $x$ (mm)")
ax_profile.set_ylabel("Intensité relative")
ax_profile.set_ylim(0, 2.2)
ax_profile.grid(True, alpha=0.3)

# 2. Rendu 2D de la figure d'interférence
I_2d_init = np.tile(I_init, (len(y), 1))
img = ax_image.imshow(
    I_2d_init,
    extent=[x[0]*1e3, x[-1]*1e3, y[0]*1e3, y[-1]*1e3],
    cmap='gray',
    aspect='auto',
    vmin=0,
    vmax=2
)
ax_image.set_title("Figure d'interférence visualisée sur l'écran")
ax_image.set_xlabel("Position $x$ (mm)")
ax_image.set_yticks([])

# --- Ajout du curseur (Slider) ---
ax_slider = plt.axes([0.2, 0.08, 0.6, 0.03])
slider_a = Slider(
    ax=ax_slider,
    label='Largeur source $a$ (mm) ',
    valmin=0.0,
    valmax=1.5,
    valinit=0.0,
    valstep=0.01,
    color='red'
)

# --- Mise à jour lors du déplacement du curseur ---
def update(val):
    a_m_high = slider_a.val * 1e-3  # Conversion mm -> m
    I_new = compute_intensity(a_m_high)

    # Mise à jour du profil 1D
    line.set_ydata(I_new)

    # Mise à jour de la figure 2D
    I_2d_new = np.tile(I_new, (len(y), 1))
    img.set_data(I_2d_new)

    fig.canvas.draw_idle()

slider_a.on_changed(update)

plt.show()