#!/usr/bin/env python3
"""Generate and verify Graviton's reproducible radical-pair demonstration.

Requirements: Python 3, NumPy and SciPy. Run:
    python3 generate_spin_data.py
Optional:
    python3 generate_spin_data.py --output /path/to/spin-data.json

No biological parameters have been fitted. No consciousness variable is defined.
The JSON contains values at exactly the listed grid points: no interpolation,
clipping, renormalization, or replacement of the triplet yield by 1-singlet.

Conventions:
  Basis order: electron 1, electron 2, one spin-1/2 nucleus; up, down per spin.
  S = Pauli/2 is dimensionless. Omega=H/hbar is in radians per second.
  Electron Zeeman term: +|gamma_e| B dot (S1+S2), |gamma_e|=1.76085963e11.
  A/(2*pi)=(2,2,4) MHz; A is diagonal in the molecular x,y,z frame.
  B direction=(sin(theta)*cos(phi),sin(theta)*sin(phi),cos(theta)).
  kS=kT=k=1/lifetime, in s^-1. Nuclear initial state is unpolarized.
  Gamma is the Bloch-vector relaxation rate of an isolated electron under
  Gamma/4 sum_a (sigma_a rho sigma_a-rho), not a generic 'coherence lifetime'.

The numerical method solves the integrated, unrecombined 64-component
Liouvillian equation. Independent verification integrates a normalized
8x8 density matrix directly with DOP853 and separately accumulates product
yields with the exponential survival probability. This checks vectorization,
reaction normalization, and numerical convergence in a different formulation.
"""
from __future__ import annotations

import argparse
from datetime import datetime, timezone
import hashlib
import json
from pathlib import Path
import platform
from time import perf_counter

import numpy as np
import scipy
from scipy.integrate import solve_ivp
from scipy.linalg import expm, solve

FIELDS = list(range(0, 501, 10))
ANGLES = list(range(0, 91, 15))
LIFETIMES = [0.1, 0.5, 2.0]
RELAXATIONS = [0.0, 1e5, 1e6]
HYPERFINE_MHZ = (2.0, 2.0, 4.0)
GAMMA_E = 1.76085963e11
I2 = np.eye(2, dtype=complex)
PAULI = [np.array([[0, 1], [1, 0]], complex),
         np.array([[0, -1j], [1j, 0]], complex),
         np.array([[1, 0], [0, -1]], complex)]


def embed(matrix, site):
    factors = [I2, I2, I2]
    factors[site] = matrix
    return np.kron(np.kron(factors[0], factors[1]), factors[2])


S = [[embed(p / 2, site) for p in PAULI] for site in range(3)]
ELECTRON_PAULI = [embed(p, site) for site in (0, 1) for p in PAULI]
I8 = np.eye(8, dtype=complex)
I64 = np.eye(64, dtype=complex)
singlet = np.array([0, 1, -1, 0], complex) / np.sqrt(2)
PS = np.kron(np.outer(singlet, singlet.conj()), I2)
PT = I8 - PS
RHO0 = PS / 2
RHO0_VEC = RHO0.ravel(order="F")
RELAXATION_SUPEROPERATOR = sum(
    (np.kron(p.conj(), p) - I64) / 4 for p in ELECTRON_PAULI
)


def omega(field_ut, angle_deg, azimuth_deg=0.0, hyperfine_mhz=HYPERFINE_MHZ):
    theta, phi = np.deg2rad([angle_deg, azimuth_deg])
    direction = [np.sin(theta) * np.cos(phi),
                 np.sin(theta) * np.sin(phi), np.cos(theta)]
    zeeman_frequency = GAMMA_E * field_ut * 1e-6
    return sum(
        zeeman_frequency * direction[a] * (S[0][a] + S[1][a])
        + 2 * np.pi * hyperfine_mhz[a] * 1e6 * (S[0][a] @ S[2][a])
        for a in range(3)
    )


def integrated_solution(field_ut, angle_deg, lifetime_us, relaxation_per_s,
                        azimuth_deg=0.0, hyperfine_mhz=HYPERFINE_MHZ):
    om = omega(field_ut, angle_deg, azimuth_deg, hyperfine_mhz)
    k = 1e6 / lifetime_us
    liouvillian = (-1j * (np.kron(I8, om) - np.kron(om.T, I8))
                   - k * I64 + relaxation_per_s * RELAXATION_SUPEROPERATOR)
    # If rho(infinity)=0, L integral(rho dt) = -rho(0).
    integral_vector = solve(-liouvillian, RHO0_VEC, assume_a="gen")
    integral = integral_vector.reshape((8, 8), order="F")
    channel_yields = np.array([k * np.trace(PS @ integral),
                              k * np.trace(PT @ integral)])
    mixture = k * integral
    residual = np.linalg.norm(-liouvillian @ integral_vector - RHO0_VEC)
    residual /= np.linalg.norm(RHO0_VEC)
    diagnostics = {
        "conservation_error": float(abs(channel_yields.sum() - 1)),
        "imaginary_yield_max": float(np.max(abs(channel_yields.imag))),
        "relative_linear_solve_residual": float(residual),
        "integrated_state_hermiticity_error": float(np.linalg.norm(mixture - mixture.conj().T)),
        "integrated_state_min_eigenvalue": float(np.linalg.eigvalsh((mixture + mixture.conj().T) / 2).min()),
    }
    return channel_yields.real, diagnostics


def time_integrated_solution(case, rtol, atol, end_lifetimes=24.0, reference_check=False):
    """Direct matrix ODE; no Kronecker/vectorized Liouvillian is used here."""
    om = omega(case["field_microtesla"], case["angle_degrees"])
    k = 1e6 / case["lifetime_us"]
    relaxation = case["relaxation_per_s"]
    # u=k*t. sigma=rho/Tr(rho), so Tr(sigma)=1 and positivity remains inspectable
    # even after most pairs have reacted. Product accumulation includes exp(-u).
    def rhs(u, augmented):
        density = augmented[:64].reshape((8, 8))
        derivative = -1j * (om @ density - density @ om) / k
        if relaxation:
            derivative += relaxation / (4 * k) * sum(
                p @ density @ p - density for p in ELECTRON_PAULI
            )
        singlet_flux = np.exp(-u) * np.trace(PS @ density).real
        triplet_flux = np.exp(-u) * np.trace(PT @ density).real
        return np.r_[derivative.ravel(), singlet_flux, triplet_flux]

    initial = np.r_[RHO0.ravel(), 0j, 0j]
    solution = solve_ivp(rhs, (0.0, end_lifetimes), initial, method="DOP853",
                         rtol=rtol, atol=atol,
                         t_eval=np.linspace(0.0, end_lifetimes, 193))
    if not solution.success:
        raise RuntimeError(solution.message)
    states = solution.y[:64].T.reshape((-1, 8, 8))
    adjoints = states.conj().transpose(0, 2, 1)
    eigenvalues = np.linalg.eigvalsh((states + adjoints) / 2)
    final = solution.y[-2:, -1].real
    diagnostics = {
        "ode_evaluations": int(solution.nfev),
        "sampled_state_count": int(len(states)),
        "normalized_state_min_eigenvalue": float(eigenvalues.min()),
        "normalized_state_trace_error_max": float(np.max(abs(np.trace(states, axis1=1, axis2=2) - 1))),
        "normalized_state_hermiticity_error_max": float(np.max(np.linalg.norm(states - adjoints, axis=(1, 2)))),
        "finite_time_total_yield_error": float(abs(final.sum() - (1 - np.exp(-end_lifetimes)))),
    }
    if reference_check:
        generator = (-1j * (np.kron(I8, om) - np.kron(om.T, I8))
                     + relaxation * RELAXATION_SUPEROPERATOR) / k
        min_eigenvalue, max_state_error, max_trace_error = 0.0, 0.0, 0.0
        # These are entries of t_eval, in dimensionless time u=k*t.
        for index in (0, 1, 8, 40, 192):
            direct_exponential = (expm(generator * solution.t[index]) @ RHO0_VEC)
            reference = direct_exponential.reshape((8, 8), order="F")
            min_eigenvalue = min(min_eigenvalue, float(np.linalg.eigvalsh((reference + reference.conj().T) / 2).min()))
            max_state_error = max(max_state_error, float(np.linalg.norm(reference - states[index])))
            max_trace_error = max(max_trace_error, float(abs(np.trace(reference) - 1)))
        diagnostics["matrix_exponential_sample_count"] = 5
        diagnostics["matrix_exponential_min_eigenvalue"] = min_eigenvalue
        diagnostics["matrix_exponential_vs_ode_state_error_max"] = max_state_error
        diagnostics["matrix_exponential_trace_error_max"] = max_trace_error
    return final, diagnostics


def verify(cases, grid_diagnostics):
    thresholds = {"conservation": 1e-10, "imaginary_yield": 1e-10,
                  "relative_linear_solve_residual": 1e-10,
                  "integrated_state_hermiticity": 1e-10,
                  "grid_integrated_state_positivity": -1e-10,
                  "symmetry": 1e-10, "time_integration_yield_agreement": 2e-9,
                  "ode_tolerance_convergence": 2e-8,
                  "ode_normalized_state_positivity": -2e-9,
                  "ode_normalized_state_trace": 1e-9,
                  "matrix_exponential_positivity": -1e-10,
                  "matrix_exponential_vs_ode_state": 1e-8,
                  "analytic_relaxation_yield": 1e-10}
    grid = {
        "grid_point_count": len(FIELDS) * len(cases),
        "max_conservation_error": max(x["conservation_error"] for x in grid_diagnostics),
        "max_imaginary_yield": max(x["imaginary_yield_max"] for x in grid_diagnostics),
        "max_relative_linear_solve_residual": max(x["relative_linear_solve_residual"] for x in grid_diagnostics),
        "max_integrated_state_hermiticity_error": max(x["integrated_state_hermiticity_error"] for x in grid_diagnostics),
        "min_integrated_state_eigenvalue": min(x["integrated_state_min_eigenvalue"] for x in grid_diagnostics),
        "minimum_yield": min(min(c[y]) for c in cases for y in ("singlet", "triplet")),
        "maximum_yield": max(max(c[y]) for c in cases for y in ("singlet", "triplet")),
    }
    assert grid["max_conservation_error"] < thresholds["conservation"]
    assert grid["max_imaginary_yield"] < thresholds["imaginary_yield"]
    assert grid["max_relative_linear_solve_residual"] < thresholds["relative_linear_solve_residual"]
    assert grid["max_integrated_state_hermiticity_error"] < thresholds["integrated_state_hermiticity"]
    assert grid["min_integrated_state_eigenvalue"] > thresholds["grid_integrated_state_positivity"]
    assert 0 <= grid["minimum_yield"] <= grid["maximum_yield"] <= 1

    zero_field_error = 0.0
    for lifetime in LIFETIMES:
        for relaxation in RELAXATIONS:
            values = [c["singlet"][0] for c in cases
                      if c["lifetime_us"] == lifetime and c["relaxation_per_s"] == relaxation]
            zero_field_error = max(zero_field_error, float(np.ptp(values)))
    axial_errors, reversal_errors, isotropic_errors = [], [], []
    for b, angle, tau, rate in [(50, 30, .5, 1e5), (500, 75, 2, 0), (130, 45, .1, 1e6)]:
        baseline, _ = integrated_solution(b, angle, tau, rate)
        for phi in (37, 91, 173):
            rotated, _ = integrated_solution(b, angle, tau, rate, azimuth_deg=phi)
            axial_errors.append(float(np.max(abs(rotated - baseline))))
        reversed_field, _ = integrated_solution(-b, angle, tau, rate)
        reversal_errors.append(float(np.max(abs(reversed_field - baseline))))
        iso0, _ = integrated_solution(b, 0, tau, rate, hyperfine_mhz=(2, 2, 2))
        iso1, _ = integrated_solution(b, angle, tau, rate, azimuth_deg=53, hyperfine_mhz=(2, 2, 2))
        isotropic_errors.append(float(np.max(abs(iso1 - iso0))))
    symmetry = {"zero_field_angle_invariance_max_error": zero_field_error,
                "axial_azimuth_invariance_max_error": max(axial_errors),
                "field_reversal_invariance_max_error": max(reversal_errors),
                "isotropic_hyperfine_orientation_invariance_max_error": max(isotropic_errors)}
    assert max(symmetry.values()) < thresholds["symmetry"]

    no_mix, _ = integrated_solution(500, 45, .5, 0, hyperfine_mhz=(0, 0, 0))
    fast_relax, _ = integrated_solution(50, 0, .5, 1e12)
    fast_react, _ = integrated_solution(50, 0, 1e-6, 0)
    no_hyperfine_relax, _ = integrated_solution(500, 45, .5, 1e5, hyperfine_mhz=(0, 0, 0))
    analytic_no_hyperfine_relax = .25 + .75 * 2e6 / (2e6 + 2e5)
    limits = {"no_hyperfine_no_relaxation_singlet_yield": float(no_mix[0]),
              "no_hyperfine_expected_singlet_yield": 1.0,
              "very_fast_relaxation_singlet_yield": float(fast_relax[0]),
              "fast_relaxation_asymptotic_singlet_yield": .25,
              "very_fast_reaction_singlet_yield": float(fast_react[0]),
              "fast_reaction_asymptotic_singlet_yield": 1.0,
              "no_hyperfine_with_relaxation_singlet_yield": float(no_hyperfine_relax[0]),
              "analytic_no_hyperfine_with_relaxation_yield": analytic_no_hyperfine_relax,
              "analytic_relaxation_formula": "Phi_S=1/4+(3/4)*k/(k+2*Gamma) when A=0",
              "analytic_relaxation_check_parameters": {"field_microtesla": 500, "angle_degrees": 45,
                                                       "lifetime_us": .5, "relaxation_per_s": 1e5},
              "finite_extreme_parameter_values": {"relaxation_per_s": 1e12,
                                                   "fast_reaction_lifetime_us": 1e-6}}
    assert abs(no_mix[0] - 1) < 1e-10
    assert abs(fast_relax[0] - .25) < 1e-5
    assert abs(fast_react[0] - 1) < 1e-8
    assert abs(no_hyperfine_relax[0] - analytic_no_hyperfine_relax) < thresholds["analytic_relaxation_yield"]

    representatives = [
        (0, 0, .1, 0), (50, 0, .5, 1e5), (50, 90, .5, 1e5),
        (500, 90, 2, 0), (130, 30, 2, 1e6), (500, 45, .1, 1e6),
        (0, 75, 2, 1e6),
    ]
    time_checks = []
    for b, angle, tau, rate in representatives:
        case = dict(field_microtesla=b, angle_degrees=angle,
                    lifetime_us=tau, relaxation_per_s=rate)
        exact, _ = integrated_solution(b, angle, tau, rate)
        coarse, _ = time_integrated_solution(case, rtol=1e-8, atol=1e-10)
        fine, diagnostics = time_integrated_solution(case, rtol=1e-12, atol=1e-14, reference_check=True)
        error = float(np.max(abs(exact - fine)))
        convergence = float(np.max(abs(fine - coarse)))
        assert error < thresholds["time_integration_yield_agreement"]
        assert convergence < thresholds["ode_tolerance_convergence"]
        assert diagnostics["normalized_state_min_eigenvalue"] > thresholds["ode_normalized_state_positivity"]
        assert diagnostics["normalized_state_trace_error_max"] < thresholds["ode_normalized_state_trace"]
        assert diagnostics["matrix_exponential_min_eigenvalue"] > thresholds["matrix_exponential_positivity"]
        assert diagnostics["matrix_exponential_vs_ode_state_error_max"] < thresholds["matrix_exponential_vs_ode_state"]
        time_checks.append({**case, "integrated_liouvillian_yields": exact.tolist(),
                            "fine_time_integral_yields": fine.tolist(),
                            "max_yield_difference": error,
                            "coarse_to_fine_yield_difference": convergence, **diagnostics})

    return {
        "status": "passed",
        "scope": "Numerical consistency of the stated idealized model; no empirical validation.",
        "thresholds": thresholds, "grid": grid, "symmetries": symmetry,
        "limiting_cases": limits,
        "independent_time_integration": {
            "method": "DOP853 direct 8x8 normalized density matrix plus separately accumulated product yields",
            "end_time_in_lifetimes": 24.0,
            "omitted_total_survival_probability_bound": float(np.exp(-24)),
            "coarse_tolerances": {"rtol": 1e-8, "atol": 1e-10},
            "fine_tolerances": {"rtol": 1e-12, "atol": 1e-14},
            "positivity_scope": "193 sampled normalized states per representative fine integration; no claim of an all-time numerical proof.",
            "max_yield_difference": max(c["max_yield_difference"] for c in time_checks),
            "max_tolerance_convergence_difference": max(c["coarse_to_fine_yield_difference"] for c in time_checks),
            "min_sampled_normalized_state_eigenvalue": min(c["normalized_state_min_eigenvalue"] for c in time_checks),
            "max_sampled_normalized_state_trace_error": max(c["normalized_state_trace_error_max"] for c in time_checks),
            "matrix_exponential_validation": {
                "scope": "Five normalized states per representative case, evaluated using the matrix exponential independently of the ODE integrator.",
                "min_eigenvalue": min(c["matrix_exponential_min_eigenvalue"] for c in time_checks),
                "max_ode_state_difference": max(c["matrix_exponential_vs_ode_state_error_max"] for c in time_checks),
                "max_trace_error": max(c["matrix_exponential_trace_error_max"] for c in time_checks),
            },
            "cases": time_checks,
        },
    }


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--output", type=Path, default=Path(__file__).with_name("spin-data.json"))
    args = parser.parse_args()
    started = perf_counter()
    cases, diagnostics = [], []
    for lifetime in LIFETIMES:
        for relaxation in RELAXATIONS:
            for angle in ANGLES:
                singlets, triplets = [], []
                for field in FIELDS:
                    values, diag = integrated_solution(field, angle, lifetime, relaxation)
                    singlets.append(float(values[0]))
                    triplets.append(float(values[1]))
                    diagnostics.append(diag)
                cases.append(dict(lifetime_us=lifetime, relaxation_per_s=relaxation,
                                  angle_degrees=angle, singlet=singlets, triplet=triplets))
    grid_seconds = perf_counter() - started
    verification = verify(cases, diagnostics)
    elapsed = perf_counter() - started
    verification["runtime_seconds"] = {"grid": grid_seconds, "total_with_verification": elapsed}
    source_path = Path(__file__)
    metadata = {
        "title": "Idealized radical-pair reaction: exact discrete numerical grid",
        "schema_version": 1,
        "status": "computational_benchmark_not_biological",
        "generated_utc": datetime.now(timezone.utc).isoformat(),
        "generation_file": source_path.name,
        "generation_source_sha256": hashlib.sha256(source_path.read_bytes()).hexdigest(),
        "dependencies": {"python": platform.python_version(), "numpy": np.__version__, "scipy": scipy.__version__},
        "hilbert_space_dimension": 8,
        "spin_order": ["electron_1", "electron_2", "spin_half_nucleus"],
        "basis_per_spin": ["up", "down"],
        "spin_operator_convention": "dimensionless S=Pauli/2",
        "initial_state": "electron singlet tensor unpolarized nuclear I/2",
        "hamiltonian": "Omega=H/hbar=+gamma_e*B dot (S1+S2)+S1 dot A dot I",
        "electron_gyromagnetic_magnitude_rad_per_s_per_tesla": GAMMA_E,
        "hyperfine_A_over_2pi_MHz": list(HYPERFINE_MHZ),
        "hyperfine_tensor_frame": "diagonal in molecular x,y,z; z is unique axial direction",
        "field_direction": "(sin(theta)*cos(phi),sin(theta)*sin(phi),cos(theta)); grid phi=0",
        "angle_meaning": "polar angle theta between magnetic field and molecular z axis",
        "reaction": "equal singlet and triplet rates k=1/lifetime; irreversible two-product branching",
        "master_equation": "rho_dot=-i[Omega,rho]-k*rho+(Gamma/4)sum_{j=1,2;a=x,y,z}(sigma_ja*rho*sigma_ja-rho)",
        "relaxation_rate_meaning": "Gamma is the isolated electron Bloch-vector relaxation rate for stipulated isotropic local spin relaxation",
        "yield_definition": "Phi_channel=k*integral_0^infinity Tr(P_channel*rho(t))dt; channels computed independently",
        "axis_units": {"fields": "microtesla", "angles": "degrees", "lifetimes": "microseconds",
                       "relaxations": "per_second", "yields": "dimensionless_probability"},
        "computation": "64-dimensional integrated Liouvillian linear solve using complex128",
        "grid_semantics": "Only explicitly listed discrete parameter points. No interpolation, clipping, renormalization, or empirical fit.",
        "orientation_semantics": "Fixed molecular orientation; freely rotating ensembles require separate averaging/dynamics.",
        "limitations": [
            "One model spin-1/2 nucleus; no fitted protein or cellular parameters.",
            "No nuclear Zeeman, exchange, dipolar, diffusion, back-reaction, or unequal-rate terms.",
            "Fixed hyperfine tensor and stipulated Markovian isotropic relaxation.",
            "No neuronal response, physiological relevance, entanglement witness, or consciousness score is inferred.",
            "Relaxation changes compare regimes of one quantum model, not proof against all conventional biochemical alternatives.",
            "Large displayed yield differences are model outcomes, not predicted human effect sizes.",
        ],
    }
    output = dict(metadata=metadata, fields=FIELDS, angles=ANGLES, lifetimes=LIFETIMES,
                  relaxations=RELAXATIONS, cases=cases, verification=verification)
    args.output.parent.mkdir(parents=True, exist_ok=True)
    temporary = args.output.with_suffix(args.output.suffix + ".tmp")
    temporary.write_text(json.dumps(output, indent=2, allow_nan=False) + "\n", encoding="utf-8")
    temporary.replace(args.output)
    print(json.dumps({"output": str(args.output.resolve()), "status": verification["status"],
                      "points": verification["grid"]["grid_point_count"],
                      "grid": verification["grid"], "symmetries": verification["symmetries"],
                      "time_integration_max_yield_error": verification["independent_time_integration"]["max_yield_difference"],
                      "time_integration_convergence_error": verification["independent_time_integration"]["max_tolerance_convergence_difference"],
                      "time_integration_min_eigenvalue": verification["independent_time_integration"]["min_sampled_normalized_state_eigenvalue"],
                      "runtime_seconds": verification["runtime_seconds"]}, indent=2))


if __name__ == "__main__":
    main()
