"""Separate a toy dephasing model, sampling uncertainty and integration error."""

from __future__ import annotations

import json
from math import exp, sqrt

import numpy as np
from numpy.typing import NDArray
from scipy.integrate import solve_ivp
from scipy.linalg import expm


def experiment() -> dict[str, object]:
    x = np.array([[0, 1], [1, 0]], dtype=complex)
    z = np.diag([1, -1])
    initial = np.array([1, 0], dtype=complex)
    duration = 0.8

    def hamiltonian(time: float) -> NDArray[np.complex128]:
        return (1.7 * x + 3.0 * time * z) / 2

    solution = solve_ivp(
        lambda time, state: -1j * hamiltonian(time) @ state,
        (0, duration),
        initial,
        method="DOP853",
        rtol=1e-11,
        atol=1e-13,
    )
    if not solution.success:
        raise RuntimeError(solution.message)
    reference = solution.y[:, -1]
    integration = []
    for steps in (4, 64):
        state = initial.copy()
        dt = duration / steps
        for step in range(steps):
            state = expm(-1j * hamiltonian((step + 0.5) * dt) * dt) @ state
        integration.append(
            {
                "steps": steps,
                "state_error_norm": float(np.linalg.norm(state - reference)),
            }
        )
    return {
        "toy_dephasing_gamma_per_us": 0.4,
        "duration_us": duration,
        "x_after_toy_dephasing": exp(-0.4 * duration),
        "sampling_se_at_p_half": {str(n): sqrt(0.25 / n) for n in (100, 10000)},
        "midpoint_integration": integration,
    }


if __name__ == "__main__":
    print(json.dumps(experiment(), sort_keys=True))
