# -*- coding: utf-8 -*-
"""
Ti:Sa (~800 nm) の 10 fs 級ガウシアン波束を、離散周波数成分の重ね合わせで可視化。
改訂点（重要）:
  - アニメ (τ=z/v) の安定表示: ani 参照保持, blit=False 既定
  - 波束中心追従 (--auto-center) で窓外逸脱を回避
  - 軸レンジを指定/自動切替 (--xlim, --ylim, --auto-ylim)
  - デバッグ表示 (--show-debug)
"""

import argparse
import numpy as np
import matplotlib.pyplot as plt
from matplotlib.animation import FuncAnimation
import time

# 物理定数
C = 299_792_458.0
LAMBDA0 = 800e-9
NU0 = C / LAMBDA0
W0 = 2*np.pi*NU0
TBP_GAUSS = 0.44

def parse_list_or_range(text, unit_scale=1.0):
    if text is None:
        return None
    s = str(text).strip()
    if ":" in s:
        a, b, n = s.split(":")
        a = float(a); b = float(b); n = int(n)
        vals = np.linspace(a, b, n)
    else:
        vals = [float(x) for x in s.split(",")]
    return [v*unit_scale for v in vals]

def gaussian_amp_from_fwhm(x, fwhm):
    sigma = fwhm / (2*np.sqrt(2*np.log(2)))
    return np.exp(-0.25 * (x/sigma)**2)

def build_discrete_spectrum(nwave, tauf_fs, span_fwhm_scale=3.0):
    tauf = tauf_fs*1e-15
    dnu_fwhm = TBP_GAUSS / tauf
    half_span = 0.5 * dnu_fwhm * span_fwhm_scale
    dnu = np.linspace(-half_span, half_span, nwave)
    nu_k = NU0 + dnu
    A_k = gaussian_amp_from_fwhm(dnu, dnu_fwhm)
    A_k /= A_k.max()
    return nu_k, A_k, dnu, dnu_fwhm

def make_phase(nu_k, mode="flat", phi1=0.0, phi2=0.0, seed=None):
    w_k = 2*np.pi*nu_k
    dw = w_k - W0
    if mode == "flat":
        phi_k = np.zeros_like(nu_k)
    elif mode == "linear":
        phi_k = phi1*dw
    elif mode == "quad":
        phi_k = phi1*dw + 0.5*phi2*dw**2
    elif mode == "random":
        rng = np.random.default_rng(seed)
        phi_k = rng.uniform(-np.pi, np.pi, size=nu_k.size)
    else:
        raise ValueError("unknown phase mode")
    return np.unwrap(phi_k)

def synthesize_discrete_sum(nu_k, A_k, phi_k, t_window_fs=200.0, nsamp=4000):
    t = np.linspace(-0.5*t_window_fs, 0.5*t_window_fs, nsamp)*1e-15
    two_pi_t = 2*np.pi*t
    comps = np.array([A_k[i]*np.cos(two_pi_t*nu_k[i] + phi_k[i]) for i in range(len(nu_k))])
    esum = comps.sum(axis=0)
    esum /= np.max(np.abs(esum))
    I = esum**2
    I /= I.max()
    return t*1e15, comps/np.max(np.abs(esum)), esum, I  # fs

def time_signal_with_dispersion(z, nu_k, A_k, base_phase_k, v0, eta, t_window_fs, nsamp,
                                auto_center=True, return_tau=False):
    """
    v(nu) = v0 / (1 + eta*(nu-NU0))
    tau_k(z) = z / v(nu_k)
    E_z(t) = Σ A_k cos( 2π nu_k (t - tau_k(z)) + base_phase_k )
    auto_center=True の場合、強度重み A_k^2 で平均遅延 <tau> を引いて中心追従。
    """
    # 時間軸（fs→s）
    t_fs = np.linspace(-0.5*t_window_fs, 0.5*t_window_fs, nsamp)
    t = t_fs * 1e-15
    two_pi = 2*np.pi

    v_k = v0 / (1.0 + eta*(nu_k - NU0))
    tau_k = z / v_k  # [s]
    if auto_center:
        w = A_k**2
        tau_c = np.sum(w * tau_k) / np.sum(w)
    else:
        tau_c = 0.0

    E = np.zeros_like(t)
    for nu, A, phi0, tau in zip(nu_k, A_k, base_phase_k, tau_k):
        E += A * np.cos(two_pi*nu*(t - (tau - tau_c)) + phi0)

    E /= max(1e-12, np.max(np.abs(E)))
    I = (E**2)
    I /= max(1e-12, I.max())
    if return_tau:
        return t_fs, E, I, tau_k*1e15, tau_c*1e15  # τ を fs で返す
    return t_fs, E, I

def main():
    ap = argparse.ArgumentParser(description="Gaussian離散スペクトルの重ね合わせと分散アニメ（τ=z/v）可視化")
    ap.add_argument("--nwave", type=int, default=101, help="単色波の本数（奇数推奨）")
    ap.add_argument("--tauf", type=float, default=10.0, help="TLパルス幅 FWHM [fs]（Gaussian）")
    ap.add_argument("--twin", type=float, default=200.0, help="時間表示窓の全幅 [fs]")
    ap.add_argument("--nsamp", type=int, default=4000, help="時間サンプル数（静的/アニメ共用）")

    ap.add_argument("--phase-mode", choices=["flat","linear","quad","random"], default="flat")
    ap.add_argument("--phi1", type=float, default=0.0, help="φ1 [s]（群遅延）")
    ap.add_argument("--phi2", type=float, default=0.0, help="φ2 [s^2]（GDD）")
    ap.add_argument("--phi2-fs2", type=float, default=None, help="φ2 を fs^2 で与える（例: 10 → 10 fs^2）")
    ap.add_argument("--seed", type=int, default=None, help="random 位相の乱数seed")

    ap.add_argument("--chirp", type=str, default=None,
                    help="φ2(GDD) 掃引（fs^2）。'a,b,c' or 'start:end:steps' 形式（静的比較）")

    # アニメ（τ=z/v モデル）
    ap.add_argument("--anim", action="store_true", help="分散による波束崩壊をアニメ表示（τ=z/v）")
    ap.add_argument("--eta", type=float, default=2e-15, help="分散強度 [1/Hz]（小さめに）")
    ap.add_argument("--zmax", type=float, default=0.01, help="最大伝搬距離 [m]")
    ap.add_argument("--nframes", type=int, default=80, help="アニメフレーム数")
    ap.add_argument("--interval", type=int, default=80, help="フレーム間隔 [ms]")
    ap.add_argument("--sleep", type=float, default=0.0, help="各フレーム後に time.sleep(s) を入れる")

    # 軸レンジ・追従オプション
    ap.add_argument("--xlim", type=float, nargs=2, default=None, help="時間軸 [fs] の範囲を指定 (tmin tmax)")
    ap.add_argument("--ylim", type=float, nargs=2, default=None, help="縦軸の範囲を指定 (ymin ymax)")
    ap.add_argument("--auto-ylim", dest="auto_ylim", action="store_true", help="縦軸を自動レンジ")
    ap.add_argument("--no-auto-ylim", dest="auto_ylim", action="store_false", help="縦軸を固定（--ylim必須）")
    ap.set_defaults(auto_ylim=True)
    ap.add_argument("--auto-center", dest="auto_center", action="store_true", help="波束中心に追従（既定）")
    ap.add_argument("--no-auto-center", dest="auto_center", action="store_false", help="追従しない")
    ap.set_defaults(auto_center=True)
    ap.add_argument("--show-debug", action="store_true", help="各フレームの τ 統計を表示")

    args = ap.parse_args()

    if args.phi2_fs2 is not None:
        args.phi2 = args.phi2_fs2 * 1e-30  # fs^2 → s^2

    chirp_list_s2 = parse_list_or_range(args.chirp, unit_scale=1e-30)

    # スペクトルと基底位相
    nu_k, A_k, dnu, dnu_fwhm = build_discrete_spectrum(args.nwave, args.tauf, span_fwhm_scale=3.0)
    base_phi_k = make_phase(nu_k, mode=args.phase_mode, phi1=args.phi1, phi2=args.phi2, seed=args.seed)

    # Fig.A: 離散強度スペクトル
    figA, axA = plt.subplots()
    for x, y in zip(dnu*1e-12, (A_k**2)):
        axA.plot([x, x], [0, y], linewidth=1)
        axA.plot(x, y, marker="o", markersize=4)
    axA.set_xlabel("Frequency offset from ν0 [THz]")
    axA.set_ylabel("Spectral intensity (discrete)")
    axA.set_title(f"Amplitude spectrum (intensity)  nwave={args.nwave}")
    axA.grid(True, alpha=0.3)

    # Fig.B: 振幅 A_k
    figB, axB = plt.subplots()
    axB.plot(dnu*1e-12, A_k, marker="o", linewidth=1.0)
    axB.set_xlabel("Frequency offset from ν0 [THz]")
    axB.set_ylabel("Spectral amplitude A_k (discrete)")
    axB.set_title("Component amplitude vs frequency offset")
    axB.grid(True, alpha=0.3)

    # Fig.C: 分解・合成・|E|^2（静的）
    t_stat, comps_stat, esum_stat, I_stat = synthesize_discrete_sum(
        nu_k, A_k, base_phi_k, t_window_fs=args.twin, nsamp=args.nsamp
    )
    figC, axC = plt.subplots()
    picks = [0, args.nwave//2, args.nwave-1]
    for p in sorted(set([max(0, p) for p in picks])):
        axC.plot(t_stat, comps_stat[p], linewidth=0.9, alpha=0.55, label=f"component #{p}")
    axC.plot(t_stat, esum_stat, linewidth=2.0, label="sum (field)")
    axC.set_xlabel("Time [fs]")
    axC.set_ylabel("Field (arb.)")
    axC.set_title("Wave packet (components + sum)")
    axC.grid(True, alpha=0.3)
    axD = axC.twinx()
    axD.plot(t_stat, I_stat, linestyle="--", linewidth=1.6, label="|E|^2 (intensity)")
    axD.set_ylabel("Intensity (normalized)")
    h1,l1 = axC.get_legend_handles_labels()
    h2,l2 = axD.get_legend_handles_labels()
    axC.legend(h1+h2, l1+l2, loc="upper right")

    # Fig.D: φ2(GDD) 掃引（静的）
    if chirp_list_s2 is not None and len(chirp_list_s2) > 0:
        figD, axD2 = plt.subplots()
        for phi2_s2 in chirp_list_s2:
            phi_k = make_phase(nu_k, mode="quad", phi1=args.phi1, phi2=phi2_s2)
            t_, _, esum_, I_ = synthesize_discrete_sum(nu_k, A_k, phi_k,
                                                       t_window_fs=args.twin, nsamp=args.nsamp)
            axD2.plot(t_, esum_, label=f"φ2={phi2_s2/1e-30:.1f} fs²")
        axD2.set_xlabel("Time [fs]")
        axD2.set_ylabel("Field (arb.)")
        axD2.set_title("Pulse shape vs GDD (chirp sweep)")
        axD2.grid(True, alpha=0.3)
        axD2.legend(loc="best")

    # アニメ（τ=z/v）
    if args.anim:
        v0 = C
        z_vals = np.linspace(0.0, args.zmax, args.nframes)

        figE, axE = plt.subplots()
        line_field, = axE.plot([], [], lw=2, label="field (E)")
        line_int,   = axE.plot([], [], lw=1.4, ls="--", label="|E|^2")
        # 軸レンジ
        if args.xlim is not None:
            axE.set_xlim(args.xlim[0], args.xlim[1])
        else:
            axE.set_xlim(t_stat[0], t_stat[-1])
        if args.auto_ylim:
            axE.set_ylim(-1.1, 1.1)
        else:
            if args.ylim is None:
                raise SystemExit("--no-auto-ylim を使う場合は --ylim ymin ymax を指定してください。")
            axE.set_ylim(args.ylim[0], args.ylim[1])

        axE.set_xlabel("Time [fs]")
        axE.set_ylabel("Field / Intensity (arb.)")
        axE.set_title("Wave packet under dispersion: τ = z / v(ν)")
        axE.grid(True, alpha=0.3)
        axE.legend(loc="upper right")

        def init():
            line_field.set_data([], [])
            line_int.set_data([], [])
            return line_field, line_int

        def update(i):
            z = z_vals[i]
            t_anim, E_anim, I_anim, tau_k_fs, tau_c_fs = time_signal_with_dispersion(
                z, nu_k, A_k, base_phi_k, v0=v0, eta=args.eta,
                t_window_fs=args.twin, nsamp=args.nsamp,
                auto_center=args.auto_center, return_tau=True
            )
            line_field.set_data(t_anim, E_anim)
            line_int.set_data(t_anim, I_anim)

            # 縦軸オート調整（希望者向け）
            if args.auto_ylim:
                ypad = 0.05
                ymin = min(E_anim.min(), I_anim.min()) - ypad
                ymax = max(E_anim.max(), I_anim.max()) + ypad
                axE.set_ylim(ymin, ymax)

            axE.set_title(f"Wave packet under dispersion: τ=z/v  (z = {z*1e3:.2f} mm)")
            if args.show_debug:
                print(f"[frame {i+1}/{len(z_vals)}] z={z:.4f} m  tau[min,mean,max]={tau_k_fs.min():.1f}, {tau_c_fs:.1f}, {tau_k_fs.max():.1f} fs")

            if args.sleep > 0:
                time.sleep(args.sleep)
            return line_field, line_int

        # ★ ani 参照を保持（重要）
        ani = FuncAnimation(figE, update, frames=len(z_vals), init_func=init,
                            interval=args.interval, blit=False)

    plt.show()

if __name__ == "__main__":
    main()
