#!/usr/bin/env python3
# SPDX-License-Identifier: MIT
"""Verified coherent dynamics for a driven two-level system.

The program computes five related benchmarks:

1. analytic and matrix-propagated Rabi oscillations with detuning;
2. full laboratory-frame dynamics versus the rotating-wave approximation;
3. square and Gaussian pulses with equal nominal pulse area;
4. ideal and finite-pulse Ramsey sequences;
5. convergence and unitarity diagnostics for every numerical propagator.

NumPy is the only dependency. Frequencies are angular frequencies in a
declared dimensionless unit, hbar = 1, and time uses the reciprocal unit.
"""

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, Callable

import numpy as np


IDENTITY = np.eye(2, dtype=complex)
SIGMA_X = np.array(
    [
        [0.0, 1.0],
        [1.0, 0.0],
    ],
    dtype=complex,
)
SIGMA_Y = np.array(
    [
        [0.0, -1.0j],
        [1.0j, 0.0],
    ],
    dtype=complex,
)
SIGMA_Z = np.array(
    [
        [1.0, 0.0],
        [0.0, -1.0],
    ],
    dtype=complex,
)

EXCITED_STATE = np.array([1.0, 0.0], dtype=complex)
GROUND_STATE = np.array([0.0, 1.0], dtype=complex)

RABI_DETUNING_RATIOS = (0.0, 0.5, 1.0)
LAB_DRIVE_RATIOS = (0.02, 0.05, 0.10, 0.20, 0.30)

PULSE_DURATION = 1.0
GAUSSIAN_SIGMA = 0.15
PULSE_TEST_DETUNING = 2.0

RAMSEY_PULSE_OMEGA = 4.0
RAMSEY_FREE_TIME = 10.0
RAMSEY_DETUNING_MIN = -1.0
RAMSEY_DETUNING_MAX = 1.0


@dataclass(frozen=True)
class LabTrace:
    times: np.ndarray
    populations: np.ndarray
    max_norm_error: float


def pauli_propagator(
    h_x: float,
    h_y: float,
    h_z: float,
    duration: float,
) -> np.ndarray:
    """Return exp[-i (hx sx + hy sy + hz sz) duration]."""
    magnitude = math.sqrt(h_x * h_x + h_y * h_y + h_z * h_z)
    if magnitude == 0.0:
        return IDENTITY.copy()
    hamiltonian = (
        h_x * SIGMA_X
        + h_y * SIGMA_Y
        + h_z * SIGMA_Z
    )
    angle = magnitude * duration
    return (
        math.cos(angle) * IDENTITY
        - 1.0j
        * math.sin(angle)
        / magnitude
        * hamiltonian
    )


def excited_population(state: np.ndarray) -> float:
    return float(abs(np.vdot(EXCITED_STATE, state)) ** 2)


def state_norm_error(state: np.ndarray) -> float:
    return abs(float(np.vdot(state, state).real) - 1.0)


def rwa_propagator(
    omega: float,
    delta: float,
    phase: float,
    duration: float,
) -> np.ndarray:
    """Return a constant-pulse RWA propagator.

    The convention is H/hbar = (delta sz + omega sigma_phase)/2, with
    delta = omega_0 - omega_L and
    sigma_phase = cos(phase) sx + sin(phase) sy.
    """
    return pauli_propagator(
        0.5 * omega * math.cos(phase),
        0.5 * omega * math.sin(phase),
        0.5 * delta,
        duration,
    )


def analytic_rabi_probability(
    time: float | np.ndarray,
    omega: float,
    delta: float,
) -> float | np.ndarray:
    generalized = math.sqrt(omega * omega + delta * delta)
    if generalized == 0.0:
        if isinstance(time, np.ndarray):
            return np.zeros_like(time)
        return 0.0
    return (
        omega * omega
        / (generalized * generalized)
        * np.sin(0.5 * generalized * time) ** 2
    )


def rabi_rows(
    samples: int = 401,
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
    omega = 1.0
    tau = np.linspace(0.0, 4.0 * math.pi, samples)
    rows: list[dict[str, Any]] = []
    matrix_error = 0.0
    norm_error = 0.0

    for value in tau:
        row: dict[str, Any] = {
            "tau": value,
            "tau_over_pi": value / math.pi,
        }
        for ratio in RABI_DETUNING_RATIOS:
            delta = ratio * omega
            label = str(ratio).replace(".", "p")
            analytic = float(
                analytic_rabi_probability(value / omega, omega, delta)
            )
            state = (
                rwa_propagator(
                    omega,
                    delta,
                    0.0,
                    value / omega,
                )
                @ GROUND_STATE
            )
            numerical = excited_population(state)
            row[f"P_delta_{label}_analytic"] = analytic
            row[f"P_delta_{label}_matrix"] = numerical
            matrix_error = max(
                matrix_error,
                abs(numerical - analytic),
            )
            norm_error = max(norm_error, state_norm_error(state))
        rows.append(row)

    validation = {
        "max_matrix_minus_analytic_probability": matrix_error,
        "max_state_norm_error": norm_error,
        "checks": {
            "analytic_population": matrix_error < 1.0e-13,
            "unit_norm": norm_error < 1.0e-13,
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(f"Rabi validation failed: {validation}")
    return rows, validation


def laboratory_trace(
    sample_times: np.ndarray,
    omega_0: float,
    omega_laser: float,
    omega_drive: float,
    phase: float,
    steps_per_carrier_cycle: int,
) -> LabTrace:
    """Propagate the full linearly driven laboratory-frame Hamiltonian.

    H/hbar = omega_0 sz/2 + omega_drive cos(omega_laser t + phase) sx.
    Each midpoint step exponentiates the frozen 2 by 2 Hamiltonian exactly.
    """
    if steps_per_carrier_cycle < 4:
        raise ValueError("steps-per-carrier-cycle must be at least 4")
    if omega_laser <= 0.0:
        raise ValueError("omega_laser must be positive")
    if np.any(np.diff(sample_times) < 0.0):
        raise ValueError("sample times must be nondecreasing")

    max_step = (
        2.0 * math.pi / omega_laser / steps_per_carrier_cycle
    )
    state = GROUND_STATE.copy()
    populations = [excited_population(state)]
    max_norm_error = state_norm_error(state)
    previous_time = float(sample_times[0])

    if previous_time != 0.0:
        raise ValueError("the first sample time must be zero")

    for sample_time_raw in sample_times[1:]:
        sample_time = float(sample_time_raw)
        interval = sample_time - previous_time
        steps = max(1, math.ceil(interval / max_step))
        step = interval / steps
        for index in range(steps):
            midpoint = previous_time + (index + 0.5) * step
            drive = omega_drive * math.cos(
                omega_laser * midpoint + phase
            )
            state = (
                pauli_propagator(
                    drive,
                    0.0,
                    0.5 * omega_0,
                    step,
                )
                @ state
            )
        populations.append(excited_population(state))
        max_norm_error = max(
            max_norm_error,
            state_norm_error(state),
        )
        previous_time = sample_time

    return LabTrace(
        times=np.asarray(sample_times, dtype=float),
        populations=np.asarray(populations, dtype=float),
        max_norm_error=max_norm_error,
    )


def drive_ratio_label(value: float) -> str:
    return f"{value:.2f}".replace(".", "p")


def rwa_comparison(
    *,
    samples: int,
    fine_steps_per_cycle: int,
) -> tuple[
    list[dict[str, Any]],
    list[dict[str, Any]],
    dict[str, Any],
]:
    omega_0 = 1.0
    omega_laser = 1.0
    tau = np.linspace(0.0, 2.0 * math.pi, samples)
    rwa_population = np.sin(0.5 * tau) ** 2
    fine_traces: dict[float, LabTrace] = {}
    convergence_rows: list[dict[str, Any]] = []

    medium_steps = fine_steps_per_cycle // 2
    coarse_steps = fine_steps_per_cycle // 4
    if coarse_steps < 4:
        raise ValueError("fine steps per cycle must be at least 16")

    for ratio in LAB_DRIVE_RATIOS:
        omega_drive = ratio * omega_0
        times = tau / omega_drive
        coarse = laboratory_trace(
            times,
            omega_0,
            omega_laser,
            omega_drive,
            0.0,
            coarse_steps,
        )
        medium = laboratory_trace(
            times,
            omega_0,
            omega_laser,
            omega_drive,
            0.0,
            medium_steps,
        )
        fine = laboratory_trace(
            times,
            omega_0,
            omega_laser,
            omega_drive,
            0.0,
            fine_steps_per_cycle,
        )
        fine_traces[ratio] = fine
        pi_index = int(np.argmin(np.abs(tau - math.pi)))
        coarse_medium = float(
            np.max(
                np.abs(
                    coarse.populations - medium.populations
                )
            )
        )
        medium_fine = float(
            np.max(
                np.abs(
                    medium.populations - fine.populations
                )
            )
        )
        convergence_rows.append(
            {
                "Omega_over_omega0": ratio,
                "coarse_steps_per_carrier_cycle": coarse_steps,
                "medium_steps_per_carrier_cycle": medium_steps,
                "fine_steps_per_carrier_cycle": (
                    fine_steps_per_cycle
                ),
                "max_abs_coarse_minus_medium": coarse_medium,
                "max_abs_medium_minus_fine": medium_fine,
                "max_abs_lab_minus_RWA": float(
                    np.max(
                        np.abs(
                            fine.populations - rwa_population
                        )
                    )
                ),
                "lab_pi_pulse_population": float(
                    fine.populations[pi_index]
                ),
                "RWA_pi_pulse_population": float(
                    rwa_population[pi_index]
                ),
                "fine_max_norm_error": fine.max_norm_error,
            }
        )

    rows: list[dict[str, Any]] = []
    for index, value in enumerate(tau):
        row: dict[str, Any] = {
            "tau": value,
            "tau_over_pi": value / math.pi,
            "P_RWA": rwa_population[index],
        }
        for ratio, trace in fine_traces.items():
            row[
                f"P_lab_Omega_over_omega0_{drive_ratio_label(ratio)}"
            ] = trace.populations[index]
        rows.append(row)

    rwa_errors = [
        row["max_abs_lab_minus_RWA"]
        for row in convergence_rows
    ]
    numerical_errors = [
        row["max_abs_medium_minus_fine"]
        for row in convergence_rows
    ]
    norm_errors = [
        row["fine_max_norm_error"]
        for row in convergence_rows
    ]
    validation = {
        "largest_medium_fine_population_difference": max(
            numerical_errors
        ),
        "largest_fine_norm_error": max(norm_errors),
        "RWA_discrepancy_at_weakest_drive": rwa_errors[0],
        "RWA_discrepancy_at_strongest_drive": rwa_errors[-1],
        "weak_drive_discrepancy_to_numerical_error_ratio": (
            rwa_errors[0] / numerical_errors[0]
        ),
        "checks": {
            "midpoint_converged": max(numerical_errors) < 6.0e-6,
            "unitary_steps": max(norm_errors) < 1.0e-11,
            "RWA_error_grows_over_test_range": all(
                upper > lower
                for lower, upper in zip(
                    rwa_errors,
                    rwa_errors[1:],
                )
            ),
            "physical_error_resolved": (
                rwa_errors[0] > 500.0 * numerical_errors[0]
            ),
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(
            f"laboratory-frame validation failed: {validation}"
        )
    return rows, convergence_rows, validation


def gaussian_normalization(
    sigma: float = GAUSSIAN_SIGMA,
    duration: float = PULSE_DURATION,
) -> float:
    half_width = 0.5 * duration
    return (
        sigma
        * math.sqrt(2.0 * math.pi)
        * math.erf(
            half_width / (math.sqrt(2.0) * sigma)
        )
    )


def gaussian_envelope(
    time: float,
    area: float,
    sigma: float = GAUSSIAN_SIGMA,
    duration: float = PULSE_DURATION,
) -> float:
    centre = 0.5 * duration
    normalization = gaussian_normalization(sigma, duration)
    return (
        area
        / normalization
        * math.exp(
            -0.5 * ((time - centre) / sigma) ** 2
        )
    )


def propagate_rwa_envelope(
    envelope: Callable[[float], float],
    *,
    delta: float,
    phase: float,
    duration: float,
    steps: int,
) -> tuple[np.ndarray, float]:
    """Propagate a time-dependent RWA envelope by midpoint exponentials."""
    if steps < 1:
        raise ValueError("steps must be positive")
    step = duration / steps
    state = GROUND_STATE.copy()
    max_norm_error = state_norm_error(state)
    for index in range(steps):
        midpoint = (index + 0.5) * step
        state = (
            rwa_propagator(
                envelope(midpoint),
                delta,
                phase,
                step,
            )
            @ state
        )
        max_norm_error = max(
            max_norm_error,
            state_norm_error(state),
        )
    return state, max_norm_error


def pulse_area_rows(
    *,
    samples: int,
    gaussian_steps: int,
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
    areas = np.linspace(0.0, 4.0 * math.pi, samples)
    rows: list[dict[str, Any]] = []
    resonant_square_error = 0.0
    resonant_gaussian_error = 0.0
    max_norm_error = 0.0
    off_resonant_shape_difference = 0.0

    for area in areas:
        area_law = math.sin(0.5 * area) ** 2

        square_resonant_state = (
            rwa_propagator(
                area / PULSE_DURATION,
                0.0,
                0.0,
                PULSE_DURATION,
            )
            @ GROUND_STATE
        )
        square_detuned_state = (
            rwa_propagator(
                area / PULSE_DURATION,
                PULSE_TEST_DETUNING,
                0.0,
                PULSE_DURATION,
            )
            @ GROUND_STATE
        )
        gaussian_resonant_state, gaussian_resonant_norm = (
            propagate_rwa_envelope(
                lambda time, selected_area=area: gaussian_envelope(
                    time,
                    selected_area,
                ),
                delta=0.0,
                phase=0.0,
                duration=PULSE_DURATION,
                steps=gaussian_steps,
            )
        )
        gaussian_detuned_state, gaussian_detuned_norm = (
            propagate_rwa_envelope(
                lambda time, selected_area=area: gaussian_envelope(
                    time,
                    selected_area,
                ),
                delta=PULSE_TEST_DETUNING,
                phase=0.0,
                duration=PULSE_DURATION,
                steps=gaussian_steps,
            )
        )

        square_resonant = excited_population(
            square_resonant_state
        )
        gaussian_resonant = excited_population(
            gaussian_resonant_state
        )
        square_detuned = excited_population(
            square_detuned_state
        )
        gaussian_detuned = excited_population(
            gaussian_detuned_state
        )

        resonant_square_error = max(
            resonant_square_error,
            abs(square_resonant - area_law),
        )
        resonant_gaussian_error = max(
            resonant_gaussian_error,
            abs(gaussian_resonant - area_law),
        )
        max_norm_error = max(
            max_norm_error,
            state_norm_error(square_resonant_state),
            state_norm_error(square_detuned_state),
            gaussian_resonant_norm,
            gaussian_detuned_norm,
        )
        off_resonant_shape_difference = max(
            off_resonant_shape_difference,
            abs(square_detuned - gaussian_detuned),
        )
        rows.append(
            {
                "area": area,
                "area_over_pi": area / math.pi,
                "P_area_law": area_law,
                "P_resonant_square": square_resonant,
                "P_resonant_gaussian": gaussian_resonant,
                "P_detuned_square": square_detuned,
                "P_detuned_gaussian": gaussian_detuned,
            }
        )

    validation = {
        "max_resonant_square_minus_area_law": (
            resonant_square_error
        ),
        "max_resonant_gaussian_minus_area_law": (
            resonant_gaussian_error
        ),
        "max_off_resonant_square_minus_gaussian": (
            off_resonant_shape_difference
        ),
        "max_state_norm_error": max_norm_error,
        "checks": {
            "square_area_theorem": resonant_square_error < 1.0e-13,
            "gaussian_area_theorem": (
                resonant_gaussian_error < 1.0e-6
            ),
            "off_resonant_shape_dependence_resolved": (
                off_resonant_shape_difference > 0.1
            ),
            "unit_norm": max_norm_error < 1.0e-11,
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(
            f"pulse-area validation failed: {validation}"
        )
    return rows, validation


def ramsey_state(
    delta: float,
    *,
    phase_1: float,
    phase_2: float,
    finite_pulses: bool,
) -> np.ndarray:
    pulse_duration = math.pi / (2.0 * RAMSEY_PULSE_OMEGA)
    pulse_delta = delta if finite_pulses else 0.0
    first = rwa_propagator(
        RAMSEY_PULSE_OMEGA,
        pulse_delta,
        phase_1,
        pulse_duration,
    )
    free = rwa_propagator(
        0.0,
        delta,
        0.0,
        RAMSEY_FREE_TIME,
    )
    second = rwa_propagator(
        RAMSEY_PULSE_OMEGA,
        pulse_delta,
        phase_2,
        pulse_duration,
    )
    return second @ free @ first @ GROUND_STATE


def ideal_ramsey_probability(
    delta: float,
    phase_1: float,
    phase_2: float,
) -> float:
    phase = (
        delta * RAMSEY_FREE_TIME
        + phase_1
        - phase_2
    )
    return 0.5 * (1.0 + math.cos(phase))


def ramsey_rows(
    samples: int = 401,
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
    detunings = np.linspace(
        RAMSEY_DETUNING_MIN,
        RAMSEY_DETUNING_MAX,
        samples,
    )
    rows: list[dict[str, Any]] = []
    ideal_matrix_error = 0.0
    max_norm_error = 0.0

    for delta in detunings:
        ideal_formula = ideal_ramsey_probability(
            float(delta),
            0.0,
            0.0,
        )
        ideal_matrix_state = ramsey_state(
            float(delta),
            phase_1=0.0,
            phase_2=0.0,
            finite_pulses=False,
        )
        finite_zero_state = ramsey_state(
            float(delta),
            phase_1=0.0,
            phase_2=0.0,
            finite_pulses=True,
        )
        finite_plus_state = ramsey_state(
            float(delta),
            phase_1=0.0,
            phase_2=0.5 * math.pi,
            finite_pulses=True,
        )
        finite_minus_state = ramsey_state(
            float(delta),
            phase_1=0.0,
            phase_2=-0.5 * math.pi,
            finite_pulses=True,
        )

        ideal_matrix = excited_population(ideal_matrix_state)
        finite_zero = excited_population(finite_zero_state)
        finite_plus = excited_population(finite_plus_state)
        finite_minus = excited_population(finite_minus_state)
        ideal_matrix_error = max(
            ideal_matrix_error,
            abs(ideal_matrix - ideal_formula),
        )
        max_norm_error = max(
            max_norm_error,
            state_norm_error(ideal_matrix_state),
            state_norm_error(finite_zero_state),
            state_norm_error(finite_plus_state),
            state_norm_error(finite_minus_state),
        )
        rows.append(
            {
                "delta": delta,
                "delta_over_pulse_Omega": (
                    delta / RAMSEY_PULSE_OMEGA
                ),
                "delta_T_over_pi": (
                    delta * RAMSEY_FREE_TIME / math.pi
                ),
                "P_ideal_formula": ideal_formula,
                "P_ideal_matrix": ideal_matrix,
                "P_finite_zero_phase": finite_zero,
                "P_finite_phase_plus_pi_over_2": finite_plus,
                "P_finite_phase_minus_pi_over_2": finite_minus,
                "finite_error_signal": finite_plus - finite_minus,
            }
        )

    centre = samples // 2
    delta_step = float(detunings[centre + 1] - detunings[centre])
    finite_slope = (
        rows[centre + 1]["finite_error_signal"]
        - rows[centre - 1]["finite_error_signal"]
    ) / (2.0 * delta_step)
    max_finite_ideal_difference = max(
        abs(
            row["P_finite_zero_phase"]
            - row["P_ideal_formula"]
        )
        for row in rows
    )
    validation = {
        "max_ideal_matrix_minus_formula": ideal_matrix_error,
        "on_resonance_finite_population": rows[centre][
            "P_finite_zero_phase"
        ],
        "on_resonance_finite_error_signal": rows[centre][
            "finite_error_signal"
        ],
        "finite_error_signal_slope_at_resonance": finite_slope,
        "max_finite_minus_instantaneous_probability": (
            max_finite_ideal_difference
        ),
        "max_state_norm_error": max_norm_error,
        "checks": {
            "ideal_formula": ideal_matrix_error < 1.0e-13,
            "resonant_sequence": (
                abs(
                    rows[centre]["P_finite_zero_phase"] - 1.0
                )
                < 1.0e-13
            ),
            "zero_crossing_error_signal": (
                abs(rows[centre]["finite_error_signal"]) < 1.0e-13
            ),
            "unit_norm": max_norm_error < 1.0e-13,
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(
            f"Ramsey 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 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("--rabi-samples", type=int, default=401)
    parser.add_argument("--lab-samples", type=int, default=401)
    parser.add_argument(
        "--lab-steps-per-cycle",
        type=int,
        default=1600,
    )
    parser.add_argument("--area-samples", type=int, default=161)
    parser.add_argument("--gaussian-steps", type=int, default=1200)
    parser.add_argument("--ramsey-samples", type=int, default=401)
    parser.add_argument(
        "--output-dir",
        type=Path,
        default=Path("time-dependent-two-level-output"),
    )
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    for name in (
        "rabi_samples",
        "lab_samples",
        "area_samples",
        "ramsey_samples",
    ):
        if getattr(args, name) < 5:
            raise ValueError(f"{name.replace('_', '-')} must be at least 5")
    if args.rabi_samples % 2 == 0:
        raise ValueError("rabi-samples must be odd")
    if args.lab_samples % 2 == 0:
        raise ValueError("lab-samples must be odd")
    if args.ramsey_samples % 2 == 0:
        raise ValueError("ramsey-samples must be odd")

    rabi, rabi_validation = rabi_rows(args.rabi_samples)
    (
        rwa,
        convergence,
        rwa_validation,
    ) = rwa_comparison(
        samples=args.lab_samples,
        fine_steps_per_cycle=args.lab_steps_per_cycle,
    )
    area, area_validation = pulse_area_rows(
        samples=args.area_samples,
        gaussian_steps=args.gaussian_steps,
    )
    ramsey, ramsey_validation = ramsey_rows(
        args.ramsey_samples
    )

    args.output_dir.mkdir(parents=True, exist_ok=True)
    paths = {
        "rabi": (
            args.output_dir
            / "time-dependent-two-level-rabi.csv"
        ),
        "rwa": (
            args.output_dir
            / "time-dependent-two-level-rwa-comparison.csv"
        ),
        "area": (
            args.output_dir
            / "time-dependent-two-level-pulse-area.csv"
        ),
        "ramsey": (
            args.output_dir
            / "time-dependent-two-level-ramsey.csv"
        ),
        "convergence": (
            args.output_dir
            / "time-dependent-two-level-convergence.csv"
        ),
        "metadata": (
            args.output_dir
            / "time-dependent-two-level-metadata.json"
        ),
    }
    write_csv(paths["rabi"], rabi)
    write_csv(paths["rwa"], rwa)
    write_csv(paths["area"], area)
    write_csv(paths["ramsey"], ramsey)
    write_csv(paths["convergence"], convergence)

    metadata = {
        "scope": {
            "system": "generic coherent two-level system",
            "basis": ["excited", "ground"],
            "sigma_z_eigenvalues": [1, -1],
            "detuning_convention": "delta = omega_0 - omega_L",
            "units": {
                "frequency": "declared angular-frequency unit",
                "time": "reciprocal angular-frequency unit",
                "hbar": 1.0,
            },
            "claim": (
                "closed-system propagator and approximation benchmark, "
                "not a species-specific prediction"
            ),
        },
        "rabi": {
            "Omega": 1.0,
            "detuning_over_Omega": list(
                RABI_DETUNING_RATIOS
            ),
            "tau_range": [0.0, 4.0 * math.pi],
            "samples": args.rabi_samples,
            "validation": rabi_validation,
        },
        "laboratory_frame": {
            "hamiltonian": (
                "H/hbar = omega_0 sigma_z/2 + "
                "Omega cos(omega_L t + phase) sigma_x"
            ),
            "omega_0": 1.0,
            "omega_L": 1.0,
            "Omega_over_omega_0": list(LAB_DRIVE_RATIOS),
            "tau_range": [0.0, 2.0 * math.pi],
            "samples": args.lab_samples,
            "integrator": (
                "exponential midpoint with exact frozen 2x2 "
                "Hamiltonian exponential"
            ),
            "integrator_global_order": 2,
            "fine_steps_per_carrier_cycle": (
                args.lab_steps_per_cycle
            ),
            "validation": rwa_validation,
        },
        "pulse_area": {
            "duration": PULSE_DURATION,
            "area_range": [0.0, 4.0 * math.pi],
            "samples": args.area_samples,
            "square_amplitude": "area / duration",
            "gaussian_sigma": GAUSSIAN_SIGMA,
            "gaussian_steps": args.gaussian_steps,
            "off_resonant_test_delta": PULSE_TEST_DETUNING,
            "validation": area_validation,
        },
        "ramsey": {
            "pulse_Omega": RAMSEY_PULSE_OMEGA,
            "pulse_duration": (
                math.pi / (2.0 * RAMSEY_PULSE_OMEGA)
            ),
            "pulse_area": 0.5 * math.pi,
            "free_time": RAMSEY_FREE_TIME,
            "detuning_range": [
                RAMSEY_DETUNING_MIN,
                RAMSEY_DETUNING_MAX,
            ],
            "samples": args.ramsey_samples,
            "validation": ramsey_validation,
        },
        "validation": {
            "all_checks_passed": all(
                all(section["checks"].values())
                for section in (
                    rabi_validation,
                    rwa_validation,
                    area_validation,
                    ramsey_validation,
                )
            ),
            "sections": {
                "rabi": rabi_validation["checks"],
                "rwa": rwa_validation["checks"],
                "pulse_area": area_validation["checks"],
                "ramsey": ramsey_validation["checks"],
            },
        },
        "limitations": [
            "two-state projection is assumed rather than derived",
            "closed-system unitary dynamics only",
            "no relaxation, dephasing, leakage, motion, or ensemble averaging",
            "the full laboratory-frame comparison uses a monochromatic drive",
            "pulse calibration and readout errors are absent",
        ],
        "provenance": {
            "Rabi": {
                "doi": "10.1103/PhysRev.51.652",
            },
            "Bloch_Siegert": {
                "doi": "10.1103/PhysRev.57.522",
            },
            "Ramsey": {
                "doi": "10.1103/PhysRev.78.695",
            },
            "Floquet_two_state": {
                "doi": "10.1103/PhysRev.138.B979",
            },
        },
        "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, indent=2, sort_keys=True) + "\n",
        encoding="utf-8",
    )

    print(
        "weak-drive max RWA error   = "
        f"{rwa_validation['RWA_discrepancy_at_weakest_drive']:.9e}"
    )
    print(
        "strong-drive max RWA error = "
        f"{rwa_validation['RWA_discrepancy_at_strongest_drive']:.9e}"
    )
    print(
        "lab numerical difference  = "
        f"{rwa_validation['largest_medium_fine_population_difference']:.9e}"
    )
    print(
        "Gaussian area-law error    = "
        f"{area_validation['max_resonant_gaussian_minus_area_law']:.9e}"
    )
    print(
        "off-resonant shape effect  = "
        f"{area_validation['max_off_resonant_square_minus_gaussian']:.9e}"
    )
    print(
        "finite Ramsey deviation    = "
        f"{ramsey_validation['max_finite_minus_instantaneous_probability']:.9e}"
    )
    print(f"outputs                    = {args.output_dir.resolve()}")
    print("validation                 = all checks passed")


if __name__ == "__main__":
    main()
