#!/usr/bin/env python3
# SPDX-License-Identifier: MIT
"""Verified optical Bloch dynamics for a driven dissipative two-level atom.

The program computes five related benchmarks:

1. transient optical Bloch trajectories from the ground state;
2. no-drive population and coherence decay checks;
3. analytic and linear-solve steady-state saturation curves;
4. power-broadened line shapes and numerically extracted linewidths;
5. RK4 convergence, positivity, and steady-state diagnostics.

NumPy is the only dependency. Gamma = 1 sets the angular-frequency unit,
time is measured in 1/Gamma, hbar = 1, and the detuning convention is
Delta = omega_0 - omega_L.
"""

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


GAMMA = 1.0
TRANSIENT_TIME_MAX = 15.0
TRANSIENT_CASES = (
    ("resonant_weak", 0.25, 0.0, 0.0),
    ("resonant_saturated", 1.0, 0.0, 0.0),
    ("resonant_oscillatory", 3.0, 0.0, 0.0),
    ("detuned", 1.0, 1.5, 0.0),
    ("dephased", 1.0, 0.5, 0.75),
)
SATURATION_DEPHASING_RATIOS = (0.0, 0.5, 2.0)
LINE_SHAPE_DRIVE_RATIOS = (0.1, 0.5, 1.0, 2.0)


@dataclass(frozen=True)
class BlochTrace:
    times: np.ndarray
    states: np.ndarray
    min_population: float
    max_population: float
    max_bloch_radius_excess: float


def gamma_2(gamma: float, gamma_phi: float) -> float:
    return 0.5 * gamma + gamma_phi


def bloch_matrix(
    omega: float,
    delta: float,
    gamma: float,
    gamma_phi: float,
) -> tuple[np.ndarray, np.ndarray]:
    """Return A and b for dr/dt = A r + b, r = (u, v, w)."""
    transverse = gamma_2(gamma, gamma_phi)
    matrix = np.array(
        [
            [-transverse, -delta, 0.0],
            [delta, -transverse, -omega],
            [0.0, omega, -gamma],
        ],
        dtype=float,
    )
    source = np.array([0.0, 0.0, -gamma], dtype=float)
    return matrix, source


def bloch_rhs(
    state: np.ndarray,
    omega: float,
    delta: float,
    gamma: float,
    gamma_phi: float,
) -> np.ndarray:
    matrix, source = bloch_matrix(
        omega,
        delta,
        gamma,
        gamma_phi,
    )
    return matrix @ state + source


def rk4_step(
    state: np.ndarray,
    step: float,
    omega: float,
    delta: float,
    gamma: float,
    gamma_phi: float,
) -> np.ndarray:
    k1 = bloch_rhs(state, omega, delta, gamma, gamma_phi)
    k2 = bloch_rhs(
        state + 0.5 * step * k1,
        omega,
        delta,
        gamma,
        gamma_phi,
    )
    k3 = bloch_rhs(
        state + 0.5 * step * k2,
        omega,
        delta,
        gamma,
        gamma_phi,
    )
    k4 = bloch_rhs(
        state + step * k3,
        omega,
        delta,
        gamma,
        gamma_phi,
    )
    return state + step * (k1 + 2.0 * k2 + 2.0 * k3 + k4) / 6.0


def excited_population(state: np.ndarray) -> float:
    return 0.5 * (1.0 + float(state[2]))


def propagate_bloch(
    sample_times: np.ndarray,
    *,
    initial_state: np.ndarray,
    omega: float,
    delta: float,
    gamma: float,
    gamma_phi: float,
    max_step: float,
) -> BlochTrace:
    if sample_times.ndim != 1 or len(sample_times) < 2:
        raise ValueError("sample_times must be a one-dimensional grid")
    if float(sample_times[0]) != 0.0:
        raise ValueError("the first sample time must be zero")
    if np.any(np.diff(sample_times) <= 0.0):
        raise ValueError("sample times must be strictly increasing")
    if max_step <= 0.0:
        raise ValueError("max_step must be positive")
    if gamma <= 0.0 or gamma_phi < 0.0:
        raise ValueError("decay rates must be physical")

    state = np.asarray(initial_state, dtype=float).copy()
    if state.shape != (3,):
        raise ValueError("initial_state must contain (u, v, w)")
    states = [state.copy()]
    populations = [excited_population(state)]
    max_radius_excess = max(0.0, float(np.linalg.norm(state)) - 1.0)
    previous_time = 0.0

    for target_raw in sample_times[1:]:
        target = float(target_raw)
        interval = target - previous_time
        steps = max(1, math.ceil(interval / max_step))
        step = interval / steps
        for _ in range(steps):
            state = rk4_step(
                state,
                step,
                omega,
                delta,
                gamma,
                gamma_phi,
            )
        states.append(state.copy())
        populations.append(excited_population(state))
        max_radius_excess = max(
            max_radius_excess,
            float(np.linalg.norm(state)) - 1.0,
        )
        previous_time = target

    population_array = np.asarray(populations, dtype=float)
    return BlochTrace(
        times=np.asarray(sample_times, dtype=float),
        states=np.asarray(states, dtype=float),
        min_population=float(np.min(population_array)),
        max_population=float(np.max(population_array)),
        max_bloch_radius_excess=max(0.0, max_radius_excess),
    )


def steady_state_analytic(
    omega: float,
    delta: float,
    gamma: float,
    gamma_phi: float,
) -> np.ndarray:
    transverse = gamma_2(gamma, gamma_phi)
    denominator = (
        gamma * (delta * delta + transverse * transverse)
        + omega * omega * transverse
    )
    u = -omega * delta * gamma / denominator
    v = omega * transverse * gamma / denominator
    w = (
        -gamma
        * (delta * delta + transverse * transverse)
        / denominator
    )
    return np.array([u, v, w], dtype=float)


def steady_state_linear_solve(
    omega: float,
    delta: float,
    gamma: float,
    gamma_phi: float,
) -> np.ndarray:
    matrix, source = bloch_matrix(
        omega,
        delta,
        gamma,
        gamma_phi,
    )
    return np.linalg.solve(matrix, -source)


def saturation_parameter(
    omega: float,
    delta: float,
    gamma: float,
    gamma_phi: float,
) -> float:
    transverse = gamma_2(gamma, gamma_phi)
    return (
        omega
        * omega
        * transverse
        / (
            gamma
            * (delta * delta + transverse * transverse)
        )
    )


def steady_population_formula(
    omega: float,
    delta: float,
    gamma: float,
    gamma_phi: float,
) -> float:
    saturation = saturation_parameter(
        omega,
        delta,
        gamma,
        gamma_phi,
    )
    return saturation / (2.0 * (1.0 + saturation))


def analytic_fwhm(
    omega: float,
    gamma: float,
    gamma_phi: float,
) -> float:
    transverse = gamma_2(gamma, gamma_phi)
    return 2.0 * math.sqrt(
        transverse * transverse
        + omega * omega * transverse / gamma
    )


def numerical_fwhm(
    omega: float,
    gamma: float,
    gamma_phi: float,
    *,
    tolerance: float = 1.0e-13,
) -> float:
    peak = excited_population(
        steady_state_linear_solve(
            omega,
            0.0,
            gamma,
            gamma_phi,
        )
    )
    target = 0.5 * peak
    lower = 0.0
    upper = max(gamma, omega, gamma_2(gamma, gamma_phi))

    def residual(delta: float) -> float:
        state = steady_state_linear_solve(
            omega,
            delta,
            gamma,
            gamma_phi,
        )
        return excited_population(state) - target

    while residual(upper) > 0.0:
        upper *= 2.0
        if upper > 1.0e8:
            raise RuntimeError("failed to bracket linewidth")

    while upper - lower > tolerance * max(1.0, upper):
        midpoint = 0.5 * (lower + upper)
        if residual(midpoint) > 0.0:
            lower = midpoint
        else:
            upper = midpoint
    return lower + upper


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


def transient_rows(
    *,
    samples: int,
    max_step: float,
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
    times = np.linspace(0.0, TRANSIENT_TIME_MAX, samples)
    traces: dict[str, BlochTrace] = {}
    validation: dict[str, Any] = {}
    minimum_population = 1.0
    maximum_population = 0.0
    radius_excess = 0.0
    steady_residual = 0.0

    for name, omega, delta, gamma_phi in TRANSIENT_CASES:
        trace = propagate_bloch(
            times,
            initial_state=np.array([0.0, 0.0, -1.0]),
            omega=omega,
            delta=delta,
            gamma=GAMMA,
            gamma_phi=gamma_phi,
            max_step=max_step,
        )
        traces[name] = trace
        minimum_population = min(
            minimum_population,
            trace.min_population,
        )
        maximum_population = max(
            maximum_population,
            trace.max_population,
        )
        radius_excess = max(
            radius_excess,
            trace.max_bloch_radius_excess,
        )
        target = steady_state_analytic(
            omega,
            delta,
            GAMMA,
            gamma_phi,
        )
        steady_residual = max(
            steady_residual,
            float(np.max(np.abs(trace.states[-1] - target))),
        )

    rows: list[dict[str, Any]] = []
    for index, time in enumerate(times):
        row: dict[str, Any] = {
            "Gamma_t": time,
        }
        for name, _, _, _ in TRANSIENT_CASES:
            state = traces[name].states[index]
            row[f"u_{name}"] = state[0]
            row[f"v_{name}"] = state[1]
            row[f"w_{name}"] = state[2]
            row[f"rhoee_{name}"] = excited_population(state)
            row[f"Rfl_over_Gamma_{name}"] = excited_population(
                state
            )
        rows.append(row)

    validation = {
        "minimum_excited_population": minimum_population,
        "maximum_excited_population": maximum_population,
        "max_bloch_radius_excess": radius_excess,
        "max_final_state_minus_steady_state": steady_residual,
        "checks": {
            "physical_population": (
                minimum_population > -1.0e-11
                and maximum_population < 1.0 + 1.0e-11
            ),
            "bloch_ball": radius_excess < 1.0e-10,
            "steady_state_reached": steady_residual < 6.0e-4,
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(
            f"transient validation failed: {validation}"
        )
    return rows, validation


def decay_validation(
    *,
    samples: int,
    max_step: float,
) -> dict[str, Any]:
    times = np.linspace(0.0, 10.0, samples)
    excited = propagate_bloch(
        times,
        initial_state=np.array([0.0, 0.0, 1.0]),
        omega=0.0,
        delta=0.0,
        gamma=GAMMA,
        gamma_phi=0.0,
        max_step=max_step,
    )
    coherence = propagate_bloch(
        times,
        initial_state=np.array([1.0, 0.0, 0.0]),
        omega=0.0,
        delta=0.0,
        gamma=GAMMA,
        gamma_phi=0.75,
        max_step=max_step,
    )
    population_numeric = 0.5 * (1.0 + excited.states[:, 2])
    population_exact = np.exp(-GAMMA * times)
    transverse = gamma_2(GAMMA, 0.75)
    coherence_exact = np.exp(-transverse * times)
    population_error = float(
        np.max(np.abs(population_numeric - population_exact))
    )
    coherence_error = float(
        np.max(
            np.abs(coherence.states[:, 0] - coherence_exact)
        )
    )
    validation = {
        "max_population_decay_error": population_error,
        "max_coherence_decay_error": coherence_error,
        "checks": {
            "population_decay": population_error < 1.0e-11,
            "coherence_decay": coherence_error < 1.0e-11,
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(
            f"decay validation failed: {validation}"
        )
    return validation


def saturation_rows(
    samples: int,
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
    omegas = np.linspace(0.0, 5.0, samples)
    rows: list[dict[str, Any]] = []
    state_error = 0.0
    population_error = 0.0
    high_drive_gap = 0.0

    for omega in omegas:
        row: dict[str, Any] = {
            "Omega_over_Gamma": omega,
        }
        for gamma_phi in SATURATION_DEPHASING_RATIOS:
            suffix = label(gamma_phi)
            analytic_state = steady_state_analytic(
                float(omega),
                0.0,
                GAMMA,
                gamma_phi,
            )
            linear_state = steady_state_linear_solve(
                float(omega),
                0.0,
                GAMMA,
                gamma_phi,
            )
            formula_population = steady_population_formula(
                float(omega),
                0.0,
                GAMMA,
                gamma_phi,
            )
            analytic_population = excited_population(
                analytic_state
            )
            linear_population = excited_population(linear_state)
            state_error = max(
                state_error,
                float(
                    np.max(
                        np.abs(analytic_state - linear_state)
                    )
                ),
            )
            population_error = max(
                population_error,
                abs(analytic_population - formula_population),
                abs(linear_population - formula_population),
            )
            row[f"s0_gamma_phi_{suffix}"] = saturation_parameter(
                float(omega),
                0.0,
                GAMMA,
                gamma_phi,
            )
            row[
                f"rhoee_formula_gamma_phi_{suffix}"
            ] = formula_population
            row[
                f"rhoee_linear_gamma_phi_{suffix}"
            ] = linear_population
            row[
                f"Rfl_over_Gamma_gamma_phi_{suffix}"
            ] = linear_population
        rows.append(row)

    for gamma_phi in SATURATION_DEPHASING_RATIOS:
        suffix = label(gamma_phi)
        high_drive_gap = max(
            high_drive_gap,
            abs(
                rows[-1][
                    f"rhoee_linear_gamma_phi_{suffix}"
                ]
                - 0.5
            ),
        )

    validation = {
        "max_analytic_minus_linear_state": state_error,
        "max_population_formula_error": population_error,
        "largest_gap_from_half_at_Omega_over_Gamma_5": (
            high_drive_gap
        ),
        "checks": {
            "steady_state_linear_solve": state_error < 1.0e-13,
            "saturation_formula": population_error < 1.0e-13,
            "high_drive_tends_to_half": high_drive_gap < 0.05,
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(
            f"saturation validation failed: {validation}"
        )
    return rows, validation


def line_shape_rows(
    samples: int,
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
    detunings = np.linspace(-6.0, 6.0, samples)
    rows: list[dict[str, Any]] = []
    population_error = 0.0
    symmetry_error = 0.0

    peaks = {
        omega: steady_population_formula(
            omega,
            0.0,
            GAMMA,
            0.0,
        )
        for omega in LINE_SHAPE_DRIVE_RATIOS
    }
    for delta in detunings:
        row: dict[str, Any] = {
            "Delta_over_Gamma": delta,
        }
        for omega in LINE_SHAPE_DRIVE_RATIOS:
            suffix = label(omega)
            formula = steady_population_formula(
                omega,
                float(delta),
                GAMMA,
                0.0,
            )
            linear = excited_population(
                steady_state_linear_solve(
                    omega,
                    float(delta),
                    GAMMA,
                    0.0,
                )
            )
            population_error = max(
                population_error,
                abs(formula - linear),
            )
            row[f"rhoee_formula_Omega_{suffix}"] = formula
            row[f"rhoee_linear_Omega_{suffix}"] = linear
            row[f"normalized_Omega_{suffix}"] = (
                linear / peaks[omega]
            )
        rows.append(row)

    for lower, upper in zip(rows, reversed(rows), strict=True):
        for omega in LINE_SHAPE_DRIVE_RATIOS:
            suffix = label(omega)
            symmetry_error = max(
                symmetry_error,
                abs(
                    lower[f"rhoee_linear_Omega_{suffix}"]
                    - upper[f"rhoee_linear_Omega_{suffix}"]
                ),
            )

    validation = {
        "max_formula_minus_linear_population": population_error,
        "max_detuning_reflection_asymmetry": symmetry_error,
        "checks": {
            "steady_line_shape": population_error < 1.0e-13,
            "even_in_detuning": symmetry_error < 1.0e-13,
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(
            f"line-shape validation failed: {validation}"
        )
    return rows, validation


def linewidth_rows(
    samples: int,
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
    omegas = np.linspace(0.05, 4.0, samples)
    rows: list[dict[str, Any]] = []
    linewidth_error = 0.0

    for omega in omegas:
        analytic = analytic_fwhm(
            float(omega),
            GAMMA,
            0.0,
        )
        numerical = numerical_fwhm(
            float(omega),
            GAMMA,
            0.0,
        )
        linewidth_error = max(
            linewidth_error,
            abs(analytic - numerical),
        )
        population = steady_population_formula(
            float(omega),
            0.0,
            GAMMA,
            0.0,
        )
        rows.append(
            {
                "Omega_over_Gamma": omega,
                "FWHM_over_Gamma_analytic": analytic,
                "FWHM_over_Gamma_numerical": numerical,
                "rhoee_on_resonance": population,
                "Rfl_over_Gamma_on_resonance": population,
            }
        )

    validation = {
        "max_analytic_minus_numerical_FWHM": linewidth_error,
        "weakest_drive_FWHM_over_Gamma": rows[0][
            "FWHM_over_Gamma_numerical"
        ],
        "strongest_drive_FWHM_over_Gamma": rows[-1][
            "FWHM_over_Gamma_numerical"
        ],
        "checks": {
            "linewidth_root": linewidth_error < 2.0e-12,
            "power_broadening": (
                rows[-1]["FWHM_over_Gamma_numerical"]
                > rows[0]["FWHM_over_Gamma_numerical"]
            ),
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(
            f"linewidth validation failed: {validation}"
        )
    return rows, validation


def convergence_rows(
    *,
    samples: int,
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
    times = np.linspace(0.0, TRANSIENT_TIME_MAX, samples)
    coarse_step = 0.04
    medium_step = 0.02
    fine_step = 0.01
    rows: list[dict[str, Any]] = []
    largest_fine_difference = 0.0
    minimum_ratio = math.inf

    for name, omega, delta, gamma_phi in TRANSIENT_CASES:
        traces = [
            propagate_bloch(
                times,
                initial_state=np.array([0.0, 0.0, -1.0]),
                omega=omega,
                delta=delta,
                gamma=GAMMA,
                gamma_phi=gamma_phi,
                max_step=step,
            )
            for step in (coarse_step, medium_step, fine_step)
        ]
        populations = [
            0.5 * (1.0 + trace.states[:, 2])
            for trace in traces
        ]
        coarse_medium = float(
            np.max(np.abs(populations[0] - populations[1]))
        )
        medium_fine = float(
            np.max(np.abs(populations[1] - populations[2]))
        )
        ratio = (
            coarse_medium / medium_fine
            if medium_fine > 0.0
            else math.inf
        )
        target = steady_state_analytic(
            omega,
            delta,
            GAMMA,
            gamma_phi,
        )
        largest_fine_difference = max(
            largest_fine_difference,
            medium_fine,
        )
        minimum_ratio = min(minimum_ratio, ratio)
        rows.append(
            {
                "case": name,
                "Omega_over_Gamma": omega,
                "Delta_over_Gamma": delta,
                "gamma_phi_over_Gamma": gamma_phi,
                "coarse_max_step": coarse_step,
                "medium_max_step": medium_step,
                "fine_max_step": fine_step,
                "max_abs_population_coarse_minus_medium": (
                    coarse_medium
                ),
                "max_abs_population_medium_minus_fine": (
                    medium_fine
                ),
                "refinement_difference_ratio": ratio,
                "fine_final_state_max_error_from_steady": float(
                    np.max(np.abs(traces[-1].states[-1] - target))
                ),
                "fine_max_bloch_radius_excess": traces[
                    -1
                ].max_bloch_radius_excess,
            }
        )

    validation = {
        "largest_medium_fine_population_difference": (
            largest_fine_difference
        ),
        "minimum_refinement_difference_ratio": minimum_ratio,
        "checks": {
            "fine_difference": largest_fine_difference < 2.0e-7,
            "fourth_order_regime": minimum_ratio > 10.0,
        },
    }
    if not all(validation["checks"].values()):
        raise RuntimeError(
            f"convergence 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("--transient-samples", type=int, default=601)
    parser.add_argument("--transient-max-step", type=float, default=0.005)
    parser.add_argument("--saturation-samples", type=int, default=251)
    parser.add_argument("--line-shape-samples", type=int, default=601)
    parser.add_argument("--linewidth-samples", type=int, default=80)
    parser.add_argument(
        "--output-dir",
        type=Path,
        default=Path("optical-bloch-equation-output"),
    )
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    for name in (
        "transient_samples",
        "saturation_samples",
        "line_shape_samples",
        "linewidth_samples",
    ):
        if getattr(args, name) < 5:
            raise ValueError(f"{name.replace('_', '-')} must be at least 5")
    if args.transient_samples % 2 == 0:
        raise ValueError("transient-samples must be odd")
    if args.line_shape_samples % 2 == 0:
        raise ValueError("line-shape-samples must be odd")

    transients, transient_validation = transient_rows(
        samples=args.transient_samples,
        max_step=args.transient_max_step,
    )
    decay = decay_validation(
        samples=args.transient_samples,
        max_step=args.transient_max_step,
    )
    saturation, saturation_validation = saturation_rows(
        args.saturation_samples
    )
    line_shape, line_shape_validation = line_shape_rows(
        args.line_shape_samples
    )
    linewidth, linewidth_validation = linewidth_rows(
        args.linewidth_samples
    )
    convergence, convergence_validation = convergence_rows(
        samples=args.transient_samples
    )

    args.output_dir.mkdir(parents=True, exist_ok=True)
    paths = {
        "transients": (
            args.output_dir / "optical-bloch-transients.csv"
        ),
        "saturation": (
            args.output_dir / "optical-bloch-saturation.csv"
        ),
        "line_shape": (
            args.output_dir / "optical-bloch-line-shapes.csv"
        ),
        "linewidth": (
            args.output_dir / "optical-bloch-linewidth.csv"
        ),
        "convergence": (
            args.output_dir / "optical-bloch-convergence.csv"
        ),
        "metadata": (
            args.output_dir / "optical-bloch-metadata.json"
        ),
    }
    write_csv(paths["transients"], transients)
    write_csv(paths["saturation"], saturation)
    write_csv(paths["line_shape"], line_shape)
    write_csv(paths["linewidth"], linewidth)
    write_csv(paths["convergence"], convergence)

    metadata = {
        "scope": {
            "system": "semiclassically driven dissipative two-level atom",
            "state": "unconditional density matrix in Bloch coordinates",
            "bloch_coordinates": {
                "u": "2 Re rho_eg",
                "v": "-2 Im rho_eg",
                "w": "rho_ee - rho_gg",
            },
            "detuning_convention": "Delta = omega_0 - omega_L",
            "units": {
                "Gamma": GAMMA,
                "frequency": "Gamma",
                "time": "1/Gamma",
                "hbar": 1.0,
            },
            "claim": (
                "optical Bloch solver and steady-response benchmark, "
                "not a species-specific fluorescence prediction"
            ),
        },
        "equations": {
            "u_dot": "-Gamma_2 u - Delta v",
            "v_dot": "Delta u - Gamma_2 v - Omega w",
            "w_dot": "Omega v - Gamma (w + 1)",
            "Gamma_2": "Gamma/2 + gamma_phi",
        },
        "transients": {
            "initial_state": {
                "u": 0.0,
                "v": 0.0,
                "w": -1.0,
            },
            "cases": [
                {
                    "name": name,
                    "Omega_over_Gamma": omega,
                    "Delta_over_Gamma": delta,
                    "gamma_phi_over_Gamma": gamma_phi,
                }
                for name, omega, delta, gamma_phi in TRANSIENT_CASES
            ],
            "Gamma_t_range": [0.0, TRANSIENT_TIME_MAX],
            "samples": args.transient_samples,
            "integrator": "classical explicit Runge-Kutta order 4",
            "max_internal_step": args.transient_max_step,
            "validation": transient_validation,
        },
        "decay_benchmarks": {
            "population": "rho_ee(t) = exp(-Gamma t)",
            "coherence": "u(t) = exp(-Gamma_2 t)",
            "validation": decay,
        },
        "saturation": {
            "Omega_over_Gamma_range": [0.0, 5.0],
            "gamma_phi_over_Gamma": list(
                SATURATION_DEPHASING_RATIOS
            ),
            "samples": args.saturation_samples,
            "validation": saturation_validation,
        },
        "line_shapes": {
            "Delta_over_Gamma_range": [-6.0, 6.0],
            "Omega_over_Gamma": list(
                LINE_SHAPE_DRIVE_RATIOS
            ),
            "gamma_phi_over_Gamma": 0.0,
            "samples": args.line_shape_samples,
            "validation": line_shape_validation,
        },
        "linewidth": {
            "definition": (
                "full width at half maximum in angular detuning"
            ),
            "Omega_over_Gamma_range": [0.05, 4.0],
            "samples": args.linewidth_samples,
            "root_solver": "deterministic bisection of linear steady state",
            "validation": linewidth_validation,
        },
        "convergence": {
            "max_internal_steps": [0.04, 0.02, 0.01],
            "observable_norm": "maximum absolute rho_ee difference",
            "validation": convergence_validation,
        },
        "validation": {
            "all_checks_passed": all(
                all(section["checks"].values())
                for section in (
                    transient_validation,
                    decay,
                    saturation_validation,
                    line_shape_validation,
                    linewidth_validation,
                    convergence_validation,
                )
            ),
            "sections": {
                "transients": transient_validation["checks"],
                "decay": decay["checks"],
                "saturation": saturation_validation["checks"],
                "line_shapes": line_shape_validation["checks"],
                "linewidth": linewidth_validation["checks"],
                "convergence": convergence_validation["checks"],
            },
        },
        "limitations": [
            "two-state and rotating-wave approximations are assumed",
            "Markovian population decay and pure dephasing only",
            "semiclassical monochromatic drive",
            "no multilevel branching or optical pumping",
            "no motion, Doppler averaging, spatial inhomogeneity, or recoil",
            "fluorescence rate excludes collection and detector efficiency",
            "no photon-counting trajectories or field-correlation spectrum",
        ],
        "provenance": {
            "optical_resonance_text": {
                "authors": "Allen and Eberly",
            },
            "quantum_optics_text": {
                "authors": "Cohen-Tannoudji, Dupont-Roc, and Grynberg",
            },
        },
        "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(
        "largest RK4 fine difference = "
        f"{convergence_validation['largest_medium_fine_population_difference']:.9e}"
    )
    print(
        "minimum RK4 refinement ratio = "
        f"{convergence_validation['minimum_refinement_difference_ratio']:.6f}"
    )
    print(
        "steady-state formula error   = "
        f"{saturation_validation['max_population_formula_error']:.9e}"
    )
    print(
        "linewidth extraction error   = "
        f"{linewidth_validation['max_analytic_minus_numerical_FWHM']:.9e}"
    )
    print(
        "largest Bloch-ball excess    = "
        f"{transient_validation['max_bloch_radius_excess']:.9e}"
    )
    print(f"outputs                       = {args.output_dir.resolve()}")
    print("validation                    = all checks passed")


if __name__ == "__main__":
    main()
