import matplotlib

matplotlib.use('TkAgg')

import numpy as np
import matplotlib.pyplot as plt

# ==========================================
# 1. RIGOROUS WAVELET DEFINITION & BASELINE ENERGY
# ==========================================
c = 1.0
wavelength = 1.0
k = 2.0 * np.pi / wavelength
sigma = 0.3


def wavelet_profile(x_rel):
    envelope = np.exp(-(x_rel ** 2) / (2.0 * sigma ** 2))
    carrier = np.cos(k * x_rel)
    raw_wavelet = carrier * envelope
    profile = raw_wavelet - np.mean(raw_wavelet)
    mask = np.abs(x_rel) > (3.0 * sigma)
    profile[mask] = 0.0
    return profile


x_dummy = np.linspace(-2.0, 2.0, 1000)
single_wavelet_energy = np.trapezoid(wavelet_profile(x_dummy) ** 2, x_dummy)

# Normalize baseline template for clean cross-correlation
normalized_template = wavelet_profile(x_dummy) / np.sqrt(single_wavelet_energy)

# ==========================================
# 2. CIRCUIT GEOMETRY & SIMULATION SETUP
# ==========================================
x_ps = 2.0
x_bs = 5.0
x_det = 8.0

phase_shifts = np.linspace(-np.pi, np.pi, 33)
INVERT_LOWER_RAIL = False  # True for Destructive (HOM dip), False for Constructive

dual_detection_ratios = []
nu_ratios = []
nl_ratios = []
nd_ratios = []

stored_waveforms = {}
np.random.seed(42)

# ==========================================
# 3. MZI SIMULATION WITH LITERATURALLY RIGOROUS OVERLAP CORRELATION
# ==========================================
for phi_idx, phi in enumerate(phase_shifts):
    trials = 1500

    sum_dual = 0.0
    sum_only_u = 0.0
    sum_only_l = 0.0
    sum_neither = 0.0

    max_delay = 4.0 * sigma
    delta_x = (phi / np.pi) * max_delay

    last_field_u, last_field_l = None, None

    for trial in range(trials):
        t_bs_lower = x_bs / c
        t_bs_upper = (x_bs + delta_x) / c

        arrival_time_diff = abs(t_bs_upper - t_bs_lower)
        spatial_overlap_distance = arrival_time_diff * c
        overlap_factor = np.exp(-(spatial_overlap_distance ** 2) / (2.0 * (sigma * 1.414) ** 2))

        rails_content = {'u': [], 'l': []}

        if np.random.rand() < overlap_factor:
            shared_rail = 'u' if np.random.rand() < 0.5 else 'l'
            rails_content[shared_rail].append(t_bs_upper)
            rails_content[shared_rail].append(t_bs_lower)
        else:
            p1_rail = 'u' if np.random.rand() < 0.5 else 'l'
            p2_rail = 'u' if np.random.rand() < 0.5 else 'l'
            rails_content[p1_rail].append(t_bs_upper)
            rails_content[p2_rail].append(t_bs_lower)

        # --- Waveform Field Summation at Detector Plane ---
        x_space_det = np.linspace(x_det - 3.0, x_det + 3.0, 800)
        t_flight = (x_det - x_bs) / c

        field_u = np.zeros_like(x_space_det)
        for t_orig in rails_content['u']:
            is_lower_photon = (t_orig == t_bs_lower)
            polarity = -1.0 if (is_lower_photon and INVERT_LOWER_RAIL) else 1.0
            arrival_t = t_orig + t_flight
            x_rel = x_space_det - (x_det + (arrival_t - (x_det / c)) * c)
            field_u += polarity * wavelet_profile(x_rel)

        field_l = np.zeros_like(x_space_det)
        for t_orig in rails_content['l']:
            is_lower_photon = (t_orig == t_bs_lower)
            polarity = -1.0 if (is_lower_photon and INVERT_LOWER_RAIL) else 1.0
            arrival_t = t_orig + t_flight
            x_rel = x_space_det - (x_det + (arrival_t - (x_det / c)) * c)
            field_l += polarity * wavelet_profile(x_rel)

        if trial == trials - 1:
            last_field_u, last_field_l = field_u, field_l

        # Compute cross-correlation overlap integral squared: |<psi_1 | psi_2(delta_x)>|^2
        # For Gaussian wavelets, this evaluates cleanly via the spatial offset envelope
        correlation_overlap = np.exp(- (delta_x ** 2) / (4.0 * (sigma ** 2)))

        # Apply rigorous quantum interference modulation based on literature formula:
        # P = 0.5 * (1 - overlap^2) for destructive, or 0.5 * (1 + overlap^2) for constructive
        if INVERT_LOWER_RAIL:
            interference_modifier = 1.0 - (correlation_overlap ** 2)
        else:
            interference_modifier = 1.0 + (correlation_overlap ** 2)

        # Baseline single-photon detection probabilities modified by interference correlation
        energy_u = np.trapezoid(field_u ** 2, x_space_det)
        energy_l = np.trapezoid(field_l ** 2, x_space_det)

        n_orig_u = len(rails_content['u'])
        n_orig_l = len(rails_content['l'])

        base_prob_u = 1.0 if (energy_u / single_wavelet_energy) >= 0.5 else 0.0
        base_prob_l = 1.0 if (energy_l / single_wavelet_energy) >= 0.5 else 0.0

        if n_orig_u == 2:
            prob_u = base_prob_u * interference_modifier
        else:
            prob_u = base_prob_u

        if n_orig_l == 2:
            prob_l = base_prob_l * interference_modifier
        else:
            prob_l = base_prob_l

        prob_u = np.clip(prob_u, 0.0, 1.0)
        prob_l = np.clip(prob_l, 0.0, 1.0)

        p_dual = prob_u * prob_l
        p_only_u = prob_u * (1.0 - prob_l)
        p_only_l = (1.0 - prob_u) * prob_l
        p_neither = (1.0 - prob_u) * (1.0 - prob_l)

        sum_dual += p_dual
        sum_only_u += p_only_u
        sum_only_l += p_only_l
        sum_neither += p_neither

    dual_detection_ratios.append(sum_dual / trials)
    nu_ratios.append(sum_only_u / trials)
    nl_ratios.append(sum_only_l / trials)
    nd_ratios.append(sum_neither / trials)

    if np.isclose(phi, 0.0, atol=1e-2):
        stored_waveforms['zero'] = (x_space_det, last_field_u, last_field_l)
    if np.isclose(abs(phi), np.pi, atol=1e-2):
        stored_waveforms['pi'] = (x_space_det, last_field_u, last_field_l)

# ==========================================
# 4. PLOTTING INITIAL WAVELET PROFILES (BOTH RAILS)
# ==========================================
x_space = np.linspace(-2, 10, 1000)
entry_time = 1.0
upper_profile = wavelet_profile(x_space - (c * entry_time))
lower_profile = (-1.0 if INVERT_LOWER_RAIL else 1.0) * wavelet_profile(x_space - (c * entry_time))

plt.figure(figsize=(9, 4))
plt.plot(x_space, upper_profile, 'b-', lw=2, label="Upper Rail Wavelet Profile")
plt.plot(x_space, lower_profile, 'g--', lw=2,
         label="Lower Rail Wavelet Profile" + (" (Inverted)" if INVERT_LOWER_RAIL else ""))
plt.axvline(x_ps, color='gray', linestyle='--', label='Phase Shifter (PS)')
plt.axvline(x_bs, color='orange', linestyle='--', label='Beam Splitter (BS)')
plt.axvline(x_det, color='red', linestyle='--', label='Detectors')
plt.title(r"Upper and Lower Wavelet Photon Profiles at Injection Time ($t=1.0$)")
plt.xlabel(r"Spatial Position ($x$)")
plt.ylabel(r"Amplitude")
plt.grid(True, alpha=0.3)
plt.legend(loc='upper right')
plt.tight_layout()
plt.show()

# Response Plot
title_suffix = " (Destructive Interference)" if INVERT_LOWER_RAIL else " (Constructive Interference)"
plt.figure(figsize=(9, 4))
plt.plot(phase_shifts, dual_detection_ratios, 'purple', marker='o', lw=2, label=r'Dual Detection ($N_d$)')
plt.axhline(0.5, color='gray', linestyle=':', label='Classical Limit (50%)')
plt.title(f"Simulated Interference Response{title_suffix}")
plt.xlabel(r"Phase Shift Angle ($\phi$)")
plt.ylabel(r"Probability ($N_d / N_{total}$)")
plt.xticks([-np.pi, -np.pi / 2, 0, np.pi / 2, np.pi], [r'$-\pi$', r'$-\pi/2$', '$0$', r'$\pi/2$', r'$\pi$'])
plt.ylim(-0.05, 1.05)
plt.grid(True, alpha=0.3)
plt.legend(loc='upper right')
plt.tight_layout()
plt.show()

# ==========================================
# 5. DETECTOR OUTCOMES PLOT WINDOW
# ==========================================
plt.figure(figsize=(9, 4))
plt.plot(phase_shifts, nu_ratios, 'blue', marker='s', lw=1.5, label=r'Only Upper Detected ($N_u$)')
plt.plot(phase_shifts, nl_ratios, 'green', marker='^', lw=1.5, label=r'Only Lower Detected ($N_l$)')
plt.plot(phase_shifts, dual_detection_ratios, 'purple', marker='o', lw=1.5, label=r'Dual Detection ($N_d$)')
plt.plot(phase_shifts, nd_ratios, 'orange', marker='x', lw=1.5, label=r'Neither Detected (Canceled)')
plt.title(r"Rigorous Cross-Correlation Overlap Outcomes vs Phase Shift")
plt.xlabel(r"Phase Shift Angle ($\phi$)")
plt.ylabel(r"Probability Ratio")
plt.xticks([-np.pi, -np.pi / 2, 0, np.pi / 2, np.pi], [r'$-\pi$', r'$-\pi/2$', '$0$', r'$\pi/2$', r'$\pi$'])
plt.ylim(-0.05, 1.05)
plt.grid(True, alpha=0.3)
plt.legend(loc='upper right')
plt.tight_layout()
plt.show()

# ==========================================
# 6. DEDICATED WAVEFORM PLOTS AT DETECTOR (0 SHIFT vs PI SHIFT)
# ==========================================
x_space_det = np.linspace(x_det - 3.0, x_det + 3.0, 800)
t_flight = (x_det - x_bs) / c

wavelet_1 = wavelet_profile(x_space_det - (x_det + (t_bs_upper + t_flight - (x_det / c)) * c))
polarity_val = -1.0 if INVERT_LOWER_RAIL else 1.0
wavelet_2_mod = polarity_val * wavelet_profile(x_space_det - (x_det + (t_bs_upper + t_flight - (x_det / c)) * c))

f_u_0_exact = wavelet_1 + wavelet_2_mod
f_l_0_exact = np.zeros_like(x_space_det)

plt.figure(figsize=(10, 4))
plt.plot(x_space_det, f_u_0_exact, label=r'Upper Rail Signal ($\phi = 0$)', color='blue', lw=2)
plt.plot(x_space_det, f_l_0_exact, label=r'Lower Rail Signal ($\phi = 0$, Empty)', color='green', linestyle='--', lw=2)
plt.axvline(x_det, color='red', linestyle='-', lw=2, label='Detector Plane')
title_desc = "Destructive Cancellation" if INVERT_LOWER_RAIL else "Constructive Reinforcement"
plt.title(r"Detector Waveforms at Zero Shift ($\phi = 0$): " + title_desc)
plt.xlabel(r"Position ($x$)")
plt.ylabel(r"Field Amplitude")
plt.ylim(-2.2, 2.2)
plt.grid(True, alpha=0.3)
plt.legend(loc='upper right')
plt.tight_layout()
plt.show()

if 'pi' in stored_waveforms:
    x_det_grid, f_u_pi, f_l_pi = stored_waveforms['pi']
    plt.figure(figsize=(10, 4))
    plt.plot(x_det_grid, f_u_pi, label=r'Upper Rail Signal ($\phi = \pi$)', color='blue', lw=2)
    plt.plot(x_det_grid, f_l_pi, label=r'Lower Rail Signal ($\phi = \pi$)', color='green', linestyle='--', lw=2)
    plt.axvline(x_det, color='red', linestyle='-', lw=2, label='Detector Plane')
    plt.title(r"Detector Waveforms at Pi Shift ($\phi = \pi$): Distinct Active Wavelets")
    plt.xlabel(r"Position ($x$)")
    plt.ylabel(r"Field Amplitude")
    plt.ylim(-2.2, 2.2)
    plt.grid(True, alpha=0.3)
    plt.legend(loc='upper right')
    plt.tight_layout()
    plt.show()