"""
How a Photon Fires a Neuron — Python simulations
Companion code for the courseshub.world lesson on the 2026 Nobel Prize
in Physiology or Medicine (optogenetics).

Three simulations, one per part of the lesson:
  1. Particle in a box: retinal's pi-electron levels and absorption wavelength
  2. Landau–Zener: numerical two-state dynamics through an avoided crossing
  3. Membrane + ChR2: a two-state channel driving a leaky integrate-and-fire neuron

Run:  python optogenetics_sims.py      (writes four PNG figures)
Needs: numpy, scipy, matplotlib
"""

import numpy as np
import matplotlib.pyplot as plt
from scipy.integrate import solve_ivp

# Physical constants (SI)
h = 6.62607015e-34       # Planck constant, J s
hbar = h / (2 * np.pi)
m_e = 9.1093837e-31      # electron mass, kg
c = 2.99792458e8         # speed of light, m/s
eV = 1.602176634e-19     # J per eV

plt.rcParams.update({"figure.dpi": 150, "font.size": 10,
                     "axes.spines.top": False, "axes.spines.right": False})


# ---------------------------------------------------------------------------
# 1. Particle in a box
# ---------------------------------------------------------------------------
def pib_energy(n, L):
    """Energy of level n (J) in a 1D infinite well of length L (m)."""
    return n**2 * h**2 / (8 * m_e * L**2)


def pib_wavelength(N, L):
    """HOMO->LUMO absorption wavelength (m) for N pi electrons in a box of length L."""
    dE = pib_energy(N // 2 + 1, L) - pib_energy(N // 2, L)
    return h * c / dE


def sim_particle_in_box():
    bond = 1.40e-10                        # average C-C bond length, m
    # Retinal Schiff base: 12 pi electrons, 11 bonds from C5 to N
    N, L = 12, 11 * bond
    lam = pib_wavelength(N, L)
    print(f"[PIB] retinal: L = {L*1e10:.1f} Å, ΔE = {h*c/lam/eV:.2f} eV, λ = {lam*1e9:.0f} nm")

    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10, 4.2))

    # Left: wavefunctions offset by their energies, filled vs empty
    x = np.linspace(0, L, 400)
    for n in range(1, 9):
        E = pib_energy(n, L) / eV
        psi = np.sqrt(2 / L) * np.sin(n * np.pi * x / L)
        scale = 0.45 / np.sqrt(2 / L)      # cosmetic amplitude in eV units
        filled = n <= N // 2
        ax1.axhline(E, color="0.85", lw=0.8, zorder=0)
        ax1.plot(x * 1e10, E + scale * psi, color="C0" if filled else "C1",
                 lw=1.6 if filled else 1.2, ls="-" if filled else "--")
        ax1.text(L * 1e10 + 0.4, E, f"n={n}", va="center", fontsize=8)
    ax1.set_xlabel("Position along the chain (Å)")
    ax1.set_ylabel("Energy (eV)")
    ax1.set_title("Wavefunctions: filled (solid) and empty (dashed)")

    # Right: wavelength vs number of conjugated double bonds
    k = np.arange(2, 13)                   # number of C=C (or C=N) double bonds
    N_k = 2 * k
    L_k = (2 * k - 1) * bond               # bonds between end atoms
    lam_k = np.array([pib_wavelength(n, l) for n, l in zip(N_k, L_k)]) * 1e9
    ax2.plot(k, lam_k, "o-", color="C0")
    ax2.axhspan(380, 750, color="0.94", zorder=0)
    ax2.text(2.2, 730, "visible band", fontsize=8, color="0.4", va="top")
    ax2.plot([6], [lam * 1e9], "s", color="C3", ms=9, label="retinal Schiff base (model)")
    ax2.axhline(470, color="C2", ls=":", lw=1.2)
    ax2.text(12, 478, "ChR2 measured ≈ 470 nm", ha="right", fontsize=8, color="C2")
    ax2.set_xlabel("Number of conjugated double bonds")
    ax2.set_ylabel("Predicted λ (nm)")
    ax2.set_title("Longer conjugation, redder absorption")
    ax2.legend(frameon=False, fontsize=8, loc="lower right")

    fig.tight_layout()
    fig.savefig("sim1_particle_in_box.png")
    plt.close(fig)


# ---------------------------------------------------------------------------
# 2. Landau–Zener: two-state dynamics through a crossing
# ---------------------------------------------------------------------------
def lz_numeric(V, alpha, T=60.0):
    """
    Solve i hbar dc/dt = H(t) c with H = [[alpha t/2, V], [V, -alpha t/2]]
    in reduced units (hbar = 1). Start in diabatic state 1 at t = -T.
    Returns times and |c1|^2 (probability of staying diabatic = hopping adiabatically).
    """
    def rhs(t, y):
        c1, c2 = y[0] + 1j * y[1], y[2] + 1j * y[3]
        d1 = -1j * (0.5 * alpha * t * c1 + V * c2)
        d2 = -1j * (V * c1 - 0.5 * alpha * t * c2)
        return [d1.real, d1.imag, d2.real, d2.imag]

    sol = solve_ivp(rhs, (-T, T), [1, 0, 0, 0], max_step=0.02, rtol=1e-8, atol=1e-10)
    p1 = sol.y[0] ** 2 + sol.y[1] ** 2
    return sol.t, p1


def lz_formula(V, alpha):
    """Landau–Zener hop probability, hbar = 1: P = exp(-2 pi V^2 / alpha)."""
    return np.exp(-2 * np.pi * V**2 / alpha)


def sim_landau_zener():
    alpha = 1.0                                   # sweep rate |d(E1-E2)/dt|
    Vs = np.linspace(0.0, 0.6, 25)
    P_num = [lz_numeric(V, alpha)[1][-1] for V in Vs]
    P_th = lz_formula(Vs, alpha)
    print(f"[LZ] max |numeric - formula| = {np.max(np.abs(np.array(P_num) - P_th)):.3f}")

    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10, 4.2))

    for V, col in [(0.05, "C3"), (0.2, "C1"), (0.5, "C0")]:
        t, p1 = lz_numeric(V, alpha)
        ax1.plot(t, p1, color=col, lw=1.4, label=f"V = {V}")
    ax1.axvline(0, color="0.7", lw=0.8, ls=":")
    ax1.set_xlabel("Time (reduced units; crossing at t = 0)")
    ax1.set_ylabel("Probability of staying on the diabatic path")
    ax1.set_title("Small coupling: the system goes straight through")
    ax1.legend(frameon=False, fontsize=8)

    ax2.plot(Vs, P_th, color="C0", lw=2, label="Landau–Zener formula")
    ax2.plot(Vs, P_num, "o", color="C3", ms=4, label="numerical solution")
    ax2.set_xlabel("Coupling V at the crossing")
    ax2.set_ylabel("Hop probability P")
    ax2.set_title("V → 0 (conical intersection): P → 1")
    ax2.legend(frameon=False, fontsize=8)

    fig.tight_layout()
    fig.savefig("sim2_landau_zener.png")
    plt.close(fig)


# ---------------------------------------------------------------------------
# 3. ChR2 + leaky integrate-and-fire neuron
# ---------------------------------------------------------------------------
# Membrane (values from Part 3 of the lesson)
C_m = 100e-12          # F
R_m = 150e6            # Ohm
V_rest = -70e-3        # V
V_th = -50e-3          # V, threshold 20 mV above rest
V_reset = -70e-3       # V
t_ref = 2e-3           # s, refractory period
E_rev = 0.0            # V, ChR2 reversal potential

# Channelrhodopsin-2, simple two-state model: closed <-> open
gamma = 50e-15         # S, single-channel conductance
N_ch = 3.0e5           # channels: ≈ 15 nS, ≈ 1 nA max photocurrent at rest
sigma = 2e-8           # µm^2, absorption cross-section
phi_qy = 0.5           # quantum yield of opening
k_close = 100.0        # 1/s, closing rate in the dark (≈ 10 ms)


def photon_flux(I_mW_mm2, lam=470e-9):
    """Photons s^-1 µm^-2 for an irradiance in mW/mm^2."""
    E_ph = h * c / lam
    return (I_mW_mm2 * 1e-3 / 1e6) / E_ph       # W/µm^2 divided by J/photon


def simulate_neuron(I_mW_mm2, pulses, T=0.25, dt=1e-5):
    """
    Forward-Euler integration of
        dp/dt = k_open(t) (1 - p) - k_close p
        C dV/dt = -(V - V_rest)/R - g_max p (V - E_rev)
    pulses: list of (t_on, t_off) in seconds when the light is on.
    Returns t, V, p, spike times.
    """
    n = int(T / dt)
    t = np.arange(n) * dt
    light = np.zeros(n, bool)
    for on, off in pulses:
        light |= (t >= on) & (t < off)
    k_open = sigma * photon_flux(I_mW_mm2) * phi_qy
    g_max = N_ch * gamma

    V = np.empty(n); p = np.empty(n)
    V[0], p[0] = V_rest, 0.0
    spikes, ref_until = [], -1.0
    for i in range(n - 1):
        ko = k_open if light[i] else 0.0
        p[i + 1] = p[i] + dt * (ko * (1 - p[i]) - k_close * p[i])
        if t[i] < ref_until:
            V[i + 1] = V_reset
            continue
        I_chr = g_max * p[i] * (E_rev - V[i])            # inward current, A
        dV = (-(V[i] - V_rest) / R_m + I_chr) / C_m
        V[i + 1] = V[i] + dt * dV
        if V[i + 1] >= V_th:
            spikes.append(t[i + 1])
            V[i] = 20e-3                                  # draw the spike peak
            V[i + 1] = V_reset
            ref_until = t[i + 1] + t_ref
    return t, V, p, np.array(spikes)


def sim_neuron():
    pulses = [(0.02 + 0.05 * k, 0.025 + 0.05 * k) for k in range(4)]   # 5 ms at 20 Hz
    fig, axes = plt.subplots(3, 1, figsize=(9, 6.5), sharex=True,
                             gridspec_kw={"height_ratios": [0.5, 1, 2]})
    for I, col in [(0.5, "C1"), (2.0, "C2"), (10.0, "C0")]:
        t, V, p, sp = simulate_neuron(I, pulses)
        axes[1].plot(t * 1e3, p, color=col, lw=1.3, label=f"{I} mW/mm²")
        axes[2].plot(t * 1e3, V * 1e3, color=col, lw=1.1)
        print(f"[LIF] {I:>4} mW/mm²: k_open = {sigma*photon_flux(I)*phi_qy:6.1f}/s, "
              f"peak p_open = {p.max():.2f}, spikes = {len(sp)}")
    for on, off in pulses:
        axes[0].axvspan(on * 1e3, off * 1e3, color="#3b82f6", alpha=0.6)
    axes[0].set_yticks([]); axes[0].set_ylabel("light")
    axes[1].set_ylabel("open fraction p")
    axes[1].legend(frameon=False, fontsize=8, ncol=3)
    axes[2].axhline(V_th * 1e3, color="0.6", ls=":", lw=1)
    axes[2].text(248, V_th * 1e3 + 1, "threshold", ha="right", fontsize=8, color="0.4")
    axes[2].set_ylabel("V (mV)")
    axes[2].set_xlabel("Time (ms)")
    axes[0].set_title("5 ms blue-light pulses at 20 Hz: only bright light fires reliably")
    fig.tight_layout()
    fig.savefig("sim3_neuron_pulses.png")
    plt.close(fig)


def sim_latency():
    """First-spike latency under continuous light, vs irradiance."""
    Is = np.logspace(-0.3, 1.5, 30)
    lat = []
    for I in Is:
        _, _, _, sp = simulate_neuron(I, [(0.0, 0.2)], T=0.2)
        lat.append(sp[0] * 1e3 if len(sp) else np.nan)
    fig, ax = plt.subplots(figsize=(6, 4))
    ax.plot(Is, lat, "o-", color="C0", ms=4)
    ax.set_xscale("log")
    ax.set_xlabel("Irradiance (mW/mm²)")
    ax.set_ylabel("First-spike latency (ms)")
    ax.set_title("Brighter light, faster spike — down to a few ms")
    i5 = np.argmin(np.abs(Is - 5))
    print(f"[LIF] latency at {Is[i5]:.1f} mW/mm² ≈ {lat[i5]:.1f} ms")
    fig.tight_layout()
    fig.savefig("sim4_latency.png")
    plt.close(fig)


if __name__ == "__main__":
    sim_particle_in_box()
    sim_landau_zener()
    sim_neuron()
    sim_latency()
    print("Figures written: sim1_particle_in_box.png, sim2_landau_zener.png, "
          "sim3_neuron_pulses.png, sim4_latency.png")
