#!/usr/bin/env python3
# SPDX-License-Identifier: MIT
"""Reproducible Magnus-truncation benchmark for a two-step qubit drive.

The first pulse rotates about sigma_x and the second about the bisector of
sigma_x and sigma_z. The exact one-period propagator is a product of two Pauli
exponentials. Homogeneous Baker-Campbell-Hausdorff terms through degree four
provide the corresponding Magnus truncations without time-discretization error.
"""

from __future__ import annotations

import argparse
import csv
import math
import platform
import sys
import time
from pathlib import Path

import numpy as np


I2 = 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)
PAULI = (SIGMA_X, SIGMA_Y, SIGMA_Z)
AXIS_A = SIGMA_X
AXIS_B = (SIGMA_X + SIGMA_Z) / math.sqrt(2.0)
B_RATIO = 0.8
PERIOD = 1.0
STROBOSCOPIC_ETA = 0.16
BRANCH_ETA = 0.64
SWEEP_ETAS = (
    0.010,
    0.015,
    0.022,
    0.033,
    0.050,
    0.075,
    0.110,
    0.160,
    0.240,
    0.320,
    0.480,
    0.640,
    0.800,
    1.000,
    1.200,
    1.500,
    1.700,
    2.000,
)
DEFAULT_OUTPUT_DIR = Path("results")


def commutator(a: np.ndarray, b: np.ndarray) -> np.ndarray:
    return a @ b - b @ a


def pauli_exponential(angle: float, axis: np.ndarray) -> np.ndarray:
    return math.cos(angle) * I2 - 1.0j * math.sin(angle) * axis


def antihermitian_exponential(omega: np.ndarray) -> np.ndarray:
    """Exponentiate an anti-Hermitian matrix through its Hermitian generator."""

    generator = 1.0j * omega
    values, vectors = np.linalg.eigh(generator)
    return (vectors * np.exp(-1.0j * values)) @ vectors.conj().T


def hermitian_exponential(generator: np.ndarray) -> np.ndarray:
    values, vectors = np.linalg.eigh(generator)
    return (vectors * np.exp(-1.0j * values)) @ vectors.conj().T


def operator_error(approximate: np.ndarray, exact: np.ndarray) -> float:
    return float(np.linalg.norm(approximate - exact, ord="fro") / math.sqrt(2.0))


def unitarity_defect(unitary: np.ndarray) -> float:
    return float(np.linalg.norm(unitary.conj().T @ unitary - I2, ord="fro") / math.sqrt(2.0))


def pulse_generators(eta: float) -> tuple[np.ndarray, np.ndarray, float, float]:
    a = eta
    b = B_RATIO * eta
    x = -1.0j * a * AXIS_A
    y = -1.0j * b * AXIS_B
    return x, y, a, b


def exact_propagator(eta: float, reverse: bool = False) -> np.ndarray:
    _, _, a, b = pulse_generators(eta)
    u_a = pauli_exponential(a, AXIS_A)
    u_b = pauli_exponential(b, AXIS_B)
    return u_a @ u_b if reverse else u_b @ u_a


def magnus_terms(eta: float, reverse: bool = False) -> tuple[np.ndarray, ...]:
    """Return homogeneous BCH terms for log(exp(Y) exp(X))."""

    x, y, _, _ = pulse_generators(eta)
    if reverse:
        x, y = y, x
    omega_1 = x + y
    omega_2 = 0.5 * commutator(y, x)
    omega_3 = (
        commutator(y, commutator(y, x))
        + commutator(x, commutator(x, y))
    ) / 12.0
    omega_4 = -commutator(x, commutator(y, commutator(y, x))) / 24.0
    return omega_1, omega_2, omega_3, omega_4


def cumulative_magnus(eta: float) -> tuple[np.ndarray, ...]:
    terms = magnus_terms(eta)
    cumulative: list[np.ndarray] = []
    running = np.zeros((2, 2), dtype=complex)
    for term in terms:
        running = running + term
        cumulative.append(running.copy())
    return tuple(cumulative)


def dyson_second_order(eta: float) -> np.ndarray:
    """Return the degree-two series of exp(Y) exp(X), without exponentiation."""

    x, y, _, _ = pulse_generators(eta)
    return I2 + x + y + 0.5 * (x @ x) + y @ x + 0.5 * (y @ y)


def principal_generator(unitary: np.ndarray) -> tuple[np.ndarray, float, np.ndarray]:
    """Return K with U=exp(-iK), principal SU(2) angle, and rotation axis."""

    scalar = float(np.clip(np.real(np.trace(unitary)) / 2.0, -1.0, 1.0))
    vector = np.array(
        [float(np.real(0.5j * np.trace(sigma @ unitary))) for sigma in PAULI]
    )
    vector_norm = float(np.linalg.norm(vector))
    angle = math.atan2(vector_norm, scalar)
    if vector_norm < 1e-14:
        axis = np.array([1.0, 0.0, 0.0])
    else:
        axis = vector / vector_norm
    generator = angle * sum(component * sigma for component, sigma in zip(axis, PAULI))
    return generator, angle, axis


def sweep_rows() -> list[dict[str, float]]:
    rows: list[dict[str, float]] = []
    for eta in SWEEP_ETAS:
        exact = exact_propagator(eta)
        cumulative = cumulative_magnus(eta)
        approximations = [antihermitian_exponential(omega) for omega in cumulative]
        terms = magnus_terms(eta)
        dyson = dyson_second_order(eta)
        generator_exact, angle, axis = principal_generator(exact)
        reverse_terms = magnus_terms(eta, reverse=True)

        row: dict[str, float] = {
            "eta": eta,
            "pulse_area_a": eta,
            "pulse_area_b": B_RATIO * eta,
            "integrated_norm": (1.0 + B_RATIO) * eta,
            "sufficient_bound_satisfied": float((1.0 + B_RATIO) * eta < math.pi),
            "principal_angle": angle,
            "axis_x": float(axis[0]),
            "axis_y": float(axis[1]),
            "axis_z": float(axis[2]),
            "exact_unitarity_defect": unitarity_defect(exact),
            "exact_determinant_defect": float(abs(np.linalg.det(exact) - 1.0)),
            "pulse_order_operator_difference": operator_error(
                exact_propagator(eta, reverse=True), exact
            ),
            "omega2_reversal_residual": float(
                np.linalg.norm(reverse_terms[1] + terms[1], ord="fro")
            ),
            "dyson2_error": operator_error(dyson, exact),
            "dyson2_unitarity_defect": unitarity_defect(dyson),
        }

        for order, (omega, approximation, term) in enumerate(
            zip(cumulative, approximations, terms), start=1
        ):
            row[f"magnus{order}_error"] = operator_error(approximation, exact)
            row[f"magnus{order}_unitarity_defect"] = unitarity_defect(approximation)
            row[f"magnus{order}_unitarity_plot"] = max(
                row[f"magnus{order}_unitarity_defect"], 1e-18
            )
            row[f"omega{order}_norm"] = float(np.linalg.norm(term, ord="fro"))
            row[f"generator{order}_error"] = float(
                np.linalg.norm(1.0j * omega - generator_exact, ord="fro")
                / math.sqrt(2.0)
            )
        rows.append(row)
    return rows


def stroboscopic_rows() -> list[dict[str, float]]:
    exact = exact_propagator(STROBOSCOPIC_ETA)
    approximations = [
        antihermitian_exponential(omega)
        for omega in cumulative_magnus(STROBOSCOPIC_ETA)
    ]
    initial = np.array([1.0, 0.0], dtype=complex)
    rows: list[dict[str, float]] = []

    for periods in range(1, 201):
        exact_n = np.linalg.matrix_power(exact, periods)
        exact_state = exact_n @ initial
        row: dict[str, float] = {
            "eta": STROBOSCOPIC_ETA,
            "periods": float(periods),
            "exact_transition_probability": float(abs(exact_state[1]) ** 2),
            "exact_unitarity_defect": unitarity_defect(exact_n),
        }
        for order, approximation in enumerate(approximations, start=1):
            approximation_n = np.linalg.matrix_power(approximation, periods)
            state_n = approximation_n @ initial
            row[f"magnus{order}_operator_error"] = operator_error(
                approximation_n, exact_n
            )
            row[f"magnus{order}_transition_probability"] = float(
                abs(state_n[1]) ** 2
            )
            row[f"magnus{order}_probability_error"] = abs(
                row[f"magnus{order}_transition_probability"]
                - row["exact_transition_probability"]
            )
            row[f"magnus{order}_unitarity_defect"] = unitarity_defect(
                approximation_n
            )
        rows.append(row)
    return rows


def branch_rows() -> list[dict[str, float]]:
    exact = exact_propagator(BRANCH_ETA)
    generator, angle, _ = principal_generator(exact)
    values, vectors = np.linalg.eigh(generator)
    rows: list[dict[str, float]] = []

    for lower_shift in (-1, 0, 1):
        for upper_shift in (-1, 0, 1):
            shifted = np.array(
                [
                    values[0] + 2.0 * math.pi * lower_shift,
                    values[1] + 2.0 * math.pi * upper_shift,
                ]
            )
            branch_generator = (vectors * shifted) @ vectors.conj().T
            reconstructed = hermitian_exponential(branch_generator)
            rows.append(
                {
                    "eta": BRANCH_ETA,
                    "principal_angle": angle,
                    "lower_shift": float(lower_shift),
                    "upper_shift": float(upper_shift),
                    "lower_quasienergy": float(shifted[0] / PERIOD),
                    "upper_quasienergy": float(shifted[1] / PERIOD),
                    "trace": float(np.real(np.trace(branch_generator))),
                    "generator_distance_from_principal": float(
                        np.linalg.norm(branch_generator - generator, ord="fro")
                        / math.sqrt(2.0)
                    ),
                    "reconstruction_error": operator_error(reconstructed, exact),
                }
            )
    return rows


def observed_orders(sweep: list[dict[str, float]]) -> dict[str, float]:
    asymptotic = [row for row in sweep if row["eta"] <= 0.075]
    log_eta = np.log([row["eta"] for row in asymptotic])
    return {
        f"magnus{order}_observed_order": float(
            np.polyfit(
                log_eta,
                np.log([row[f"magnus{order}_error"] for row in asymptotic]),
                1,
            )[0]
        )
        for order in range(1, 5)
    }


def commuting_control_error(eta: float = 0.9) -> tuple[float, float]:
    a = eta
    b = B_RATIO * eta
    x = -1.0j * a * SIGMA_X
    y = -1.0j * b * SIGMA_X
    exact = pauli_exponential(a + b, SIGMA_X)
    first_magnus = antihermitian_exponential(x + y)
    commutator_norm = float(np.linalg.norm(commutator(y, x), ord="fro"))
    return operator_error(first_magnus, exact), commutator_norm


def validate(
    sweep: list[dict[str, float]],
    stroboscopic: list[dict[str, float]],
    branches: list[dict[str, float]],
) -> dict[str, float]:
    orders = observed_orders(sweep)
    by_eta = {row["eta"]: row for row in sweep}
    commuting_error, commuting_commutator = commuting_control_error()
    metrics = {
        **orders,
        "max_magnus_unitarity_defect": max(
            row[f"magnus{order}_unitarity_defect"]
            for row in sweep
            for order in range(1, 5)
        ),
        "max_exact_unitarity_defect": max(
            row["exact_unitarity_defect"] for row in sweep
        ),
        "max_exact_determinant_defect": max(
            row["exact_determinant_defect"] for row in sweep
        ),
        "max_omega2_reversal_residual": max(
            row["omega2_reversal_residual"] for row in sweep
        ),
        "commuting_control_error": commuting_error,
        "commuting_control_commutator": commuting_commutator,
        "dyson2_unitarity_defect_eta_0p64": by_eta[0.64][
            "dyson2_unitarity_defect"
        ],
        "magnus4_error_eta_0p32": by_eta[0.32]["magnus4_error"],
        "magnus3_error_eta_0p32": by_eta[0.32]["magnus3_error"],
        "magnus4_error_eta_1p5": by_eta[1.5]["magnus4_error"],
        "magnus3_error_eta_1p5": by_eta[1.5]["magnus3_error"],
        "max_branch_reconstruction_error": max(
            row["reconstruction_error"] for row in branches
        ),
        "max_stroboscopic_magnus_unitarity_defect": max(
            row[f"magnus{order}_unitarity_defect"]
            for row in stroboscopic
            for order in range(1, 5)
        ),
        "magnus4_error_200_periods": stroboscopic[-1][
            "magnus4_operator_error"
        ],
        "magnus3_error_200_periods": stroboscopic[-1][
            "magnus3_operator_error"
        ],
    }

    checks = {
        "Magnus-1 order": abs(orders["magnus1_observed_order"] - 2.0) < 0.02,
        "Magnus-2 order": abs(orders["magnus2_observed_order"] - 3.0) < 0.02,
        "Magnus-3 order": abs(orders["magnus3_observed_order"] - 4.0) < 0.02,
        "Magnus-4 order": abs(orders["magnus4_observed_order"] - 5.0) < 0.03,
        "Magnus unitarity": metrics["max_magnus_unitarity_defect"] < 2e-14,
        "exact unitarity": metrics["max_exact_unitarity_defect"] < 2e-14,
        "exact determinant": metrics["max_exact_determinant_defect"] < 2e-14,
        "pulse-order sign": metrics["max_omega2_reversal_residual"] < 2e-14,
        "commuting control": commuting_error < 2e-14 and commuting_commutator < 2e-14,
        "Dyson diagnostic active": metrics["dyson2_unitarity_defect_eta_0p64"] > 0.1,
        "small-eta order hierarchy": metrics["magnus4_error_eta_0p32"]
        < metrics["magnus3_error_eta_0p32"],
        "low-order nonmonotonicity visible": metrics["magnus4_error_eta_1p5"]
        > metrics["magnus3_error_eta_1p5"],
        "branch reconstruction": metrics["max_branch_reconstruction_error"] < 2e-14,
        "stroboscopic unitarity": metrics[
            "max_stroboscopic_magnus_unitarity_defect"
        ]
        < 2e-13,
        "bound classification": by_eta[1.7]["sufficient_bound_satisfied"] == 1.0
        and by_eta[2.0]["sufficient_bound_satisfied"] == 0.0,
    }
    failed = [name for name, passed in checks.items() if not passed]
    if failed:
        raise RuntimeError("Validation failed: " + ", ".join(failed))
    return metrics


def write_csv(path: Path, rows: list[dict[str, float]]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)


def make_quicklook_plots(
    output_dir: Path,
    sweep: list[dict[str, float]],
    stroboscopic: list[dict[str, float]],
) -> None:
    try:
        import matplotlib.pyplot as plt
    except ImportError as exc:
        raise SystemExit("--plot requires Matplotlib.") from exc

    eta = np.array([row["eta"] for row in sweep])
    figure, axes = plt.subplots(2, 1, figsize=(7.0, 8.0), constrained_layout=True)
    for order in range(1, 5):
        axes[0].loglog(
            eta,
            [row[f"magnus{order}_error"] for row in sweep],
            marker="o",
            label=f"Magnus {order}",
        )
    axes[0].set(xlabel="pulse scale eta", ylabel="one-period operator error")
    axes[0].legend()
    axes[0].grid(alpha=0.25)
    axes[1].loglog(
        eta,
        [row["dyson2_unitarity_defect"] for row in sweep],
        "s-",
        label="Dyson degree 2",
    )
    axes[1].loglog(
        eta,
        [
            max(row[f"magnus{order}_unitarity_plot"] for order in range(1, 5))
            for row in sweep
        ],
        "o-",
        label="largest Magnus defect",
    )
    axes[1].set(xlabel="pulse scale eta", ylabel="unitarity defect")
    axes[1].legend()
    axes[1].grid(alpha=0.25)
    figure.savefig(output_dir / "magnus-expansion-error-orders-quicklook.png", dpi=180)
    plt.close(figure)

    periods = np.array([row["periods"] for row in stroboscopic])
    figure, axes = plt.subplots(2, 1, figsize=(7.0, 8.0), constrained_layout=True)
    for order in range(1, 5):
        axes[0].semilogy(
            periods,
            [row[f"magnus{order}_operator_error"] for row in stroboscopic],
            label=f"Magnus {order}",
        )
    axes[0].set(xlabel="periods", ylabel="operator error")
    axes[0].legend()
    axes[0].grid(alpha=0.25)
    axes[1].plot(
        periods,
        [row["exact_transition_probability"] for row in stroboscopic],
        color="black",
        label="exact",
    )
    for order in (2, 3, 4):
        axes[1].plot(
            periods,
            [row[f"magnus{order}_transition_probability"] for row in stroboscopic],
            label=f"Magnus {order}",
        )
    axes[1].set(xlabel="periods", ylabel="transition probability")
    axes[1].legend()
    axes[1].grid(alpha=0.25)
    figure.savefig(output_dir / "magnus-expansion-stroboscopic-quicklook.png", dpi=180)
    plt.close(figure)


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--output-dir",
        type=Path,
        default=DEFAULT_OUTPUT_DIR,
        help="Directory for retained CSV data and optional quick-look plots (default: results in the working directory).",
    )
    parser.add_argument(
        "--plot",
        action="store_true",
        help="Generate optional Matplotlib quick-look PNGs.",
    )
    return parser.parse_args()


def main() -> int:
    args = parse_args()
    started = time.perf_counter()
    sweep = sweep_rows()
    stroboscopic = stroboscopic_rows()
    branches = branch_rows()
    metrics = validate(sweep, stroboscopic, branches)

    sweep_path = args.output_dir / "magnus-expansion-error-sweep.csv"
    stroboscopic_path = args.output_dir / "magnus-expansion-stroboscopic.csv"
    branch_path = args.output_dir / "magnus-effective-hamiltonian-branches.csv"
    write_csv(sweep_path, sweep)
    write_csv(stroboscopic_path, stroboscopic)
    write_csv(branch_path, branches)

    if args.plot:
        make_quicklook_plots(args.output_dir, sweep, stroboscopic)

    elapsed = time.perf_counter() - started
    print("Magnus expansion error benchmark")
    print(f"Python: {platform.python_version()}")
    print(f"NumPy: {np.__version__}")
    print(f"Platform: {platform.platform()}")
    print(f"Runtime: {elapsed:.3f} s")
    print(f"Retained error sweep: {sweep_path}")
    print(f"Retained stroboscopic data: {stroboscopic_path}")
    print(f"Retained branch data: {branch_path}")
    for name, value in metrics.items():
        print(f"{name}: {value:.6e}")
    print("Validation: PASS")
    return 0


if __name__ == "__main__":
    sys.exit(main())
