import numpy as np
import matplotlib.pyplot as plt
from matplotlib.widgets import Slider

# --- Paramètres physiques initiaux ---
lambda0 = 600e-9       # Longueur d'onde centrale (600 nm)
d = 0.2e-3             # Distance entre les deux trous d'Young (0.2 mm)
D2 = 1.0               # Distance Trous -> Écran (1.0 m)

# --- Grille de calcul ---
# Plage large (-30 mm à +30 mm) pour observer l'atténuation en s'éloignant du centre
x = np.linspace(-0.050, 0.050, 1200)
y = np.linspace(-0.005, 0.005, 200)

# --- Fonction de calcul d'intensité ---
def compute_intensity(dlambda_val, n_spectre=100):
    """
    Somme les contributions d'un spectre rectangulaire [lambda0 - dlambda/2, lambda0 + dlambda/2].
    """
    if dlambda_val == 0:
        wavelengths = [lambda0]
    else:
        wavelengths = np.linspace(lambda0 - dlambda_val / 2, lambda0 + dlambda_val / 2, n_spectre)

    I_total = np.zeros_like(x)

    for lmb in wavelengths:
        # Déphasage d phi = (2 * pi / lambda) * delta avec delta = d * x / D2
        phi = (2 * np.pi / lmb) * (d * x / D2)
        I_total += 1 + np.cos(phi)

    return I_total / len(wavelengths)

# --- 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(dlambda_val=0)
line, = ax_profile.plot(x * 1e3, I_init, color='blue', lw=1.5)
ax_profile.set_title(r"Profil d'intensité lumineuse $I(x)$ ($\lambda_0 = 600$ nm)")
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_dlambda = Slider(
    ax=ax_slider,
    label=r'Largeur spectrale $\Delta\lambda$ (nm) ',
    valmin=0.0,
    valmax=150.0,
    valinit=0.0,
    valstep=1.0,
    color='blue'
)

# --- Mise à jour lors du déplacement du curseur ---
def update(val):
    dlambda_m = slider_dlambda.val * 1e-9  # Conversion nm -> m
    I_new = compute_intensity(dlambda_m)

    # 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_dlambda.on_changed(update)

plt.show()