#!/usr/bin/env python3
# SPDX-License-Identifier: MIT
"""Verified Jaynes-Cummings and small open-cavity benchmarks.

The program computes six related artifacts:

1. numerical and analytic Jaynes-Cummings dressed doublets;
2. vacuum Rabi exchange from |e,0>;
3. coherent-state collapse and revival;
4. photon-cutoff convergence and omitted Poisson weight;
5. an optional unconditional Lindblad cavity-loss extension;
6. machine-readable validation and convention metadata.

NumPy is the only dependency. The closed dynamics use hbar = g = 1 unless
otherwise stated. The open extension reports kappa as the cavity
energy-decay rate and gamma as the emitter population-decay rate.
"""

from __future__ import annotations

import argparse
import csv
import json
import math
import platform
from dataclasses import dataclass
from pathlib import Path
from typing import Any

import numpy as np


IDENTITY_ATOM = np.eye(2, dtype=complex)
SIGMA_PLUS = np.array(
    [
        [0.0, 1.0],
        [0.0, 0.0],
    ],
    dtype=complex,
)
SIGMA_MINUS = SIGMA_PLUS.conj().T
SIGMA_Z = np.array(
    [
        [1.0, 0.0],
        [0.0, -1.0],
    ],
    dtype=complex,
)
PROJECTOR_EXCITED_ATOM = np.array(
    [
        [1.0, 0.0],
        [0.0, 0.0],
    ],
    dtype=complex,
)

G_COUPLING = 1.0
COHERENT_NBAR = 16.0
COHERENT_CUTOFFS = (20, 30, 40, 60)
COHERENT_REFERENCE_CUTOFF = 100
LOSS_CASES = (
    ("closed", 0.0, 0.0),
    ("balanced", 0.2, 0.2),
    ("leaky_cavity", 1.0, 0.2),
)


@dataclass(frozen=True)
class PureEvolution:
    states: np.ndarray
    max_norm_error: float
    max_excitation_drift: float


@dataclass(frozen=True)
class DensityEvolution:
    density_matrices: np.ndarray
    max_trace_error: float
    max_hermiticity_error: float
    minimum_eigenvalue: float


def annihilation_operator(cutoff: int) -> np.ndarray:
    if cutoff < 2:
        raise ValueError("photon cutoff must be at least 2")
    operator = np.zeros((cutoff, cutoff), dtype=complex)
    for number in range(1, cutoff):
        operator[number - 1, number] = math.sqrt(number)
    return operator


def field_number_operator(cutoff: int) -> np.ndarray:
    return np.diag(np.arange(cutoff, dtype=float)).astype(complex)


def jc_interaction_hamiltonian(
    cutoff: int,
    *,
    g: float,
    delta: float,
) -> np.ndarray:
    """Return the rotating-frame Jaynes-Cummings Hamiltonian divided by hbar."""
    identity_field = np.eye(cutoff, dtype=complex)
    annihilation = annihilation_operator(cutoff)
    creation = annihilation.conj().T
    return (
        0.5 * delta * np.kron(SIGMA_Z, identity_field)
        + g
        * (
            np.kron(SIGMA_PLUS, annihilation)
            + np.kron(SIGMA_MINUS, creation)
        )
    )


def full_jc_hamiltonian(
    cutoff: int,
    *,
    omega_c: float,
    omega_a: float,
    g: float,
) -> np.ndarray:
    identity_field = np.eye(cutoff, dtype=complex)
    number = field_number_operator(cutoff)
    annihilation = annihilation_operator(cutoff)
    creation = annihilation.conj().T
    return (
        omega_c * np.kron(IDENTITY_ATOM, number)
        + 0.5 * omega_a * np.kron(SIGMA_Z, identity_field)
        + g
        * (
            np.kron(SIGMA_PLUS, annihilation)
            + np.kron(SIGMA_MINUS, creation)
        )
    )


def total_excitation_operator(cutoff: int) -> np.ndarray:
    identity_field = np.eye(cutoff, dtype=complex)
    number = field_number_operator(cutoff)
    return (
        np.kron(IDENTITY_ATOM, number)
        + np.kron(PROJECTOR_EXCITED_ATOM, identity_field)
    )


def excited_projector(cutoff: int) -> np.ndarray:
    return np.kron(
        PROJECTOR_EXCITED_ATOM,
        np.eye(cutoff, dtype=complex),
    )


def photon_number_composite(cutoff: int) -> np.ndarray:
    return np.kron(
        IDENTITY_ATOM,
        field_number_operator(cutoff),
    )


def basis_state(atom_excited: bool, number: int, cutoff: int) -> np.ndarray:
    if number < 0 or number >= cutoff:
        raise ValueError("photon number lies outside cutoff")
    atom = np.array(
        [1.0, 0.0] if atom_excited else [0.0, 1.0],
        dtype=complex,
    )
    field = np.zeros(cutoff, dtype=complex)
    field[number] = 1.0
    return np.kron(atom, field)


def expectation(
    states: np.ndarray,
    operator: np.ndarray,
) -> np.ndarray:
    return np.einsum(
        "it,ij,jt->t",
        states.conj(),
        operator,
        states,
        optimize=True,
    ).real


def evolve_pure_state(
    hamiltonian: np.ndarray,
    initial_state: np.ndarray,
    times: np.ndarray,
    excitation: np.ndarray,
) -> PureEvolution:
    eigenvalues, eigenvectors = np.linalg.eigh(hamiltonian)
    coefficients = eigenvectors.conj().T @ initial_state
    phases = np.exp(
        -1.0j
        * eigenvalues[:, np.newaxis]
        * times[np.newaxis, :]
    )
    states = eigenvectors @ (coefficients[:, np.newaxis] * phases)
    norms = np.sum(np.abs(states) ** 2, axis=0)
    excitations = expectation(states, excitation)
    return PureEvolution(
        states=states,
        max_norm_error=float(np.max(np.abs(norms - 1.0))),
        max_excitation_drift=float(
            np.max(np.abs(excitations - excitations[0]))
        ),
    )


def dressed_spectrum_rows(
    samples: int,
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
    g = 0.1
    omega_c = 1.0
    detuning_ratios = np.linspace(-6.0, 6.0, samples)
    rows: list[dict[str, Any]] = []
    energy_error = 0.0
    eigenvector_norm_error = 0.0
    resonance_splitting_errors: dict[str, float] = {}

    for ratio in detuning_ratios:
        delta = float(ratio * g)
        row: dict[str, Any] = {
            "Delta_over_g": ratio,
        }
        for manifold_n in (0, 1):
            coupling = g * math.sqrt(manifold_n + 1.0)
            block = np.array(
                [
                    [0.5 * delta, coupling],
                    [coupling, -0.5 * delta],
                ],
                dtype=float,
            )
            eigenvalues, eigenvectors = np.linalg.eigh(block)
            generalized = math.sqrt(
                delta * delta
                + 4.0 * g * g * (manifold_n + 1.0)
            )
            analytic = np.array(
                [-0.5 * generalized, 0.5 * generalized]
            )
            energy_error = max(
                energy_error,
                float(np.max(np.abs(eigenvalues - analytic))),
            )
            eigenvector_norm_error = max(
                eigenvector_norm_error,
                float(
                    np.max(
                        np.abs(
                            eigenvectors.conj().T @ eigenvectors
                            - np.eye(2)
                        )
                    )
                ),
            )
            suffix = f"n{manifold_n}"
            centre = omega_c * (manifold_n + 0.5)
            row[f"E_lower_{suffix}"] = centre + eigenvalues[0]
            row[f"E_upper_{suffix}"] = centre + eigenvalues[1]
            row[
                f"shift_lower_over_g_{suffix}"
            ] = eigenvalues[0] / g
            row[
                f"shift_upper_over_g_{suffix}"
            ] = eigenvalues[1] / g
            row[
                f"atomic_fraction_lower_{suffix}"
            ] = abs(eigenvectors[0, 0]) ** 2
            row[
                f"atomic_fraction_upper_{suffix}"
            ] = abs(eigenvectors[0, 1]) ** 2
            row[
                f"analytic_lower_over_g_{suffix}"
            ] = analytic[0] / g
            row[
                f"analytic_upper_over_g_{suffix}"
            ] = analytic[1] / g
        rows.append(row)

    centre_index = samples // 2
    for manifold_n in (0, 1):
        suffix = f"n{manifold_n}"
        splitting = (
            rows[centre_index][f"E_upper_{suffix}"]
            - rows[centre_index][f"E_lower_{suffix}"]
        )
        expected = 2.0 * g * math.sqrt(manifold_n + 1.0)
        resonance_splitting_errors[suffix] = abs(
            splitting - expected
        )

    validation = {
        "max_block_energy_error": energy_error,
        "max_eigenvector_orthonormality_error": (
            eigenvector_norm_error
        ),
        "resonance_splitting_errors": resonance_splitting_errors,
        "checks": {
            "analytic_dressed_energies": energy_error < 1.0e-14,
            "orthonormal_eigenvectors": (
                eigenvector_norm_error < 1.0e-14
            ),
            "vacuum_splitting": (
                resonance_splitting_errors["n0"] < 1.0e-14
            ),
            "sqrt_two_splitting": (
                resonance_splitting_errors["n1"] < 1.0e-14
            ),
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(
            f"dressed-spectrum validation failed: {validation}"
        )
    return rows, validation


def vacuum_rabi_rows(
    samples: int,
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
    cutoff = 4
    times = np.linspace(0.0, 4.0 * math.pi, samples)
    hamiltonian = jc_interaction_hamiltonian(
        cutoff,
        g=G_COUPLING,
        delta=0.0,
    )
    excitation = total_excitation_operator(cutoff)
    evolution = evolve_pure_state(
        hamiltonian,
        basis_state(True, 0, cutoff),
        times,
        excitation,
    )
    excited = expectation(
        evolution.states,
        excited_projector(cutoff),
    )
    photons = expectation(
        evolution.states,
        photon_number_composite(cutoff),
    )
    analytic_excited = np.cos(G_COUPLING * times) ** 2
    analytic_photons = np.sin(G_COUPLING * times) ** 2
    population_error = float(
        np.max(np.abs(excited - analytic_excited))
    )
    photon_error = float(
        np.max(np.abs(photons - analytic_photons))
    )
    exchange_error = float(
        np.max(np.abs(excited + photons - 1.0))
    )
    rows = [
        {
            "g_t": time,
            "g_t_over_pi": time / math.pi,
            "P_e_numeric": excited[index],
            "P_e_analytic": analytic_excited[index],
            "mean_photon_numeric": photons[index],
            "mean_photon_analytic": analytic_photons[index],
            "total_excitation_numeric": (
                excited[index] + photons[index]
            ),
        }
        for index, time in enumerate(times)
    ]
    validation = {
        "max_excited_population_error": population_error,
        "max_photon_number_error": photon_error,
        "max_exchange_sum_error": exchange_error,
        "max_state_norm_error": evolution.max_norm_error,
        "max_total_excitation_drift": (
            evolution.max_excitation_drift
        ),
        "checks": {
            "vacuum_rabi_probability": population_error < 1.0e-13,
            "vacuum_photon_probability": photon_error < 1.0e-13,
            "one_excitation_exchange": exchange_error < 1.0e-13,
            "unit_norm": evolution.max_norm_error < 1.0e-13,
            "excitation_conservation": (
                evolution.max_excitation_drift < 1.0e-13
            ),
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(
            f"vacuum-Rabi validation failed: {validation}"
        )
    return rows, validation


def poisson_probabilities(
    nbar: float,
    maximum_number: int,
) -> np.ndarray:
    if nbar < 0.0:
        raise ValueError("mean photon number must be nonnegative")
    probabilities = np.empty(maximum_number + 1, dtype=float)
    probabilities[0] = math.exp(-nbar)
    for number in range(maximum_number):
        probabilities[number + 1] = (
            probabilities[number]
            * nbar
            / (number + 1.0)
        )
    return probabilities


def truncated_coherent_state(
    cutoff: int,
    nbar: float,
) -> tuple[np.ndarray, float]:
    probabilities = poisson_probabilities(nbar, cutoff - 1)
    retained = float(np.sum(probabilities))
    if retained <= 0.0:
        raise ValueError("truncated coherent state has zero norm")
    amplitudes = np.sqrt(probabilities / retained).astype(complex)
    return amplitudes, max(0.0, 1.0 - retained)


def coherent_evolution(
    cutoff: int,
    times: np.ndarray,
) -> tuple[np.ndarray, dict[str, float]]:
    field, omitted_tail = truncated_coherent_state(
        cutoff,
        COHERENT_NBAR,
    )
    initial_state = np.kron(
        np.array([1.0, 0.0], dtype=complex),
        field,
    )
    hamiltonian = jc_interaction_hamiltonian(
        cutoff,
        g=G_COUPLING,
        delta=0.0,
    )
    excitation = total_excitation_operator(cutoff)
    evolution = evolve_pure_state(
        hamiltonian,
        initial_state,
        times,
        excitation,
    )
    excited = expectation(
        evolution.states,
        excited_projector(cutoff),
    )
    inversion = 2.0 * excited - 1.0
    return inversion, {
        "omitted_poisson_tail": omitted_tail,
        "max_state_norm_error": evolution.max_norm_error,
        "max_total_excitation_drift": (
            evolution.max_excitation_drift
        ),
    }


def infinite_coherent_inversion(times: np.ndarray) -> np.ndarray:
    probabilities = poisson_probabilities(COHERENT_NBAR, 199)
    tail = 1.0 - float(np.sum(probabilities))
    if abs(tail) > 1.0e-14:
        raise RuntimeError(
            f"infinite-sum Poisson tail is too large: {tail}"
        )
    frequencies = (
        2.0
        * G_COUPLING
        * np.sqrt(np.arange(200, dtype=float) + 1.0)
    )
    return probabilities @ np.cos(
        frequencies[:, np.newaxis] * times[np.newaxis, :]
    )


def collapse_revival_rows(
    samples: int,
) -> tuple[
    list[dict[str, Any]],
    list[dict[str, Any]],
    dict[str, Any],
]:
    revival_estimate = (
        2.0 * math.pi * math.sqrt(COHERENT_NBAR) / G_COUPLING
    )
    times = np.linspace(0.0, 2.1 * revival_estimate, samples)
    infinite = infinite_coherent_inversion(times)
    cutoffs = (*COHERENT_CUTOFFS, COHERENT_REFERENCE_CUTOFF)
    inversions: dict[int, np.ndarray] = {}
    diagnostics: dict[int, dict[str, float]] = {}

    for cutoff in cutoffs:
        inversions[cutoff], diagnostics[cutoff] = coherent_evolution(
            cutoff,
            times,
        )

    reference = inversions[COHERENT_REFERENCE_CUTOFF]
    truncation_rows: list[dict[str, Any]] = []
    largest_norm_error = 0.0
    largest_excitation_drift = 0.0
    for cutoff in COHERENT_CUTOFFS:
        diagnostic = diagnostics[cutoff]
        largest_norm_error = max(
            largest_norm_error,
            diagnostic["max_state_norm_error"],
        )
        largest_excitation_drift = max(
            largest_excitation_drift,
            diagnostic["max_total_excitation_drift"],
        )
        truncation_rows.append(
            {
                "photon_cutoff": cutoff,
                "highest_retained_photon_number": cutoff - 1,
                "omitted_poisson_probability": diagnostic[
                    "omitted_poisson_tail"
                ],
                "max_abs_inversion_error_vs_N100": float(
                    np.max(np.abs(inversions[cutoff] - reference))
                ),
                "max_abs_inversion_error_vs_infinite_sum": float(
                    np.max(np.abs(inversions[cutoff] - infinite))
                ),
                "max_state_norm_error": diagnostic[
                    "max_state_norm_error"
                ],
                "max_total_excitation_drift": diagnostic[
                    "max_total_excitation_drift"
                ],
            }
        )

    reference_error = float(np.max(np.abs(reference - infinite)))
    search = (
        (times >= 0.75 * revival_estimate)
        & (times <= 1.25 * revival_estimate)
    )
    search_indices = np.flatnonzero(search)
    local_index = int(
        np.argmax(np.abs(infinite[search_indices]))
    )
    revival_index = int(search_indices[local_index])
    rows: list[dict[str, Any]] = []
    for index, time in enumerate(times):
        row: dict[str, Any] = {
            "g_t": time,
            "g_t_over_revival_estimate": time / revival_estimate,
            "W_infinite_poisson_sum": infinite[index],
            "P_e_infinite_poisson_sum": 0.5 * (
                1.0 + infinite[index]
            ),
        }
        for cutoff in COHERENT_CUTOFFS:
            row[f"W_cutoff_{cutoff}"] = inversions[cutoff][index]
            row[f"P_e_cutoff_{cutoff}"] = 0.5 * (
                1.0 + inversions[cutoff][index]
            )
        rows.append(row)

    validation = {
        "mean_photon_number": COHERENT_NBAR,
        "collapse_time_1_over_e_estimate": (
            math.sqrt(2.0) / G_COUPLING
        ),
        "revival_time_estimate": revival_estimate,
        "largest_revival_window_abs_inversion": float(
            abs(infinite[revival_index])
        ),
        "largest_revival_window_time": float(times[revival_index]),
        "reference_cutoff": COHERENT_REFERENCE_CUTOFF,
        "max_reference_minus_infinite_sum": reference_error,
        "largest_state_norm_error": largest_norm_error,
        "largest_total_excitation_drift": largest_excitation_drift,
        "checks": {
            "reference_cutoff": reference_error < 1.0e-12,
            "unit_norm": largest_norm_error < 1.0e-12,
            "excitation_conservation": (
                largest_excitation_drift < 1.0e-11
            ),
            "revival_resolved": (
                abs(infinite[revival_index]) > 0.45
            ),
            "cutoff_converges": all(
                upper[
                    "max_abs_inversion_error_vs_infinite_sum"
                ]
                < lower[
                    "max_abs_inversion_error_vs_infinite_sum"
                ]
                for lower, upper in zip(
                    truncation_rows,
                    truncation_rows[1:],
                )
            ),
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(
            f"collapse-revival validation failed: {validation}"
        )
    return rows, truncation_rows, validation


def lindblad_rhs(
    density: np.ndarray,
    hamiltonian: np.ndarray,
    collapse_operators: tuple[np.ndarray, ...],
) -> np.ndarray:
    derivative = -1.0j * (
        hamiltonian @ density - density @ hamiltonian
    )
    for operator in collapse_operators:
        product = operator.conj().T @ operator
        derivative += (
            operator @ density @ operator.conj().T
            - 0.5 * (product @ density + density @ product)
        )
    return derivative


def rk4_density_step(
    density: np.ndarray,
    step: float,
    hamiltonian: np.ndarray,
    collapse_operators: tuple[np.ndarray, ...],
) -> np.ndarray:
    k1 = lindblad_rhs(density, hamiltonian, collapse_operators)
    k2 = lindblad_rhs(
        density + 0.5 * step * k1,
        hamiltonian,
        collapse_operators,
    )
    k3 = lindblad_rhs(
        density + 0.5 * step * k2,
        hamiltonian,
        collapse_operators,
    )
    k4 = lindblad_rhs(
        density + step * k3,
        hamiltonian,
        collapse_operators,
    )
    return density + step * (k1 + 2.0 * k2 + 2.0 * k3 + k4) / 6.0


def evolve_density(
    times: np.ndarray,
    *,
    initial_density: np.ndarray,
    hamiltonian: np.ndarray,
    collapse_operators: tuple[np.ndarray, ...],
    max_step: float,
) -> DensityEvolution:
    density = initial_density.copy()
    matrices = [density.copy()]
    trace_error = abs(np.trace(density) - 1.0)
    hermiticity_error = float(
        np.max(np.abs(density - density.conj().T))
    )
    minimum_eigenvalue = float(
        np.min(np.linalg.eigvalsh(0.5 * (density + density.conj().T)))
    )
    previous = float(times[0])

    for target_raw in times[1:]:
        target = float(target_raw)
        interval = target - previous
        steps = max(1, math.ceil(interval / max_step))
        step = interval / steps
        for _ in range(steps):
            density = rk4_density_step(
                density,
                step,
                hamiltonian,
                collapse_operators,
            )
        matrices.append(density.copy())
        trace_error = max(
            trace_error,
            abs(np.trace(density) - 1.0),
        )
        hermiticity_error = max(
            hermiticity_error,
            float(
                np.max(
                    np.abs(density - density.conj().T)
                )
            ),
        )
        hermitian = 0.5 * (density + density.conj().T)
        minimum_eigenvalue = min(
            minimum_eigenvalue,
            float(np.min(np.linalg.eigvalsh(hermitian))),
        )
        previous = target

    return DensityEvolution(
        density_matrices=np.asarray(matrices, dtype=complex),
        max_trace_error=float(abs(trace_error)),
        max_hermiticity_error=hermiticity_error,
        minimum_eigenvalue=minimum_eigenvalue,
    )


def density_expectation(
    density_matrices: np.ndarray,
    operator: np.ndarray,
) -> np.ndarray:
    return np.einsum(
        "tij,ji->t",
        density_matrices,
        operator,
        optimize=True,
    ).real


def loss_rows(
    samples: int,
    max_step: float,
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
    cutoff = 2
    times = np.linspace(0.0, 10.0, samples)
    hamiltonian = jc_interaction_hamiltonian(
        cutoff,
        g=G_COUPLING,
        delta=0.0,
    )
    initial = basis_state(True, 0, cutoff)
    initial_density = np.outer(initial, initial.conj())
    annihilation = np.kron(
        IDENTITY_ATOM,
        annihilation_operator(cutoff),
    )
    emitter_lowering = np.kron(
        SIGMA_MINUS,
        np.eye(cutoff, dtype=complex),
    )
    excited_operator = excited_projector(cutoff)
    photon_operator = photon_number_composite(cutoff)
    traces: dict[str, dict[str, Any]] = {}
    largest_trace_error = 0.0
    largest_hermiticity_error = 0.0
    minimum_eigenvalue = 1.0

    for name, kappa, gamma in LOSS_CASES:
        collapse_operators = tuple(
            operator
            for rate, operator in (
                (kappa, annihilation),
                (gamma, emitter_lowering),
            )
            if rate > 0.0
            for operator in (math.sqrt(rate) * operator,)
        )
        evolution = evolve_density(
            times,
            initial_density=initial_density,
            hamiltonian=hamiltonian,
            collapse_operators=collapse_operators,
            max_step=max_step,
        )
        excited = density_expectation(
            evolution.density_matrices,
            excited_operator,
        )
        photons = density_expectation(
            evolution.density_matrices,
            photon_operator,
        )
        traces[name] = {
            "excited": excited,
            "photons": photons,
            "total_excitation": excited + photons,
            "evolution": evolution,
        }
        largest_trace_error = max(
            largest_trace_error,
            evolution.max_trace_error,
        )
        largest_hermiticity_error = max(
            largest_hermiticity_error,
            evolution.max_hermiticity_error,
        )
        minimum_eigenvalue = min(
            minimum_eigenvalue,
            evolution.minimum_eigenvalue,
        )

    analytic_closed = np.cos(times) ** 2
    closed_error = float(
        np.max(
            np.abs(
                traces["closed"]["excited"] - analytic_closed
            )
        )
    )
    rows: list[dict[str, Any]] = []
    for index, time in enumerate(times):
        row: dict[str, Any] = {
            "g_t": time,
            "P_e_closed_analytic": analytic_closed[index],
        }
        for name, kappa, gamma in LOSS_CASES:
            row[f"P_e_{name}"] = traces[name]["excited"][index]
            row[f"mean_photon_{name}"] = traces[name]["photons"][
                index
            ]
            row[
                f"total_excitation_{name}"
            ] = traces[name]["total_excitation"][index]
            row[f"kappa_over_g_{name}"] = kappa
            row[f"gamma_over_g_{name}"] = gamma
        rows.append(row)

    validation = {
        "max_closed_population_error": closed_error,
        "largest_trace_error": largest_trace_error,
        "largest_hermiticity_error": largest_hermiticity_error,
        "minimum_density_eigenvalue": minimum_eigenvalue,
        "balanced_final_total_excitation": float(
            traces["balanced"]["total_excitation"][-1]
        ),
        "leaky_final_total_excitation": float(
            traces["leaky_cavity"]["total_excitation"][-1]
        ),
        "checks": {
            "closed_limit": closed_error < 2.0e-10,
            "trace_one": largest_trace_error < 1.0e-11,
            "hermitian": largest_hermiticity_error < 1.0e-11,
            "positive": minimum_eigenvalue > -1.0e-10,
            "loss_removes_excitation": bool(
                traces["leaky_cavity"]["total_excitation"][-1]
                < traces["balanced"]["total_excitation"][-1]
                < 1.0
            ),
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(
            f"cavity-loss validation failed: {validation}"
        )
    return rows, validation


def format_csv_value(value: Any) -> Any:
    if value is None:
        return ""
    if isinstance(value, (float, np.floating)):
        return f"{float(value):.16g}"
    if isinstance(value, (int, np.integer)):
        return int(value)
    return value


def json_scalar(value: Any) -> Any:
    """Convert NumPy scalar values without silently accepting other objects."""
    if isinstance(value, np.generic):
        return value.item()
    raise TypeError(
        f"Object of type {value.__class__.__name__} is not JSON serializable"
    )


def write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
    if not rows:
        raise ValueError(f"cannot write empty CSV: {path}")
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=list(rows[0]))
        writer.writeheader()
        for row in rows:
            writer.writerow(
                {
                    key: format_csv_value(value)
                    for key, value in row.items()
                }
            )


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--spectrum-samples", type=int, default=401)
    parser.add_argument("--vacuum-samples", type=int, default=801)
    parser.add_argument("--collapse-samples", type=int, default=1201)
    parser.add_argument("--loss-samples", type=int, default=501)
    parser.add_argument("--loss-max-step", type=float, default=0.0025)
    parser.add_argument(
        "--output-dir",
        type=Path,
        default=Path("cavity-qed-simulation-output"),
    )
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    for name in (
        "spectrum_samples",
        "vacuum_samples",
        "collapse_samples",
        "loss_samples",
    ):
        if getattr(args, name) < 5:
            raise ValueError(f"{name.replace('_', '-')} must be at least 5")
        if getattr(args, name) % 2 == 0:
            raise ValueError(f"{name.replace('_', '-')} must be odd")
    if args.loss_max_step <= 0.0:
        raise ValueError("loss-max-step must be positive")

    spectrum, spectrum_validation = dressed_spectrum_rows(
        args.spectrum_samples
    )
    vacuum, vacuum_validation = vacuum_rabi_rows(
        args.vacuum_samples
    )
    (
        collapse,
        truncation,
        collapse_validation,
    ) = collapse_revival_rows(args.collapse_samples)
    loss, loss_validation = loss_rows(
        args.loss_samples,
        args.loss_max_step,
    )

    args.output_dir.mkdir(parents=True, exist_ok=True)
    paths = {
        "spectrum": (
            args.output_dir / "cavity-qed-dressed-spectrum.csv"
        ),
        "vacuum": (
            args.output_dir / "cavity-qed-vacuum-rabi.csv"
        ),
        "collapse": (
            args.output_dir / "cavity-qed-collapse-revival.csv"
        ),
        "truncation": (
            args.output_dir / "cavity-qed-truncation.csv"
        ),
        "loss": (
            args.output_dir / "cavity-qed-loss.csv"
        ),
        "metadata": (
            args.output_dir / "cavity-qed-metadata.json"
        ),
    }
    write_csv(paths["spectrum"], spectrum)
    write_csv(paths["vacuum"], vacuum)
    write_csv(paths["collapse"], collapse)
    write_csv(paths["truncation"], truncation)
    write_csv(paths["loss"], loss)

    metadata = {
        "scope": {
            "closed_model": "Jaynes-Cummings Hamiltonian",
            "basis_order": (
                "atom [excited, ground] tensor photon [0,...,N-1]"
            ),
            "detuning": "Delta = omega_a - omega_c",
            "units": {
                "hbar": 1.0,
                "closed_dynamics_g": G_COUPLING,
                "time": "1/g",
            },
            "claim": (
                "finite-basis Jaynes-Cummings and small Lindblad "
                "benchmark, not a device-specific cavity-QED prediction"
            ),
        },
        "dressed_spectrum": {
            "g": 0.1,
            "omega_c": 1.0,
            "Delta_over_g_range": [-6.0, 6.0],
            "manifold_n": [0, 1],
            "samples": args.spectrum_samples,
            "validation": spectrum_validation,
        },
        "vacuum_rabi": {
            "g": G_COUPLING,
            "photon_cutoff": 4,
            "initial_state": "|e,0>",
            "g_t_range": [0.0, 4.0 * math.pi],
            "samples": args.vacuum_samples,
            "validation": vacuum_validation,
        },
        "collapse_revival": {
            "g": G_COUPLING,
            "initial_atom": "excited",
            "initial_field": "coherent state with real amplitude",
            "mean_photon_number": COHERENT_NBAR,
            "photon_cutoffs": list(COHERENT_CUTOFFS),
            "reference_cutoff": COHERENT_REFERENCE_CUTOFF,
            "samples": args.collapse_samples,
            "validation": collapse_validation,
        },
        "loss_extension": {
            "model": (
                "unconditional Lindblad master equation in the "
                "zero- and one-excitation sectors"
            ),
            "g": G_COUPLING,
            "photon_cutoff": 2,
            "cases": [
                {
                    "name": name,
                    "kappa_over_g": kappa,
                    "gamma_over_g": gamma,
                }
                for name, kappa, gamma in LOSS_CASES
            ],
            "kappa_convention": "cavity energy-decay rate",
            "gamma_convention": "emitter population-decay rate",
            "integrator": "classical explicit Runge-Kutta order 4",
            "max_internal_step": args.loss_max_step,
            "samples": args.loss_samples,
            "validation": loss_validation,
        },
        "validation": {
            "all_checks_passed": all(
                all(section["checks"].values())
                for section in (
                    spectrum_validation,
                    vacuum_validation,
                    collapse_validation,
                    loss_validation,
                )
            ),
            "sections": {
                "dressed_spectrum": spectrum_validation["checks"],
                "vacuum_rabi": vacuum_validation["checks"],
                "collapse_revival": collapse_validation["checks"],
                "loss_extension": loss_validation["checks"],
            },
        },
        "limitations": [
            "one two-level emitter and one bosonic mode",
            "rotating-wave Jaynes-Cummings interaction",
            "closed calculations omit all loss and drive",
            "loss extension is restricted to zero and one excitation",
            "Markovian cavity and emitter decay only",
            "no pure dephasing, coherent cavity drive, or thermal photons",
            "no input-output spectrum or detector model",
            "no counter-rotating or diamagnetic terms",
        ],
        "provenance": {
            "Jaynes_Cummings": {
                "doi": "10.1109/PROC.1963.1664",
            },
            "collapse_revival": {
                "doi": "10.1103/PhysRevLett.44.1323",
            },
        },
        "runtime": {
            "python": platform.python_version(),
            "numpy": np.__version__,
            "platform": platform.platform(),
            "random_seed": None,
        },
        "license": "MIT",
        "outputs": [path.name for path in paths.values()],
    }
    if not metadata["validation"]["all_checks_passed"]:
        raise RuntimeError("at least one validation check failed")
    paths["metadata"].write_text(
        json.dumps(
            metadata,
            default=json_scalar,
            indent=2,
            sort_keys=True,
        )
        + "\n",
        encoding="utf-8",
    )

    print(
        "dressed energy error         = "
        f"{spectrum_validation['max_block_energy_error']:.9e}"
    )
    print(
        "vacuum Rabi population error = "
        f"{vacuum_validation['max_excited_population_error']:.9e}"
    )
    print(
        "N=100 versus Poisson sum     = "
        f"{collapse_validation['max_reference_minus_infinite_sum']:.9e}"
    )
    print(
        "revival-window |W|           = "
        f"{collapse_validation['largest_revival_window_abs_inversion']:.9f}"
    )
    print(
        "open-model minimum eigenvalue= "
        f"{loss_validation['minimum_density_eigenvalue']:.9e}"
    )
    print(f"outputs                      = {args.output_dir.resolve()}")
    print("validation                   = all checks passed")


if __name__ == "__main__":
    main()
